use std::io::{self, IsTerminal, Write};
use std::path::Path;
use anyhow::{Context, Result};
use crate::config::schema::LlmProviderType;
use crate::config::{
create_default_config, load_config, load_default_config_with, USER_CONFIG_FILE,
};
use crate::project::ProjectRoot;
pub fn run(env: bool, config_path: Option<&Path>, root: &ProjectRoot) -> Result<()> {
run_with_io(
env,
config_path,
root,
&crate::config::ensure_global_config_dir()?.0,
std::io::stdin().is_terminal(),
&mut read_stdin_line,
)
}
fn read_stdin_line() -> io::Result<String> {
let mut line = String::new();
std::io::stdin().read_line(&mut line)?;
Ok(line)
}
fn run_with_io(
env: bool,
config_path: Option<&Path>,
root: &ProjectRoot,
global_dir: &Path,
is_tty: bool,
read_line: &mut dyn FnMut() -> io::Result<String>,
) -> Result<()> {
let target = global_dir.join(USER_CONFIG_FILE);
if !target.exists() {
create_default_config(&target)?;
}
let provider_cfg = match config_path {
Some(p) => load_config(p)?,
None => load_default_config_with(root, global_dir)?.1,
};
if provider_cfg.llm.provider == LlmProviderType::Mock {
println!("mock provider 无需 API key(本地模拟,不触网)");
return Ok(());
}
let env_name = &provider_cfg.llm.api_key_env;
if std::env::var(env_name).is_ok() {
println!("已通过环境变量 {} 配置,无需重复设置", env_name);
return Ok(());
}
if env {
let suggested = suggested_env_name(&provider_cfg.llm.provider);
write_field(&target, "api_key_env", suggested)?;
println!(
"已写入环境变量引用 api_key_env = \"{suggested}\"(用户级配置,不随 Git 共享)"
);
println!(
"请设置环境变量 {suggested}(如 export {suggested}=sk-... 或 setx {suggested} sk-...),重启终端后生效"
);
return Ok(());
}
if !is_tty {
println!("{}", guidance_text());
return Ok(());
}
print!("请输入 API key(直接回车取消): ");
io::stdout().flush()?;
let line = read_line()?;
let key = line.trim();
if key.is_empty() {
println!("未输入内容,取消");
return Ok(());
}
write_field(&target, "api_key", key)?;
println!(
"已写入 {} 的 [llm] api_key(用户级配置,不随 Git 共享)",
target.display()
);
Ok(())
}
fn suggested_env_name(provider: &LlmProviderType) -> &'static str {
match provider {
LlmProviderType::Anthropic => "ANTHROPIC_API_KEY",
_ => "DEEPSEEK_API_KEY",
}
}
pub(crate) fn guidance_text() -> String {
[
"当前环境非交互式终端(管道/CI/外部 Agent),无法读取键盘输入。可用方式:",
" 1. 在交互式终端运行 `code-repo-wiki key` 直接输入明文 API key(写入用户级 config.toml,不随 Git 共享)",
" 2. 运行 `code-repo-wiki key --env` 改用环境变量引用(不落明文,key 由 shell 环境提供)",
]
.join("\n")
}
fn write_field(target: &Path, field: &str, value: &str) -> Result<()> {
let text = std::fs::read_to_string(target)
.with_context(|| format!("读取配置失败: {}", target.display()))?;
let updated = set_llm_field(&text, field, value)?;
crate::fs::write_file_atomic(target, &updated)?;
let cfg = load_config(target)
.with_context(|| format!("写入后配置解析失败: {}", target.display()))?;
let effective = if field == "api_key" {
cfg.llm.api_key.as_deref() == Some(value)
} else {
cfg.llm.api_key_env == value
};
if !effective {
anyhow::bail!("验证失败:{field} 写入后未生效,请检查 {}", target.display());
}
Ok(())
}
fn set_llm_field(text: &str, field: &str, value: &str) -> Result<String> {
let escaped = escape_toml_string(value);
let mut lines: Vec<String> = text.lines().map(String::from).collect();
let Some(llm_idx) = lines.iter().position(|l| l.trim() == "[llm]") else {
let mut doc: toml::Value =
toml::from_str(text).with_context(|| "解析配置文本失败".to_string())?;
let llm = doc
.as_table_mut()
.ok_or_else(|| anyhow::anyhow!("配置根不是表"))?
.entry("llm")
.or_insert_with(|| toml::Value::Table(Default::default()))
.as_table_mut()
.ok_or_else(|| anyhow::anyhow!("[llm] 不是表"))?;
llm.insert(field.to_string(), toml::Value::String(value.to_string()));
let out = toml::to_string(&doc).with_context(|| "配置序列化失败".to_string())?;
return Ok(out);
};
let end = lines
.iter()
.enumerate()
.skip(llm_idx + 1)
.find(|(_, l)| {
let t = l.trim();
t.starts_with('[') && t.ends_with(']')
})
.map(|(i, _)| i)
.unwrap_or(lines.len());
let seg = &lines[llm_idx + 1..end];
let matches_field = |l: &str| -> bool {
l.trim_start()
.strip_prefix(field)
.is_some_and(|rest| rest.trim_start().starts_with('='))
};
let matches_comment_field = |l: &str| -> bool {
l.trim_start().strip_prefix('#').is_some_and(matches_field)
};
if let Some((off, line)) = seg.iter().enumerate().find(|(_, l)| matches_field(l)) {
let indent = line[..line.len() - line.trim_start().len()].to_string();
lines[llm_idx + 1 + off] = format!("{indent}{field} = \"{escaped}\"");
return Ok(join_preserving_newline(&lines, text));
}
if let Some((off, line)) = seg
.iter()
.enumerate()
.find(|(_, l)| matches_comment_field(l))
{
let indent = line[..line.len() - line.trim_start().len()].to_string();
lines[llm_idx + 1 + off] = format!("{indent}{field} = \"{escaped}\"");
return Ok(join_preserving_newline(&lines, text));
}
let insert_at = match seg.iter().rposition(|l| !l.trim().is_empty()) {
Some(rel) => llm_idx + 1 + rel + 1,
None => llm_idx + 1, };
lines.insert(insert_at, format!("{field} = \"{escaped}\""));
Ok(join_preserving_newline(&lines, text))
}
fn join_preserving_newline(lines: &[String], original: &str) -> String {
let mut out = lines.join("\n");
if original.ends_with('\n') {
out.push('\n');
}
out
}
fn escape_toml_string(value: &str) -> String {
value.replace('\\', "\\\\").replace('"', "\\\"")
}
#[cfg(test)]
mod tests {
use super::*;
use std::path::PathBuf;
fn temp_global(tag: &str, provider: &str) -> (PathBuf, PathBuf) {
let dir = std::env::temp_dir().join(format!(
"code_repo_wiki_key_{}_{}",
tag,
std::process::id()
));
let _ = std::fs::remove_dir_all(&dir);
std::fs::create_dir_all(&dir).unwrap();
let global_dir = std::env::temp_dir().join(format!(
"code_repo_wiki_key_global_{}_{}",
tag,
std::process::id()
));
let _ = std::fs::remove_dir_all(&global_dir);
std::fs::create_dir_all(&global_dir).unwrap();
let user_text = include_str!("../config.toml")
.lines()
.map(|l| {
if l.starts_with("provider = ") {
format!("provider = \"{provider}\"")
} else if l.contains("api_key_env = \"OPENCODEGO2_API_KEY\"") {
"api_key_env = \"REPO_WIKI_TEST_ENV_NONE\"".to_string()
} else {
l.to_string()
}
})
.collect::<Vec<_>>()
.join("\n");
std::fs::write(global_dir.join(USER_CONFIG_FILE), &user_text).unwrap();
(dir, global_dir)
}
#[test]
fn test_key_writes_plain_api_key_to_user_config() {
let (dir, global_dir) = temp_global("plain", "openai");
let root = ProjectRoot::new(dir.clone());
let mut input = || Ok("sk-test-123".to_string());
run_with_io(false, None, &root, &global_dir, true, &mut input).unwrap();
let target = global_dir.join(USER_CONFIG_FILE);
let text = std::fs::read_to_string(&target).unwrap();
assert!(text.contains("api_key = \"sk-test-123\""), "未写入明文: {text}");
let cfg = load_config(&target).unwrap();
assert_eq!(cfg.llm.api_key.as_deref(), Some("sk-test-123"));
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn test_key_env_mode_writes_env_reference() {
let (dir, global_dir) = temp_global("env", "anthropic");
let root = ProjectRoot::new(dir.clone());
let mut input = || Ok(String::new());
run_with_io(true, None, &root, &global_dir, false, &mut input).unwrap();
let target = global_dir.join(USER_CONFIG_FILE);
let text = std::fs::read_to_string(&target).unwrap();
assert!(
text.contains("api_key_env = \"ANTHROPIC_API_KEY\""),
"--env 应写建议 env 名: {text}"
);
let plain_lines = text
.lines()
.filter(|l| l.trim_start().starts_with("api_key ="))
.count();
assert_eq!(plain_lines, 0, "不应写明文 api_key: {text}");
let cfg = load_config(&target).unwrap();
assert_eq!(cfg.llm.api_key_env, "ANTHROPIC_API_KEY");
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn test_key_no_tty_prints_guidance() {
let g = guidance_text();
assert!(g.contains("非交互"), "应说明非交互: {g}");
assert!(g.contains("--env"), "应引导 --env 模式: {g}");
let (dir, global_dir) = temp_global("tty", "openai");
let root = ProjectRoot::new(dir.clone());
let mut input = || Ok("sk-test-123".to_string());
run_with_io(false, None, &root, &global_dir, false, &mut input).unwrap();
let text = std::fs::read_to_string(global_dir.join(USER_CONFIG_FILE)).unwrap();
assert!(!text.contains("sk-test-123"), "非 TTY 不应写入: {text}");
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn test_set_llm_field_appends_when_no_placeholder() {
let text = "[llm]\nprovider = \"openai\"\napi_key_env = \"X\"\n\n[embed]\nmodel = \"m\"\n";
let out = set_llm_field(text, "api_key", "sk-abc").unwrap();
assert!(
out.contains("api_key_env = \"X\"\napi_key = \"sk-abc\""),
"应追加在段末: {out}"
);
assert!(out.contains("[embed]"), "应保留后续段: {out}");
let v: toml::Value = toml::from_str(&out).unwrap();
assert_eq!(v["llm"]["api_key"].as_str(), Some("sk-abc"));
assert_eq!(v["embed"]["model"].as_str(), Some("m"));
}
#[test]
fn test_set_llm_field_escapes_quotes_and_backslashes() {
let text = "[llm]\nprovider = \"openai\"\n";
let out = set_llm_field(text, "api_key", "sk-a\"b\\c").unwrap();
assert!(out.contains("api_key = \"sk-a\\\"b\\\\c\""), "{out}");
let v: toml::Value = toml::from_str(&out).unwrap();
assert_eq!(v["llm"]["api_key"].as_str(), Some("sk-a\"b\\c"));
}
#[test]
fn test_set_llm_field_falls_back_without_llm_section() {
let text = "[llm]\nprovider = \"mock\"\n";
let out = set_llm_field(text, "api_key", "sk-1").unwrap();
let v: toml::Value = toml::from_str(&out).unwrap();
assert_eq!(v["llm"]["api_key"].as_str(), Some("sk-1"));
}
}