1use std::io::{self, IsTerminal, Write};
11use std::path::Path;
12
13use anyhow::{Context, Result};
14
15use crate::config::schema::LlmProviderType;
16use crate::config::{
17 create_default_config, load_config, load_default_config_with, USER_CONFIG_FILE,
18};
19use crate::project::ProjectRoot;
20
21pub fn run(env: bool, config_path: Option<&Path>, root: &ProjectRoot) -> Result<()> {
26 run_with_io(
27 env,
28 config_path,
29 root,
30 &crate::config::ensure_global_config_dir()?.0,
31 std::io::stdin().is_terminal(),
32 &mut read_stdin_line,
33 )
34}
35
36fn read_stdin_line() -> io::Result<String> {
38 let mut line = String::new();
39 std::io::stdin().read_line(&mut line)?;
40 Ok(line)
41}
42
43fn run_with_io(
46 env: bool,
47 config_path: Option<&Path>,
48 root: &ProjectRoot,
49 global_dir: &Path,
50 is_tty: bool,
51 read_line: &mut dyn FnMut() -> io::Result<String>,
52) -> Result<()> {
53 let target = global_dir.join(USER_CONFIG_FILE);
56 if !target.exists() {
57 create_default_config(&target)?;
58 }
59
60 let provider_cfg = match config_path {
64 Some(p) => load_config(p)?,
65 None => load_default_config_with(root, global_dir)?.1,
66 };
67 if provider_cfg.llm.provider == LlmProviderType::Mock {
68 println!("mock provider 无需 API key(本地模拟,不触网)");
69 return Ok(());
70 }
71
72 let env_name = &provider_cfg.llm.api_key_env;
75 if std::env::var(env_name).is_ok() {
76 println!("已通过环境变量 {} 配置,无需重复设置", env_name);
77 return Ok(());
78 }
79
80 if env {
82 let suggested = suggested_env_name(&provider_cfg.llm.provider);
83 write_field(&target, "api_key_env", suggested)?;
84 println!(
85 "已写入环境变量引用 api_key_env = \"{suggested}\"(用户级配置,不随 Git 共享)"
86 );
87 println!(
88 "请设置环境变量 {suggested}(如 export {suggested}=sk-... 或 setx {suggested} sk-...),重启终端后生效"
89 );
90 return Ok(());
91 }
92
93 if !is_tty {
95 println!("{}", guidance_text());
96 return Ok(());
97 }
98
99 print!("请输入 API key(直接回车取消): ");
101 io::stdout().flush()?;
102 let line = read_line()?;
103 let key = line.trim();
104 if key.is_empty() {
105 println!("未输入内容,取消");
106 return Ok(());
107 }
108
109 write_field(&target, "api_key", key)?;
111 println!(
112 "已写入 {} 的 [llm] api_key(用户级配置,不随 Git 共享)",
113 target.display()
114 );
115 Ok(())
116}
117
118fn suggested_env_name(provider: &LlmProviderType) -> &'static str {
122 match provider {
123 LlmProviderType::Anthropic => "ANTHROPIC_API_KEY",
124 _ => "DEEPSEEK_API_KEY",
125 }
126}
127
128pub(crate) fn guidance_text() -> String {
130 [
131 "当前环境非交互式终端(管道/CI/外部 Agent),无法读取键盘输入。可用方式:",
132 " 1. 在交互式终端运行 `code-repo-wiki key` 直接输入明文 API key(写入用户级 config.toml,不随 Git 共享)",
133 " 2. 运行 `code-repo-wiki key --env` 改用环境变量引用(不落明文,key 由 shell 环境提供)",
134 ]
135 .join("\n")
136}
137
138fn write_field(target: &Path, field: &str, value: &str) -> Result<()> {
140 let text = std::fs::read_to_string(target)
141 .with_context(|| format!("读取配置失败: {}", target.display()))?;
142 let updated = set_llm_field(&text, field, value)?;
143 crate::fs::write_file_atomic(target, &updated)?;
144 let cfg = load_config(target)
147 .with_context(|| format!("写入后配置解析失败: {}", target.display()))?;
148 let effective = if field == "api_key" {
149 cfg.llm.api_key.as_deref() == Some(value)
150 } else {
151 cfg.llm.api_key_env == value
152 };
153 if !effective {
154 anyhow::bail!("验证失败:{field} 写入后未生效,请检查 {}", target.display());
155 }
156 Ok(())
157}
158
159fn set_llm_field(text: &str, field: &str, value: &str) -> Result<String> {
165 let escaped = escape_toml_string(value);
166 let mut lines: Vec<String> = text.lines().map(String::from).collect();
167
168 let Some(llm_idx) = lines.iter().position(|l| l.trim() == "[llm]") else {
169 let mut doc: toml::Value =
172 toml::from_str(text).with_context(|| "解析配置文本失败".to_string())?;
173 let llm = doc
174 .as_table_mut()
175 .ok_or_else(|| anyhow::anyhow!("配置根不是表"))?
176 .entry("llm")
177 .or_insert_with(|| toml::Value::Table(Default::default()))
178 .as_table_mut()
179 .ok_or_else(|| anyhow::anyhow!("[llm] 不是表"))?;
180 llm.insert(field.to_string(), toml::Value::String(value.to_string()));
181 let out = toml::to_string(&doc).with_context(|| "配置序列化失败".to_string())?;
182 return Ok(out);
183 };
184
185 let end = lines
188 .iter()
189 .enumerate()
190 .skip(llm_idx + 1)
191 .find(|(_, l)| {
192 let t = l.trim();
193 t.starts_with('[') && t.ends_with(']')
194 })
195 .map(|(i, _)| i)
196 .unwrap_or(lines.len());
197 let seg = &lines[llm_idx + 1..end];
198
199 let matches_field = |l: &str| -> bool {
201 l.trim_start()
202 .strip_prefix(field)
203 .is_some_and(|rest| rest.trim_start().starts_with('='))
204 };
205 let matches_comment_field = |l: &str| -> bool {
207 l.trim_start().strip_prefix('#').is_some_and(matches_field)
208 };
209
210 if let Some((off, line)) = seg.iter().enumerate().find(|(_, l)| matches_field(l)) {
212 let indent = line[..line.len() - line.trim_start().len()].to_string();
213 lines[llm_idx + 1 + off] = format!("{indent}{field} = \"{escaped}\"");
214 return Ok(join_preserving_newline(&lines, text));
215 }
216 if let Some((off, line)) = seg
218 .iter()
219 .enumerate()
220 .find(|(_, l)| matches_comment_field(l))
221 {
222 let indent = line[..line.len() - line.trim_start().len()].to_string();
223 lines[llm_idx + 1 + off] = format!("{indent}{field} = \"{escaped}\"");
224 return Ok(join_preserving_newline(&lines, text));
225 }
226 let insert_at = match seg.iter().rposition(|l| !l.trim().is_empty()) {
228 Some(rel) => llm_idx + 1 + rel + 1,
229 None => llm_idx + 1, };
231 lines.insert(insert_at, format!("{field} = \"{escaped}\""));
232 Ok(join_preserving_newline(&lines, text))
233}
234
235fn join_preserving_newline(lines: &[String], original: &str) -> String {
237 let mut out = lines.join("\n");
238 if original.ends_with('\n') {
239 out.push('\n');
240 }
241 out
242}
243
244fn escape_toml_string(value: &str) -> String {
247 value.replace('\\', "\\\\").replace('"', "\\\"")
248}
249
250#[cfg(test)]
251mod tests {
252 use super::*;
253 use std::path::PathBuf;
254
255 fn temp_global(tag: &str, provider: &str) -> (PathBuf, PathBuf) {
265 let dir = std::env::temp_dir().join(format!(
266 "code_repo_wiki_key_{}_{}",
267 tag,
268 std::process::id()
269 ));
270 let _ = std::fs::remove_dir_all(&dir);
271 std::fs::create_dir_all(&dir).unwrap();
272 let global_dir = std::env::temp_dir().join(format!(
277 "code_repo_wiki_key_global_{}_{}",
278 tag,
279 std::process::id()
280 ));
281 let _ = std::fs::remove_dir_all(&global_dir);
282 std::fs::create_dir_all(&global_dir).unwrap();
283 let user_text = include_str!("../config.toml")
284 .lines()
285 .map(|l| {
286 if l.starts_with("provider = ") {
287 format!("provider = \"{provider}\"")
288 } else if l.contains("api_key_env = \"OPENCODEGO2_API_KEY\"") {
289 "api_key_env = \"REPO_WIKI_TEST_ENV_NONE\"".to_string()
290 } else {
291 l.to_string()
292 }
293 })
294 .collect::<Vec<_>>()
295 .join("\n");
296 std::fs::write(global_dir.join(USER_CONFIG_FILE), &user_text).unwrap();
297 (dir, global_dir)
298 }
299
300 #[test]
302 fn test_key_writes_plain_api_key_to_user_config() {
303 let (dir, global_dir) = temp_global("plain", "openai");
304 let root = ProjectRoot::new(dir.clone());
305 let mut input = || Ok("sk-test-123".to_string());
306 run_with_io(false, None, &root, &global_dir, true, &mut input).unwrap();
307
308 let target = global_dir.join(USER_CONFIG_FILE);
309 let text = std::fs::read_to_string(&target).unwrap();
310 assert!(text.contains("api_key = \"sk-test-123\""), "未写入明文: {text}");
312 let cfg = load_config(&target).unwrap();
314 assert_eq!(cfg.llm.api_key.as_deref(), Some("sk-test-123"));
315 let _ = std::fs::remove_dir_all(&dir);
316 }
317
318 #[test]
320 fn test_key_env_mode_writes_env_reference() {
321 let (dir, global_dir) = temp_global("env", "anthropic");
322 let root = ProjectRoot::new(dir.clone());
323 let mut input = || Ok(String::new());
325 run_with_io(true, None, &root, &global_dir, false, &mut input).unwrap();
326
327 let target = global_dir.join(USER_CONFIG_FILE);
328 let text = std::fs::read_to_string(&target).unwrap();
329 assert!(
331 text.contains("api_key_env = \"ANTHROPIC_API_KEY\""),
332 "--env 应写建议 env 名: {text}"
333 );
334 let plain_lines = text
336 .lines()
337 .filter(|l| l.trim_start().starts_with("api_key ="))
338 .count();
339 assert_eq!(plain_lines, 0, "不应写明文 api_key: {text}");
340 let cfg = load_config(&target).unwrap();
341 assert_eq!(cfg.llm.api_key_env, "ANTHROPIC_API_KEY");
342 let _ = std::fs::remove_dir_all(&dir);
343 }
344
345 #[test]
348 fn test_key_no_tty_prints_guidance() {
349 let g = guidance_text();
350 assert!(g.contains("非交互"), "应说明非交互: {g}");
351 assert!(g.contains("--env"), "应引导 --env 模式: {g}");
352
353 let (dir, global_dir) = temp_global("tty", "openai");
354 let root = ProjectRoot::new(dir.clone());
355 let mut input = || Ok("sk-test-123".to_string());
356 run_with_io(false, None, &root, &global_dir, false, &mut input).unwrap();
357 let text = std::fs::read_to_string(global_dir.join(USER_CONFIG_FILE)).unwrap();
358 assert!(!text.contains("sk-test-123"), "非 TTY 不应写入: {text}");
359 let _ = std::fs::remove_dir_all(&dir);
360 }
361
362 #[test]
364 fn test_set_llm_field_appends_when_no_placeholder() {
365 let text = "[llm]\nprovider = \"openai\"\napi_key_env = \"X\"\n\n[embed]\nmodel = \"m\"\n";
366 let out = set_llm_field(text, "api_key", "sk-abc").unwrap();
367 assert!(
368 out.contains("api_key_env = \"X\"\napi_key = \"sk-abc\""),
369 "应追加在段末: {out}"
370 );
371 assert!(out.contains("[embed]"), "应保留后续段: {out}");
372 let v: toml::Value = toml::from_str(&out).unwrap();
373 assert_eq!(v["llm"]["api_key"].as_str(), Some("sk-abc"));
374 assert_eq!(v["embed"]["model"].as_str(), Some("m"));
375 }
376
377 #[test]
379 fn test_set_llm_field_escapes_quotes_and_backslashes() {
380 let text = "[llm]\nprovider = \"openai\"\n";
381 let out = set_llm_field(text, "api_key", "sk-a\"b\\c").unwrap();
382 assert!(out.contains("api_key = \"sk-a\\\"b\\\\c\""), "{out}");
383 let v: toml::Value = toml::from_str(&out).unwrap();
384 assert_eq!(v["llm"]["api_key"].as_str(), Some("sk-a\"b\\c"));
385 }
386
387 #[test]
389 fn test_set_llm_field_falls_back_without_llm_section() {
390 let text = "[llm]\nprovider = \"mock\"\n";
391 let out = set_llm_field(text, "api_key", "sk-1").unwrap();
392 let v: toml::Value = toml::from_str(&out).unwrap();
393 assert_eq!(v["llm"]["api_key"].as_str(), Some("sk-1"));
394 }
395}