vtcode_llm/providers/openai/
custom_provider_auth.rs1use std::path::PathBuf;
2use std::process::Stdio;
3use std::sync::Arc;
4use std::time::{Duration, Instant};
5
6use anyhow::{Context, Result, anyhow, bail};
7use tokio::process::Command;
8use tokio::sync::Mutex;
9use tokio::time::timeout;
10use vtcode_commons::sanitizer::sanitize_provider_diagnostic;
11use vtcode_config::core::CustomProviderCommandAuthConfig;
12
13use crate::process_env::sanitize_tokio_command_environment;
14
15#[derive(Clone, Debug)]
22pub struct CustomProviderAuthHandle {
23 config: CustomProviderCommandAuthConfig,
24 workspace_root: Option<PathBuf>,
25 state: Arc<Mutex<CustomProviderAuthState>>,
26}
27
28#[derive(Debug, Default)]
29struct CustomProviderAuthState {
30 cached_token: Option<CachedToken>,
31}
32
33#[derive(Clone)]
34struct CachedToken {
35 value: String,
36 fetched_at: Instant,
37}
38
39impl std::fmt::Debug for CachedToken {
40 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
41 f.debug_struct("CachedToken")
42 .field("value", &"<redacted>")
43 .field("fetched_at", &self.fetched_at)
44 .finish()
45 }
46}
47
48impl CustomProviderAuthHandle {
49 pub fn new(config: CustomProviderCommandAuthConfig, workspace_root: Option<PathBuf>) -> Self {
50 Self {
51 config,
52 workspace_root,
53 state: Arc::new(Mutex::new(CustomProviderAuthState::default())),
54 }
55 }
56
57 pub(crate) async fn current_token(&self) -> Result<String> {
58 {
59 let state = self.state.lock().await;
60 if let Some(token) = state.cached_token.as_ref()
61 && token.fetched_at.elapsed() < self.refresh_interval()
62 {
63 return Ok(token.value.clone());
64 }
65 }
66
67 let token = self.fetch_token().await?;
68 let mut state = self.state.lock().await;
69 state.cached_token = Some(CachedToken { value: token.clone(), fetched_at: Instant::now() });
70 Ok(token)
71 }
72
73 pub(crate) async fn force_refresh(&self) -> Result<String> {
74 let token = self.fetch_token().await?;
75 let mut state = self.state.lock().await;
76 state.cached_token = Some(CachedToken { value: token.clone(), fetched_at: Instant::now() });
77 Ok(token)
78 }
79
80 fn refresh_interval(&self) -> Duration {
81 Duration::from_millis(self.config.refresh_interval_ms)
82 }
83
84 fn timeout(&self) -> Duration {
85 Duration::from_millis(self.config.timeout_ms)
86 }
87
88 fn resolve_cwd(&self) -> Option<PathBuf> {
89 let cwd = self.config.cwd.as_ref()?;
90 if cwd.is_absolute() {
91 return Some(cwd.clone());
92 }
93
94 self.workspace_root
95 .as_ref()
96 .map(|workspace_root| workspace_root.join(cwd))
97 .or_else(|| std::env::current_dir().ok().map(|cwd_root| cwd_root.join(cwd)))
98 }
99
100 async fn fetch_token(&self) -> Result<String> {
101 let mut command = Command::new(&self.config.command);
102 sanitize_tokio_command_environment(&mut command, &[]);
103 command
104 .args(&self.config.args)
105 .stdin(Stdio::null())
106 .stdout(Stdio::piped())
107 .stderr(Stdio::piped());
108
109 if let Some(cwd) = self.resolve_cwd() {
110 command.current_dir(cwd);
111 }
112
113 let output = timeout(self.timeout(), command.output())
114 .await
115 .with_context(|| format!("provider auth command timed out after {}ms", self.config.timeout_ms))?
116 .with_context(|| format!("failed to execute provider auth command `{}`", self.config.command))?;
117
118 if !output.status.success() {
119 let stderr = sanitize_provider_diagnostic(&output.stderr);
120 let stderr = stderr.trim();
121 if stderr.is_empty() {
122 bail!("provider auth command `{}` exited with status {}", self.config.command, output.status);
123 }
124 bail!("provider auth command `{}` exited with status {}: {}", self.config.command, output.status, stderr);
125 }
126
127 let stdout = String::from_utf8(output.stdout).map_err(|err| {
128 anyhow!("provider auth command `{}` returned non-utf8 stdout: {err}", self.config.command)
129 })?;
130 let token = stdout.trim();
131 if token.is_empty() {
132 bail!("provider auth command `{}` returned an empty token", self.config.command);
133 }
134
135 Ok(token.to_string())
136 }
137}
138
139#[cfg(test)]
140mod tests {
141 use super::CustomProviderAuthHandle;
142 use std::path::Path;
143 use tempfile::TempDir;
144 use vtcode_config::core::CustomProviderCommandAuthConfig;
145
146 fn write_tokens_file(dir: &Path, tokens: &[&str]) {
147 std::fs::write(dir.join("tokens.txt"), tokens.join("\n")).expect("write tokens file");
148 }
149
150 #[cfg(unix)]
151 fn build_fixture(dir: &TempDir, tokens: &[&str]) -> CustomProviderCommandAuthConfig {
152 use std::os::unix::fs::PermissionsExt;
153
154 write_tokens_file(dir.path(), tokens);
155 let script_path = dir.path().join("print-token.sh");
156 std::fs::write(
157 &script_path,
158 r#"#!/bin/sh
159first_line=$(sed -n '1p' tokens.txt)
160printf ' %s \n' "$first_line"
161tail -n +2 tokens.txt > tokens.next
162mv tokens.next tokens.txt
163"#,
164 )
165 .expect("write script");
166 let mut permissions = std::fs::metadata(&script_path).expect("script metadata").permissions();
167 permissions.set_mode(0o755);
168 std::fs::set_permissions(&script_path, permissions).expect("set permissions");
169
170 CustomProviderCommandAuthConfig {
171 command: "./print-token.sh".to_string(),
172 args: Vec::new(),
173 cwd: Some(dir.path().to_path_buf()),
174 timeout_ms: 1_000,
175 refresh_interval_ms: 60_000,
176 }
177 }
178
179 #[cfg(windows)]
180 fn build_fixture(dir: &TempDir, tokens: &[&str]) -> CustomProviderCommandAuthConfig {
181 write_tokens_file(dir.path(), tokens);
182 let script_path = dir.path().join("print-token.ps1");
183 std::fs::write(
184 &script_path,
185 r#"$lines = Get-Content -Path tokens.txt
186if ($lines.Count -eq 0) { exit 1 }
187Write-Output (" " + $lines[0] + " ")
188$lines | Select-Object -Skip 1 | Set-Content -Path tokens.txt
189"#,
190 )
191 .expect("write script");
192
193 CustomProviderCommandAuthConfig {
194 command: "powershell".to_string(),
195 args: vec![
196 "-NoProfile".to_string(),
197 "-ExecutionPolicy".to_string(),
198 "Bypass".to_string(),
199 "-File".to_string(),
200 script_path.to_string_lossy().into_owned(),
201 ],
202 cwd: Some(dir.path().to_path_buf()),
203 timeout_ms: 1_000,
204 refresh_interval_ms: 60_000,
205 }
206 }
207
208 #[tokio::test]
209 async fn current_token_trims_stdout_and_uses_cache() {
210 let dir = TempDir::new().expect("tempdir");
211 let handle = CustomProviderAuthHandle::new(build_fixture(&dir, &["first", "second"]), None);
212
213 let first = handle.current_token().await.expect("first token");
214 let second = handle.current_token().await.expect("cached token");
215
216 assert_eq!(first, "first");
217 assert_eq!(second, "first");
218 let remaining = std::fs::read_to_string(dir.path().join("tokens.txt")).expect("tokens");
219 assert_eq!(remaining.trim(), "second");
220 }
221
222 #[tokio::test]
223 async fn force_refresh_reruns_command() {
224 let dir = TempDir::new().expect("tempdir");
225 let handle = CustomProviderAuthHandle::new(build_fixture(&dir, &["first", "second"]), None);
226
227 let first = handle.current_token().await.expect("first token");
228 let refreshed = handle.force_refresh().await.expect("refreshed token");
229
230 assert_eq!(first, "first");
231 assert_eq!(refreshed, "second");
232 }
233}