Skip to main content

vtcode_llm/providers/lmstudio/
client.rs

1/// High-level LM Studio client for server interaction and model management.
2///
3/// Supports both LM Studio native REST API (`/api/v0/*`) and OpenAI-compatible
4/// endpoints (`/v1/*`). The native API provides enhanced features like model
5/// management (load/unload), richer model metadata, and TTL-based auto-evict.
6///
7/// See: https://lmstudio.ai/docs/developer
8use std::io;
9use std::path::Path;
10use std::time::Duration;
11
12use serde_json::Value as JsonValue;
13
14use crate::process_env::sanitize_std_command_environment;
15
16pub(crate) const LMSTUDIO_CONNECTION_ERROR: &str =
17    "LM Studio is not responding. Install from https://lmstudio.ai/download and run 'lms server start'.";
18
19/// Client for interacting with a local LM Studio instance.
20///
21/// Supports both native REST API (`/api/v0/*`) and OpenAI-compatible endpoints (`/v1/*`).
22#[derive(Clone, Debug)]
23pub struct LMStudioClient {
24    client: reqwest::Client,
25    base_url: String,
26    /// Use native REST API endpoints (default: false, uses OpenAI-compatible endpoints)
27    use_native_api: bool,
28}
29
30impl LMStudioClient {
31    /// Create a client from a base URL and verify the server is reachable.
32    async fn try_from_base_url(base_url: &str) -> io::Result<Self> {
33        Self::try_from_base_url_with_api_version(base_url, false).await
34    }
35
36    /// Create a client with explicit API version selection.
37    ///
38    /// - `use_native_api = false`: Use OpenAI-compatible endpoints at `/v1/*` (default)
39    /// - `use_native_api = true`: Use native REST API at `/api/v0/*`
40    async fn try_from_base_url_with_api_version(base_url: &str, use_native_api: bool) -> io::Result<Self> {
41        let client = reqwest::Client::builder()
42            .connect_timeout(Duration::from_secs(5))
43            .build()
44            .unwrap_or_else(|_| reqwest::Client::new());
45
46        let instance = Self {
47            client,
48            base_url: base_url.to_string(),
49            use_native_api,
50        };
51
52        instance.check_server().await?;
53        Ok(instance)
54    }
55
56    /// Get the models endpoint URL based on API version.
57    fn models_endpoint(&self) -> String {
58        let base = self.base_url.trim_end_matches('/');
59        if self.use_native_api {
60            format!("{base}/api/v0/models")
61        } else {
62            format!("{base}/v1/models")
63        }
64    }
65
66    /// Verify that the server is reachable.
67    async fn check_server(&self) -> io::Result<()> {
68        let url = self.models_endpoint();
69        let response = self.client.get(&url).send().await;
70
71        if let Ok(resp) = response {
72            if resp.status().is_success() {
73                Ok(())
74            } else {
75                Err(io::Error::other(format!("Server returned error: {} {LMSTUDIO_CONNECTION_ERROR}", resp.status())))
76            }
77        } else {
78            Err(io::Error::other(LMSTUDIO_CONNECTION_ERROR))
79        }
80    }
81
82    /// Fetch the list of model IDs available on the server.
83    async fn fetch_models(&self) -> io::Result<Vec<String>> {
84        let url = self.models_endpoint();
85        let response = self
86            .client
87            .get(&url)
88            .send()
89            .await
90            .map_err(|e| io::Error::other(format!("Request failed: {e}")))?;
91
92        if response.status().is_success() {
93            let json: JsonValue = response
94                .json()
95                .await
96                .map_err(|e| io::Error::new(io::ErrorKind::InvalidData, format!("JSON parse error: {e}")))?;
97
98            let models = json["data"]
99                .as_array()
100                .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "No 'data' array in response"))?
101                .iter()
102                .filter_map(|model| model["id"].as_str())
103                .map(ToString::to_string)
104                .collect();
105
106            Ok(models)
107        } else {
108            Err(io::Error::other(format!("Failed to fetch models: {}", response.status())))
109        }
110    }
111
112    /// Load a model into memory (pre-loads for faster inference).
113    ///
114    /// Uses native REST API `/api/v0/models/load` when `use_native_api` is true,
115    /// otherwise sends a minimal request via `/v1/chat/completions`.
116    pub async fn load_model(&self, model: &str) -> io::Result<()> {
117        if self.use_native_api {
118            let url = format!("{}/api/v0/models/load", self.base_url.trim_end_matches('/'));
119            let request_body = serde_json::json!({
120                "model": model
121            });
122
123            let response = self
124                .client
125                .post(&url)
126                .header("Content-Type", "application/json")
127                .json(&request_body)
128                .send()
129                .await
130                .map_err(|e| io::Error::other(format!("Request failed: {e}")))?;
131
132            if response.status().is_success() {
133                tracing::info!("Successfully loaded model '{model}' via native API");
134                Ok(())
135            } else {
136                Err(io::Error::other(format!("Failed to load model: {}", response.status())))
137            }
138        } else {
139            // Use OpenAI-compatible endpoint with minimal chat completion
140            let url = format!("{}/v1/chat/completions", self.base_url.trim_end_matches('/'));
141            let request_body = serde_json::json!({
142                "model": model,
143                "messages": [{"role": "user", "content": "hi"}],
144                "max_tokens": 1
145            });
146
147            let response = self
148                .client
149                .post(&url)
150                .header("Content-Type", "application/json")
151                .json(&request_body)
152                .send()
153                .await
154                .map_err(|e| io::Error::other(format!("Request failed: {e}")))?;
155
156            if response.status().is_success() {
157                tracing::info!("Successfully loaded model '{model}'");
158                Ok(())
159            } else {
160                Err(io::Error::other(format!("Failed to load model: {}", response.status())))
161            }
162        }
163    }
164
165    /// Unload a model from memory (native REST API only).
166    ///
167    /// This endpoint requires `use_native_api = true`.
168    pub async fn unload_model(&self, model: &str) -> io::Result<()> {
169        if !self.use_native_api {
170            return Err(io::Error::other("Model unload requires native API (use_native_api = true)"));
171        }
172
173        let url = format!("{}/api/v0/models/unload", self.base_url.trim_end_matches('/'));
174        let request_body = serde_json::json!({
175            "model": model
176        });
177
178        let response = self
179            .client
180            .post(&url)
181            .header("Content-Type", "application/json")
182            .json(&request_body)
183            .send()
184            .await
185            .map_err(|e| io::Error::other(format!("Request failed: {e}")))?;
186
187        if response.status().is_success() {
188            tracing::info!("Successfully unloaded model '{model}'");
189            Ok(())
190        } else {
191            Err(io::Error::other(format!("Failed to unload model: {}", response.status())))
192        }
193    }
194
195    /// Find the `lms` CLI tool, checking PATH and fallback locations.
196    fn find_lms() -> io::Result<String> {
197        Self::find_lms_with_home_dir(None)
198    }
199
200    /// Find `lms` CLI with an optional home directory override (for testing).
201    fn find_lms_with_home_dir(home_dir: Option<&str>) -> io::Result<String> {
202        // First try 'lms' in PATH
203        if which::which("lms").is_ok() {
204            return Ok("lms".to_string());
205        }
206
207        // Platform-specific fallback paths
208        let home = match home_dir {
209            Some(dir) => dir.to_string(),
210            None => {
211                #[cfg(unix)]
212                {
213                    std::env::var("HOME").unwrap_or_default()
214                }
215                #[cfg(windows)]
216                {
217                    std::env::var("USERPROFILE").unwrap_or_default()
218                }
219            }
220        };
221
222        #[cfg(unix)]
223        let fallback_path = format!("{home}/.lmstudio/bin/lms");
224        #[cfg(windows)]
225        let fallback_path = format!("{home}/.lmstudio/bin/lms.exe");
226
227        if Path::new(&fallback_path).exists() {
228            Ok(fallback_path)
229        } else {
230            Err(io::Error::new(
231                io::ErrorKind::NotFound,
232                "LM Studio not found. Please install LM Studio from https://lmstudio.ai/",
233            ))
234        }
235    }
236
237    /// Download a model using the `lms` CLI tool.
238    pub async fn download_model(&self, model: &str) -> io::Result<()> {
239        let lms = Self::find_lms()?;
240        tracing::info!(model, "downloading model");
241
242        // `Command::status()` blocks until the `lms get` subprocess exits; run
243        // it off the async executor so it cannot stall Tokio workers. See `# Blocking`
244        // docs in `src/agent/runloop/git.rs`.
245        let model = model.to_string();
246        let lms_for_task = lms.clone();
247        let model_for_task = model.clone();
248        let status = tokio::task::spawn_blocking(move || {
249            let mut command = std::process::Command::new(&lms_for_task);
250            sanitize_std_command_environment(&mut command, &[]);
251            command
252                .args(["get", "--yes", &model_for_task])
253                .stdout(std::process::Stdio::inherit())
254                .stderr(std::process::Stdio::null())
255                .status()
256        })
257        .await
258        .map_err(|e| io::Error::other(format!("Model download task panicked: {e}")))?
259        .map_err(|e| io::Error::other(format!("Failed to execute '{lms} get --yes {model}': {e}")))?;
260
261        if !status.success() {
262            return Err(io::Error::other(format!(
263                "Model download failed with exit code: {}",
264                status.code().unwrap_or(-1)
265            )));
266        }
267
268        tracing::info!("Successfully downloaded model '{model}'");
269        Ok(())
270    }
271}
272
273#[cfg(test)]
274mod tests {
275    use super::*;
276
277    use crate::providers::test_support::start_mock_server_or_skip;
278
279    #[test]
280    fn test_find_lms() {
281        let result = LMStudioClient::find_lms();
282        match result {
283            Ok(_) => {
284                // lms was found in PATH - that's fine
285            }
286            Err(e) => {
287                // Expected error when LM Studio not installed
288                assert!(e.to_string().contains("LM Studio not found"));
289            }
290        }
291    }
292
293    #[test]
294    fn test_find_lms_with_mock_home() {
295        // Test fallback path construction without touching env vars
296        #[cfg(unix)]
297        {
298            let result = LMStudioClient::find_lms_with_home_dir(Some("/test/home"));
299            if let Err(e) = result {
300                assert!(e.to_string().contains("LM Studio not found"));
301            }
302        }
303        #[cfg(windows)]
304        {
305            let result = LMStudioClient::find_lms_with_home_dir(Some("C:\\test\\home"));
306            if let Err(e) = result {
307                assert!(e.to_string().contains("LM Studio not found"));
308            }
309        }
310    }
311
312    #[tokio::test]
313    async fn test_fetch_models_happy_path() {
314        if std::env::var("CODEX_SANDBOX_NETWORK_DISABLED").is_ok() {
315            return;
316        }
317
318        let Some(server) = start_mock_server_or_skip().await else {
319            return;
320        };
321        wiremock::Mock::given(wiremock::matchers::method("GET"))
322            .and(wiremock::matchers::path("/v1/models"))
323            .respond_with(
324                wiremock::ResponseTemplate::new(200).set_body_raw(
325                    serde_json::json!({
326                        "data": [
327                            {"id": "openai/gpt-oss-20b"},
328                        ]
329                    })
330                    .to_string(),
331                    "application/json",
332                ),
333            )
334            .mount(&server)
335            .await;
336
337        let client = LMStudioClient::try_from_base_url(&server.uri()).await;
338        assert!(client.is_ok());
339
340        let client = client.unwrap();
341        let models = client.fetch_models().await.expect("fetch models");
342        assert!(models.contains(&"openai/gpt-oss-20b".to_string()));
343    }
344
345    #[tokio::test]
346    async fn test_fetch_models_native_api() {
347        if std::env::var("CODEX_SANDBOX_NETWORK_DISABLED").is_ok() {
348            return;
349        }
350
351        let Some(server) = start_mock_server_or_skip().await else {
352            return;
353        };
354        wiremock::Mock::given(wiremock::matchers::method("GET"))
355            .and(wiremock::matchers::path("/api/v0/models"))
356            .respond_with(
357                wiremock::ResponseTemplate::new(200).set_body_raw(
358                    serde_json::json!({
359                        "data": [
360                            {"id": "lmstudio-community/meta-llama-3.1-8b-instruct"},
361                        ]
362                    })
363                    .to_string(),
364                    "application/json",
365                ),
366            )
367            .mount(&server)
368            .await;
369
370        let client = LMStudioClient::try_from_base_url_with_api_version(&server.uri(), true).await;
371        assert!(client.is_ok());
372
373        let client = client.unwrap();
374        let models = client.fetch_models().await.expect("fetch models");
375        assert!(models.contains(&"lmstudio-community/meta-llama-3.1-8b-instruct".to_string()));
376    }
377
378    #[tokio::test]
379    async fn test_fetch_models_no_data_array() {
380        if std::env::var("CODEX_SANDBOX_NETWORK_DISABLED").is_ok() {
381            return;
382        }
383
384        let Some(server) = start_mock_server_or_skip().await else {
385            return;
386        };
387        wiremock::Mock::given(wiremock::matchers::method("GET"))
388            .and(wiremock::matchers::path("/v1/models"))
389            .respond_with(
390                wiremock::ResponseTemplate::new(200)
391                    .set_body_raw(serde_json::json!({}).to_string(), "application/json"),
392            )
393            .mount(&server)
394            .await;
395
396        let client = LMStudioClient::try_from_base_url(&server.uri()).await;
397        let client = client.unwrap();
398        let result = client.fetch_models().await;
399
400        assert!(result.is_err());
401        assert!(result.unwrap_err().to_string().contains("No 'data' array in response"));
402    }
403
404    #[tokio::test]
405    async fn test_check_server_happy_path() {
406        if std::env::var("CODEX_SANDBOX_NETWORK_DISABLED").is_ok() {
407            return;
408        }
409
410        let Some(server) = start_mock_server_or_skip().await else {
411            return;
412        };
413        wiremock::Mock::given(wiremock::matchers::method("GET"))
414            .and(wiremock::matchers::path("/v1/models"))
415            .respond_with(wiremock::ResponseTemplate::new(200))
416            .mount(&server)
417            .await;
418
419        let result = LMStudioClient::try_from_base_url(&server.uri()).await;
420        result.unwrap();
421    }
422
423    #[tokio::test]
424    async fn test_check_server_error() {
425        if std::env::var("CODEX_SANDBOX_NETWORK_DISABLED").is_ok() {
426            return;
427        }
428
429        let Some(server) = start_mock_server_or_skip().await else {
430            return;
431        };
432        wiremock::Mock::given(wiremock::matchers::method("GET"))
433            .and(wiremock::matchers::path("/v1/models"))
434            .respond_with(wiremock::ResponseTemplate::new(404))
435            .mount(&server)
436            .await;
437
438        let result = LMStudioClient::try_from_base_url(&server.uri()).await;
439        assert!(result.is_err());
440        assert!(result.unwrap_err().to_string().contains("Server returned error: 404"));
441    }
442}