Skip to main content

vtcode_llm/providers/openai/
custom_provider_auth.rs

1use 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// Retained custom auth transport for OpenAI-compatible providers.
16// Rig does not model VT Code's configured command-token provider, refresh
17// interval, workspace-relative cwd, timeout, stdout trimming, or forced refresh
18// on 401. Protected by this module's auth-command tests and
19// `custom_provider_auth_retries_with_refreshed_tokens_after_401`. Remove only
20// if Rig gains an equivalent command-auth hook with retry/cache parity.
21#[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}