vtcode_llm/providers/lmstudio/
client.rs1use 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#[derive(Clone, Debug)]
23pub struct LMStudioClient {
24 client: reqwest::Client,
25 base_url: String,
26 use_native_api: bool,
28}
29
30impl LMStudioClient {
31 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 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 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 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 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 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 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 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 fn find_lms() -> io::Result<String> {
197 Self::find_lms_with_home_dir(None)
198 }
199
200 fn find_lms_with_home_dir(home_dir: Option<&str>) -> io::Result<String> {
202 if which::which("lms").is_ok() {
204 return Ok("lms".to_string());
205 }
206
207 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 pub async fn download_model(&self, model: &str) -> io::Result<()> {
239 let lms = Self::find_lms()?;
240 tracing::info!(model, "downloading model");
241
242 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 }
286 Err(e) => {
287 assert!(e.to_string().contains("LM Studio not found"));
289 }
290 }
291 }
292
293 #[test]
294 fn test_find_lms_with_mock_home() {
295 #[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}