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

1"""配置系统 - 支持多层级配置合并 + provider 自动发现""" 

2 

3import os 

4import json 

5import logging 

6from pathlib import Path 

7from dataclasses import dataclass, field 

8 

9logger = logging.getLogger("merco.config") 

10 

11# ── Provider 注册表:内置平台的完整元数据 ── 

12# 新增平台只需加一条 ProviderInfo,setup 向导自动适配。 

13 

14 

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 # 一句话介绍 

26 

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) 

34 

35 

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} 

89 

90 

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 

103 

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 ) 

118 

119 

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) 

153 

154 @classmethod 

155 def load(cls, config_path: str | None = None) -> "MercoConfig": 

156 if config_path is None: 

157 config_path = cls._find_config() 

158 

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() 

165 

166 # 后处理:根据 provider 补全未填字段 

167 cfg.model.resolve() 

168 return cfg 

169 

170 def save(self, config_path: str): 

171 with open(config_path, "w") as f: 

172 json.dump(self._to_dict(), f, indent=2) 

173 

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) 

178 

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 } 

226 

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 ) 

281 

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