1use anyhow::{Result, anyhow};
2
3const DEFAULT_MODEL: &str = "copilot";
4const DEFAULT_BASE_URL: &str = "https://api.openai.com/v1";
5
6#[derive(Clone, Debug)]
8pub struct LlmConfig {
9 pub api_key: String,
11 pub model: String,
13 pub base_url: String,
15}
16
17pub fn resolve_llm_config(model: Option<&str>, base_url: Option<&str>) -> Result<LlmConfig> {
21 let api_key = super::optional_env("LLM_API_KEY")
22 .or_else(|| super::optional_env("OPENAI_API_KEY"))
23 .ok_or_else(|| anyhow!("Missing environment variable LLM_API_KEY. Please configure it in .env."))?;
24
25 let resolved_model = model
26 .map(|s| s.to_string())
27 .or_else(|| super::optional_env("LLM_MODEL"))
28 .or_else(|| super::optional_env("OPENAI_MODEL"))
29 .unwrap_or_else(|| DEFAULT_MODEL.to_string());
30
31 let resolved_base_url = base_url
32 .map(|s| s.to_string())
33 .or_else(|| super::optional_env("LLM_BASE_URL"))
34 .or_else(|| super::optional_env("OPENAI_BASE_URL"))
35 .unwrap_or_else(|| DEFAULT_BASE_URL.to_string());
36
37 Ok(LlmConfig { api_key, model: resolved_model, base_url: resolved_base_url })
38}
39
40#[cfg(test)]
41mod tests {
42 use super::*;
43
44 struct EnvGuard {
45 keys: Vec<&'static str>,
46 saved: Vec<Option<String>>,
47 }
48
49 impl EnvGuard {
50 fn new(keys: &[&'static str]) -> Self {
51 let saved: Vec<Option<String>> = keys.iter().map(|k| std::env::var(k).ok()).collect();
52 for k in keys {
53 unsafe { std::env::remove_var(k) };
54 }
55 Self { keys: keys.to_vec(), saved }
56 }
57 }
58
59 impl Drop for EnvGuard {
60 fn drop(&mut self) {
61 for (i, k) in self.keys.iter().enumerate() {
62 unsafe { std::env::remove_var(k) };
63 if let Some(ref v) = self.saved[i] {
64 unsafe { std::env::set_var(k, v) };
65 }
66 }
67 }
68 }
69
70 #[test]
72 fn test_env_var_resolution_chain() {
73 let vars = &[
74 "LLM_API_KEY", "OPENAI_API_KEY", "LLM_MODEL", "OPENAI_MODEL",
75 "LLM_BASE_URL", "OPENAI_BASE_URL",
76 ];
77 let _guard = EnvGuard::new(vars);
78
79 let set = |k: &str, v: &str| unsafe { std::env::set_var(k, v) };
80 let rm = |k: &str| unsafe { std::env::remove_var(k) };
81
82 assert!(resolve_llm_config(None, None).is_err());
84
85 set("LLM_API_KEY", "sk-llm");
87 let cfg = resolve_llm_config(None, None).unwrap();
88 assert_eq!(cfg.api_key, "sk-llm");
89 assert_eq!(cfg.model, DEFAULT_MODEL);
90 rm("LLM_API_KEY");
91
92 set("OPENAI_API_KEY", "sk-openai");
94 let cfg = resolve_llm_config(None, None).unwrap();
95 assert_eq!(cfg.api_key, "sk-openai");
96 rm("OPENAI_API_KEY");
97
98 set("LLM_API_KEY", "sk-llm");
100 set("OPENAI_API_KEY", "sk-openai");
101 let cfg = resolve_llm_config(None, None).unwrap();
102 assert_eq!(cfg.api_key, "sk-llm");
103 rm("LLM_API_KEY");
104 rm("OPENAI_API_KEY");
105
106 set("LLM_API_KEY", "sk-test");
108 set("LLM_MODEL", "env-model");
109 let cfg = resolve_llm_config(Some("cli-model"), None).unwrap();
110 assert_eq!(cfg.model, "cli-model");
111 rm("LLM_MODEL");
112
113 set("LLM_MODEL", "gpt-4");
115 let cfg = resolve_llm_config(None, None).unwrap();
116 assert_eq!(cfg.model, "gpt-4");
117 rm("LLM_MODEL");
118
119 set("OPENAI_MODEL", "gpt-3.5");
121 let cfg = resolve_llm_config(None, None).unwrap();
122 assert_eq!(cfg.model, "gpt-3.5");
123 rm("OPENAI_MODEL");
124
125 let cfg = resolve_llm_config(None, None).unwrap();
127 assert_eq!(cfg.model, DEFAULT_MODEL);
128
129 set("LLM_BASE_URL", "https://env.example.com/v1");
131 let cfg = resolve_llm_config(None, Some("https://cli.example.com/v1")).unwrap();
132 assert_eq!(cfg.base_url, "https://cli.example.com/v1");
133 rm("LLM_BASE_URL");
134
135 set("LLM_BASE_URL", "https://llm.example.com/v1");
137 let cfg = resolve_llm_config(None, None).unwrap();
138 assert_eq!(cfg.base_url, "https://llm.example.com/v1");
139 rm("LLM_BASE_URL");
140
141 let cfg = resolve_llm_config(None, None).unwrap();
143 assert_eq!(cfg.base_url, DEFAULT_BASE_URL);
144
145 set("LLM_MODEL", "");
147 let cfg = resolve_llm_config(None, None).unwrap();
148 assert_eq!(cfg.model, DEFAULT_MODEL);
149 rm("LLM_MODEL");
150
151 set("OPENAI_BASE_URL", "https://openai.example.com/v1");
153 let cfg = resolve_llm_config(None, None).unwrap();
154 assert_eq!(cfg.base_url, "https://openai.example.com/v1");
155 rm("OPENAI_BASE_URL");
156
157 set("LLM_API_KEY", "sk-test");
158 }
160
161 #[test]
162 fn test_llm_config_debug_clone() {
163 let cfg = LlmConfig {
164 api_key: "sk-test".into(),
165 model: "gpt-4".into(),
166 base_url: "https://api.openai.com/v1".into(),
167 };
168 let cloned = cfg.clone();
169 assert_eq!(cloned.api_key, "sk-test");
170 assert_eq!(cloned.model, "gpt-4");
171 let _ = format!("{:?}", cfg);
173 }
174}