Coverage for merco/core/config.py: 84%
115 statements
« prev ^ index » next coverage.py v7.15.0, created at 2026-07-07 14:04 +0800
« prev ^ index » next coverage.py v7.15.0, created at 2026-07-07 14:04 +0800
1"""配置系统 - 支持多层级配置合并 + provider 自动发现"""
3import os
4import json
5import logging
6from pathlib import Path
7from dataclasses import dataclass, field
9logger = logging.getLogger("merco.config")
11# ── Provider 注册表:内置平台的完整元数据 ──
12# 新增平台只需加一条 ProviderInfo,setup 向导自动适配。
15@dataclass
16class ProviderInfo:
17 """平台元数据 — 一条记录即可驱动配置向导和自动补全"""
18 key: str # provider id: "openai", "minimax", ...
19 name: str # 显示名: "OpenAI", "MiniMax"
20 base_url: str # 默认 API 端点
21 key_env: str # 环境变量名
22 default_model: str # 推荐模型
23 models: list[str] # 已知模型列表(空 = 用户自行输入)
24 key_help: str # 获取 API key 的链接
25 description: str # 一句话介绍
27 # 向后兼容:支持 dict-style 访问(旧代码用 PROVIDER_REGISTRY["openai"]["base_url"])
28 def __getitem__(self, key: str):
29 if key == "base_url":
30 return self.base_url
31 if key == "key_env":
32 return self.key_env
33 raise KeyError(key)
36PROVIDER_REGISTRY: dict[str, ProviderInfo] = {
37 "openai": ProviderInfo(
38 key="openai",
39 name="OpenAI",
40 base_url="https://api.openai.com/v1",
41 key_env="OPENAI_API_KEY",
42 default_model="gpt-4o",
43 models=["gpt-4o", "gpt-4o-mini", "gpt-4-turbo", "o3-mini", "o1"],
44 key_help="https://platform.openai.com/api-keys",
45 description="最通用的平台,GPT-4o / o3 系列",
46 ),
47 "minimax": ProviderInfo(
48 key="minimax",
49 name="MiniMax",
50 base_url="https://api.minimaxi.com/v1",
51 key_env="MINIMAX_API_KEY",
52 default_model="MiniMax-M2.7",
53 models=["MiniMax-M2.7", "MiniMax-Text-01", "abab7-chat"],
54 key_help="https://platform.minimaxi.com/user-center/basic-information",
55 description="国产平台,MiniMax-M2.7 性价比高",
56 ),
57 "anthropic": ProviderInfo(
58 key="anthropic",
59 name="Anthropic",
60 base_url="https://api.anthropic.com",
61 key_env="ANTHROPIC_API_KEY",
62 default_model="claude-sonnet-4-20250514",
63 models=["claude-sonnet-4-20250514", "claude-3-5-haiku-20241022",
64 "claude-3-opus-20240229", "claude-3-5-sonnet-20241022"],
65 key_help="https://console.anthropic.com/settings/keys",
66 description="Claude 系列,代码能力优秀",
67 ),
68 "openrouter": ProviderInfo(
69 key="openrouter",
70 name="OpenRouter",
71 base_url="https://openrouter.ai/api/v1",
72 key_env="OPENROUTER_API_KEY",
73 default_model="anthropic/claude-sonnet-4",
74 models=[], # 模型太多,用户自行输入
75 key_help="https://openrouter.ai/keys",
76 description="模型聚合平台,一个 key 调用上百种模型",
77 ),
78 "deepseek": ProviderInfo(
79 key="deepseek",
80 name="DeepSeek",
81 base_url="https://api.deepseek.com/v1",
82 key_env="DEEPSEEK_API_KEY",
83 default_model="deepseek-chat",
84 models=["deepseek-chat", "deepseek-reasoner"],
85 key_help="https://platform.deepseek.com/api_keys",
86 description="国产平台,deepseek-reasoner 推理能力强",
87 ),
88}
91@dataclass
92class ModelConfig:
93 """模型配置"""
94 provider: str = "openai"
95 model: str = "gpt-4"
96 api_key: str | None = None
97 base_url: str | None = None
98 temperature: float = 0.7
99 max_tokens: int = 4096
100 extra_params: dict = field(default_factory=dict)
101 headers: dict = field(default_factory=dict)
102 stream_options: dict | None = None
104 def resolve(self):
105 """后处理:根据 provider 名补齐未填的字段"""
106 entry = PROVIDER_REGISTRY.get(self.provider)
107 if entry:
108 if not self.base_url:
109 self.base_url = entry.base_url
110 if not self.api_key:
111 self.api_key = os.environ.get(entry.key_env, "")
112 elif not self.base_url:
113 # 未收录的 provider —— 用户必须显式写 base_url
114 logger.warning(
115 "Provider '%s' 未在注册表中。请确认 base_url 已正确填写,"
116 "或使用 provider: 'custom' 明确标注。", self.provider
117 )
120@dataclass
121class MercoConfig:
122 """主配置类"""
123 username: str = "user"
124 model: ModelConfig = field(default_factory=ModelConfig)
125 max_tool_calls: int = 50
126 max_input_tokens: int = 64000
127 compression_threshold: float = 0.75
128 skills_paths: list = field(default_factory=lambda: ["./.merco/skills", "~/.config/merco/skills"])
129 memory_enabled: bool = True
130 memory_path: str = "~/.merco/memory"
131 memory_backend: str = "json"
132 log_level: str = "INFO"
133 sandbox_mode: str = "ask"
134 sandbox_rules: list = field(default_factory=list)
135 streaming: bool = False
136 stream_thinking: bool = True
137 stream_content: bool = True
138 stream_thinking_transient: bool = False # 流式思考内容是否仅临时显示(不保留思考面板)
139 stream_render_interval: float = 0.3 # 流式 reasoning 面板最小渲染间隔(秒),0=不限制
140 diff_view: str = "unified"
141 memory_recall_enabled: bool = True
142 memory_recall_limit: int = 3
143 memory_recall_max_chars: int = 300
144 memory_recall_threshold: float = 0.0
145 memory_auto_extract_on_session_end: bool = False
146 memory_extract_max_per_session: int = 3
147 memory_extract_min_messages: int = 5
148 fork_enabled: bool = True
149 fork_auto_on_compress: bool = True
150 fork_reset_observer: bool = False # default: inherit observer acc on fork
151 mcp_servers: dict = field(default_factory=dict)
152 plugins: dict = field(default_factory=dict)
154 @classmethod
155 def load(cls, config_path: str | None = None) -> "MercoConfig":
156 if config_path is None:
157 config_path = cls._find_config()
159 if config_path and Path(config_path).exists():
160 with open(config_path) as f:
161 data = json.load(f)
162 cfg = cls._from_dict(data)
163 else:
164 cfg = cls()
166 # 后处理:根据 provider 补全未填字段
167 cfg.model.resolve()
168 return cfg
170 def save(self, config_path: str):
171 with open(config_path, "w") as f:
172 json.dump(self._to_dict(), f, indent=2)
174 def merge(self, other: "MercoConfig"):
175 for key, value in vars(other).items():
176 if value is not None:
177 setattr(self, key, value)
179 def _to_dict(self) -> dict:
180 return {
181 "username": self.username,
182 "model": {
183 "provider": self.model.provider,
184 "model": self.model.model,
185 "api_key": self.model.api_key,
186 "base_url": self.model.base_url,
187 "temperature": self.model.temperature,
188 "max_tokens": self.model.max_tokens,
189 "extra_params": self.model.extra_params or None,
190 "headers": self.model.headers or None,
191 "stream_options": self.model.stream_options,
192 },
193 "max_tool_calls": self.max_tool_calls,
194 "max_input_tokens": self.max_input_tokens,
195 "compression_threshold": self.compression_threshold,
196 "skills_paths": self.skills_paths,
197 "log_level": self.log_level,
198 "sandbox_mode": self.sandbox_mode,
199 "sandbox_rules": self.sandbox_rules,
200 "streaming": self.streaming,
201 "stream_thinking": self.stream_thinking,
202 "stream_content": self.stream_content,
203 "stream_thinking_transient": self.stream_thinking_transient,
204 "stream_render_interval": self.stream_render_interval,
205 "diff_view": self.diff_view,
206 "memory": {
207 "enabled": self.memory_enabled,
208 "path": self.memory_path,
209 "backend": self.memory_backend,
210 "recall_enabled": self.memory_recall_enabled,
211 "recall_limit": self.memory_recall_limit,
212 "recall_max_chars": self.memory_recall_max_chars,
213 "recall_threshold": self.memory_recall_threshold,
214 "auto_extract_on_session_end": self.memory_auto_extract_on_session_end,
215 "extract_max_per_session": self.memory_extract_max_per_session,
216 "extract_min_messages": self.memory_extract_min_messages,
217 },
218 "session": {
219 "fork_enabled": self.fork_enabled,
220 "fork_auto_on_compress": self.fork_auto_on_compress,
221 "fork_reset_observer": self.fork_reset_observer,
222 },
223 "mcp_servers": self.mcp_servers,
224 "plugins": self.plugins,
225 }
227 @classmethod
228 def _from_dict(cls, data: dict) -> "MercoConfig":
229 model_data = data.get("model", {})
230 if not isinstance(model_data, dict):
231 model_data = {}
232 model = ModelConfig(
233 provider=model_data.get("provider", "openai"),
234 model=model_data.get("model", "gpt-4"),
235 api_key=model_data.get("api_key"),
236 base_url=model_data.get("base_url"),
237 temperature=model_data.get("temperature", 0.7),
238 max_tokens=model_data.get("max_tokens", 4096),
239 extra_params=model_data.get("extra_params", {}),
240 headers=model_data.get("headers", {}),
241 stream_options=model_data.get("stream_options"),
242 )
243 memory_data = data.get("memory", {})
244 if not isinstance(memory_data, dict):
245 memory_data = {}
246 sess = data.get("session", {})
247 if not isinstance(sess, dict):
248 sess = {}
249 return cls(
250 username=data.get("username", "user"),
251 model=model,
252 max_tool_calls=data.get("max_tool_calls", 50),
253 max_input_tokens=data.get("max_input_tokens", 64000),
254 compression_threshold=data.get("compression_threshold", 0.75),
255 skills_paths=data.get("skills_paths", ["./.merco/skills", "~/.config/merco/skills"]),
256 memory_enabled=memory_data.get("enabled", data.get("memory_enabled", True)),
257 memory_path=memory_data.get("path", data.get("memory_path", "~/.merco/memory")),
258 memory_backend=memory_data.get("backend", "json"),
259 log_level=data.get("log_level", "INFO"),
260 sandbox_mode=data.get("sandbox_mode", "ask"),
261 sandbox_rules=data.get("sandbox_rules", []),
262 streaming=data.get("streaming", False),
263 stream_thinking=data.get("stream_thinking", True),
264 stream_content=data.get("stream_content", True),
265 stream_thinking_transient=data.get("stream_thinking_transient", False),
266 stream_render_interval=data.get("stream_render_interval", 0.05),
267 diff_view=data.get("diff_view", "unified"),
268 memory_recall_enabled=memory_data.get("recall_enabled", True),
269 memory_recall_limit=memory_data.get("recall_limit", 3),
270 memory_recall_max_chars=memory_data.get("recall_max_chars", 300),
271 memory_recall_threshold=memory_data.get("recall_threshold", 0.0),
272 memory_auto_extract_on_session_end=memory_data.get("auto_extract_on_session_end", False),
273 memory_extract_max_per_session=memory_data.get("extract_max_per_session", 3),
274 memory_extract_min_messages=memory_data.get("extract_min_messages", 5),
275 fork_enabled=sess.get("fork_enabled", True),
276 fork_auto_on_compress=sess.get("fork_auto_on_compress", True),
277 fork_reset_observer=isinstance(sess, dict) and sess.get("fork_reset_observer", False),
278 mcp_servers=data.get("mcp_servers", {}),
279 plugins=data.get("plugins", {}),
280 )
282 @staticmethod
283 def _find_config() -> str | None:
284 candidates = [
285 "./merco.json",
286 "./.merco/merco.json",
287 os.path.expanduser("~/.config/merco/config.json"),
288 ]
289 for path in candidates:
290 if Path(path).exists():
291 return path
292 return None