Skip to main content

harn_vm/llm/
readiness.rs

1use serde::Serialize;
2
3use super::api::apply_auth_headers;
4use super::resolve_api_key;
5use crate::llm_config::{self, ProviderDef};
6
7#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
8#[serde(rename_all = "snake_case")]
9pub enum ReadinessStatus {
10    Ok,
11    UnknownProvider,
12    Unsupported,
13    InvalidUrl,
14    Unreachable,
15    BadStatus,
16    BadResponse,
17    ModelMissing,
18}
19
20#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
21pub struct ProviderReadiness {
22    pub provider: String,
23    pub ok: bool,
24    pub status: ReadinessStatus,
25    pub message: String,
26    pub base_url: Option<String>,
27    pub url: Option<String>,
28    pub model: Option<String>,
29    pub requested_model: Option<String>,
30    pub served_models: Vec<String>,
31    pub http_status: Option<u16>,
32}
33
34impl ProviderReadiness {
35    fn fail(
36        provider: &str,
37        status: ReadinessStatus,
38        message: String,
39        base_url: Option<String>,
40        url: Option<String>,
41        model: Option<String>,
42        requested_model: Option<String>,
43        http_status: Option<u16>,
44    ) -> Self {
45        Self {
46            provider: provider.to_string(),
47            ok: false,
48            status,
49            message,
50            base_url,
51            url,
52            model,
53            requested_model,
54            served_models: Vec::new(),
55            http_status,
56        }
57    }
58}
59
60#[derive(Debug, Clone, Copy, Default)]
61pub struct ProviderReadinessOptions<'a> {
62    pub requested_model: Option<&'a str>,
63    pub base_url_override: Option<&'a str>,
64    pub api_key_override: Option<&'a str>,
65}
66
67pub fn supports_model_readiness_probe(def: &ProviderDef) -> bool {
68    let healthcheck_lists_models = def.healthcheck.as_ref().is_some_and(|hc| {
69        hc.method.eq_ignore_ascii_case("GET") && {
70            hc.path.as_deref().is_some_and(is_model_inventory_endpoint)
71                || hc.url.as_deref().is_some_and(is_model_inventory_endpoint)
72        }
73    });
74    healthcheck_lists_models || openai_compatible_models_path(&def.chat_endpoint).is_some()
75}
76
77pub fn selected_model_for_provider(provider: &str) -> Option<String> {
78    configured_model_for_provider(provider).map(|model| {
79        let (resolved, _) = llm_config::resolve_model(model.trim());
80        resolved
81    })
82}
83
84pub fn build_models_url(def: &ProviderDef) -> Result<String, String> {
85    let base_url = llm_config::resolve_base_url(def);
86    models_url(def, &base_url)
87}
88
89pub async fn probe_provider_readiness(
90    provider: &str,
91    requested_model: Option<&str>,
92    base_url_override: Option<&str>,
93) -> ProviderReadiness {
94    probe_provider_readiness_with_options(
95        provider,
96        ProviderReadinessOptions {
97            requested_model,
98            base_url_override,
99            api_key_override: None,
100        },
101    )
102    .await
103}
104
105pub async fn probe_provider_readiness_with_options(
106    provider: &str,
107    options: ProviderReadinessOptions<'_>,
108) -> ProviderReadiness {
109    let Some(def) = llm_config::provider_config(provider) else {
110        return ProviderReadiness::fail(
111            provider,
112            ReadinessStatus::UnknownProvider,
113            format!("Unknown provider: {provider}"),
114            None,
115            None,
116            options.requested_model.map(ToOwned::to_owned),
117            options.requested_model.map(ToOwned::to_owned),
118            None,
119        );
120    };
121
122    let base_url = options
123        .base_url_override
124        .filter(|value| !value.trim().is_empty())
125        .map(|value| value.trim().to_string())
126        .unwrap_or_else(|| llm_config::resolve_base_url(&def));
127    let diagnostic_base_url = crate::egress::redact_diagnostic_text(&base_url);
128    let url = match models_url(&def, &base_url) {
129        Ok(url) => url,
130        Err(message) => {
131            let message = crate::egress::redact_diagnostic_text(&message);
132            let status = if supports_model_readiness_probe(&def) {
133                ReadinessStatus::InvalidUrl
134            } else {
135                ReadinessStatus::Unsupported
136            };
137            return ProviderReadiness::fail(
138                provider,
139                status,
140                message,
141                Some(diagnostic_base_url),
142                None,
143                options.requested_model.map(ToOwned::to_owned),
144                options.requested_model.map(ToOwned::to_owned),
145                None,
146            );
147        }
148    };
149    let diagnostic_url = crate::egress::redact_diagnostic_text(&url);
150
151    let (raw_model, resolved_model) = options
152        .requested_model
153        .filter(|model| !model.trim().is_empty())
154        .map(|model| {
155            let trimmed = model.trim();
156            let (resolved, _) = llm_config::resolve_model(trimmed);
157            (Some(trimmed.to_string()), Some(resolved))
158        })
159        .unwrap_or_else(|| match configured_model_for_provider(provider) {
160            Some(model) => {
161                let (resolved, _) = llm_config::resolve_model(&model);
162                (Some(model), Some(resolved))
163            }
164            None => (None, None),
165        });
166
167    let client = super::utility_client_for_base_url(&base_url);
168    let api_key = options
169        .api_key_override
170        .filter(|value| !value.trim().is_empty())
171        .map(|value| value.trim().to_string())
172        .unwrap_or_else(|| resolve_api_key(provider).unwrap_or_default());
173    let request = client.get(&url).header("Content-Type", "application/json");
174    let request = apply_auth_headers(request, &api_key, Some(&def));
175    let request = def
176        .extra_headers
177        .iter()
178        .fold(request, |request, (name, value)| {
179            request.header(name.as_str(), value.as_str())
180        });
181
182    let response = match request.send().await {
183        Ok(response) => response,
184        Err(error) => {
185            let error = crate::egress::redact_reqwest_error(&error);
186            return ProviderReadiness::fail(
187                provider,
188                ReadinessStatus::Unreachable,
189                format!("{provider} server is not reachable at {diagnostic_base_url}: {error}"),
190                Some(diagnostic_base_url),
191                Some(diagnostic_url),
192                resolved_model,
193                raw_model,
194                None,
195            );
196        }
197    };
198
199    let http_status = response.status().as_u16();
200    if !response.status().is_success() {
201        return ProviderReadiness::fail(
202            provider,
203            ReadinessStatus::BadStatus,
204            format!("{provider} returned HTTP {http_status} at {diagnostic_url}"),
205            Some(diagnostic_base_url),
206            Some(diagnostic_url),
207            resolved_model,
208            raw_model,
209            Some(http_status),
210        );
211    }
212
213    let body = match response.text().await {
214        Ok(body) => body,
215        Err(error) => {
216            let error = crate::egress::redact_reqwest_error(&error);
217            return ProviderReadiness::fail(
218                provider,
219                ReadinessStatus::BadResponse,
220                format!("{provider} returned an unreadable /models response: {error}"),
221                Some(diagnostic_base_url),
222                Some(diagnostic_url),
223                resolved_model,
224                raw_model,
225                Some(http_status),
226            );
227        }
228    };
229    let served_models = match parse_model_ids(&body) {
230        Ok(models) if !models.is_empty() => models,
231        Ok(_) => {
232            return ProviderReadiness::fail(
233                provider,
234                ReadinessStatus::BadResponse,
235                format!("{provider} /models response did not include any model ids"),
236                Some(diagnostic_base_url),
237                Some(diagnostic_url),
238                resolved_model,
239                raw_model,
240                Some(http_status),
241            );
242        }
243        Err(error) => {
244            return ProviderReadiness::fail(
245                provider,
246                ReadinessStatus::BadResponse,
247                format!("{provider} returned an unparsable /models response: {error}"),
248                Some(diagnostic_base_url),
249                Some(diagnostic_url),
250                resolved_model,
251                raw_model,
252                Some(http_status),
253            );
254        }
255    };
256
257    let readiness_model = resolved_model.as_deref().map(llm_config::wire_model_id);
258
259    if let Some(model) = readiness_model.as_deref() {
260        if !model_is_served(model, &served_models) {
261            let display_model = resolved_model.as_deref().unwrap_or(model);
262            let model_label = if display_model == model {
263                format!("Model '{display_model}'")
264            } else {
265                format!("Model '{display_model}' (wire model '{model}')")
266            };
267            return ProviderReadiness {
268                provider: provider.to_string(),
269                ok: false,
270                status: ReadinessStatus::ModelMissing,
271                message: format!(
272                    "{model_label} is not served by {provider} at {diagnostic_base_url}. Currently served: {}",
273                    served_models.join(", ")
274                ),
275                base_url: Some(diagnostic_base_url),
276                url: Some(diagnostic_url),
277                model: resolved_model,
278                requested_model: raw_model,
279                served_models,
280                http_status: Some(http_status),
281            };
282        }
283    }
284
285    let message = match (resolved_model.as_deref(), readiness_model.as_deref()) {
286        (Some(model), Some(wire_model)) if model != wire_model => format!(
287            "{provider} is ready at {diagnostic_base_url}; model '{model}' is served as '{wire_model}'"
288        ),
289        (Some(model), _) => {
290            format!("{provider} is ready at {diagnostic_base_url}; model '{model}' is served")
291        }
292        (None, _) => format!(
293            "{provider} is reachable at {diagnostic_base_url}; served models: {}",
294            served_models.join(", ")
295        ),
296    };
297
298    ProviderReadiness {
299        provider: provider.to_string(),
300        ok: true,
301        status: ReadinessStatus::Ok,
302        message,
303        base_url: Some(diagnostic_base_url),
304        url: Some(diagnostic_url),
305        model: resolved_model,
306        requested_model: raw_model,
307        served_models,
308        http_status: Some(http_status),
309    }
310}
311
312pub fn parse_model_ids(body: &str) -> Result<Vec<String>, serde_json::Error> {
313    let payload: serde_json::Value = serde_json::from_str(body)?;
314    Ok(parse_model_ids_from_value(&payload))
315}
316
317pub fn parse_model_ids_from_value(payload: &serde_json::Value) -> Vec<String> {
318    let mut models = Vec::new();
319    if let Some(entries) = payload.as_array() {
320        collect_model_ids(entries, &mut models);
321    }
322    if let Some(entries) = payload.get("data").and_then(|value| value.as_array()) {
323        collect_model_ids(entries, &mut models);
324    }
325    if let Some(entries) = payload.get("models").and_then(|value| value.as_array()) {
326        collect_model_ids(entries, &mut models);
327    }
328    models.sort();
329    models.dedup();
330    models
331}
332
333fn collect_model_ids(entries: &[serde_json::Value], models: &mut Vec<String>) {
334    for entry in entries {
335        if let Some(id) = entry.as_str().or_else(|| {
336            entry
337                .get("id")
338                .or_else(|| entry.get("name"))
339                .and_then(|value| value.as_str())
340        }) {
341            models.push(id.to_string());
342        }
343    }
344}
345
346pub fn model_is_served(model: &str, served_models: &[String]) -> bool {
347    served_models.iter().any(|served| {
348        served == model
349            || served
350                .strip_prefix(model)
351                .is_some_and(|suffix| suffix.starts_with(':'))
352    })
353}
354
355pub fn configured_model_for_provider(provider: &str) -> Option<String> {
356    if provider == "mlx" {
357        if let Ok(model) = std::env::var("MLX_MODEL_ID") {
358            if !model.trim().is_empty() {
359                return Some(model);
360            }
361        }
362    }
363    if provider == "local" {
364        if let Ok(model) = std::env::var("LOCAL_LLM_MODEL") {
365            if !model.trim().is_empty() {
366                return Some(model);
367            }
368        }
369    }
370    let harn_provider = std::env::var("HARN_LLM_PROVIDER").ok();
371    let model = std::env::var("HARN_LLM_MODEL")
372        .ok()
373        .filter(|model| !model.trim().is_empty())?;
374    let (_, resolved_provider) = llm_config::resolve_model(&model);
375    if resolved_provider.as_deref() == Some(provider)
376        || (resolved_provider.is_none() && harn_provider.as_deref() == Some(provider))
377    {
378        return Some(model);
379    }
380    None
381}
382
383fn models_url(def: &ProviderDef, base_url: &str) -> Result<String, String> {
384    if let Some(url) = def.healthcheck.as_ref().and_then(|healthcheck| {
385        if healthcheck.method.eq_ignore_ascii_case("GET") {
386            healthcheck
387                .url
388                .as_deref()
389                .filter(|url| is_model_inventory_endpoint(url))
390        } else {
391            None
392        }
393    }) {
394        return reqwest::Url::parse(url)
395            .map(|_| normalize_loopback(url))
396            .map_err(|error| format!("Invalid provider models URL '{url}': {error}"));
397    }
398
399    let path = def
400        .healthcheck
401        .as_ref()
402        .and_then(|healthcheck| {
403            if healthcheck.method.eq_ignore_ascii_case("GET") {
404                healthcheck
405                    .path
406                    .as_deref()
407                    .filter(|path| is_model_inventory_endpoint(path))
408            } else {
409                None
410            }
411        })
412        .map(ToOwned::to_owned);
413    let path = match path.or_else(|| openai_compatible_models_path(&def.chat_endpoint)) {
414        Some(path) => path,
415        None => {
416            return Err(
417                "Provider does not expose a model readiness endpoint; configure a GET healthcheck path/url that lists models or use an OpenAI-compatible /chat/completions endpoint".to_string(),
418            );
419        }
420    };
421    let url = super::healthcheck::join_base_and_path(base_url, &path);
422    reqwest::Url::parse(&url)
423        .map(|_| normalize_loopback(&url))
424        .map_err(|error| format!("Invalid provider models URL '{url}': {error}"))
425}
426
427fn is_model_inventory_endpoint(endpoint: &str) -> bool {
428    let path = reqwest::Url::parse(endpoint)
429        .ok()
430        .map(|url| url.path().to_string())
431        .unwrap_or_else(|| endpoint.split('?').next().unwrap_or(endpoint).to_string());
432    let path = path.trim_end_matches('/');
433    path == "models" || path.ends_with("/models") || path.ends_with("/api/tags")
434}
435
436fn openai_compatible_models_path(chat_endpoint: &str) -> Option<String> {
437    let prefix = chat_endpoint.strip_suffix("/chat/completions")?;
438    Some(if prefix.is_empty() {
439        "/models".to_string()
440    } else {
441        format!("{prefix}/models")
442    })
443}
444
445fn normalize_loopback(url: &str) -> String {
446    url.replace("://localhost:", "://127.0.0.1:")
447}
448
449#[cfg(test)]
450mod tests {
451    use super::*;
452    use std::io::{Read, Write};
453    use std::net::TcpListener;
454
455    #[test]
456    fn parse_model_ids_reads_openai_compatible_data() {
457        let models =
458            parse_model_ids(r#"{"object":"list","data":[{"id":"qwen"},{"id":"mlx-model"}]}"#)
459                .expect("parse models");
460        assert_eq!(models, vec!["mlx-model".to_string(), "qwen".to_string()]);
461    }
462
463    #[test]
464    fn parse_model_ids_reads_together_top_level_array() {
465        let models = parse_model_ids(r#"[{"id":"deepseek-ai/DeepSeek-V4-Pro"},{"name":"qwen"}]"#)
466            .expect("parse models");
467        assert_eq!(
468            models,
469            vec![
470                "deepseek-ai/DeepSeek-V4-Pro".to_string(),
471                "qwen".to_string()
472            ]
473        );
474    }
475
476    #[test]
477    fn models_url_does_not_duplicate_version_prefix_in_base_url() {
478        let def = ProviderDef {
479            base_url: "https://openrouter.ai/api/v1".to_string(),
480            chat_endpoint: "/chat/completions".to_string(),
481            healthcheck: Some(crate::llm_config::HealthcheckDef {
482                method: "GET".to_string(),
483                path: Some("/auth/key".to_string()),
484                url: None,
485                body: None,
486            }),
487            ..Default::default()
488        };
489
490        assert_eq!(
491            models_url(&def, &def.base_url).expect("models url"),
492            "https://openrouter.ai/api/v1/models"
493        );
494    }
495
496    #[test]
497    fn anthropic_models_url_uses_catalog_healthcheck_path() {
498        let def = llm_config::provider_config("anthropic").expect("anthropic provider");
499
500        assert!(supports_model_readiness_probe(&def));
501        assert_eq!(
502            models_url(&def, &def.base_url).expect("models url"),
503            "https://api.anthropic.com/v1/models"
504        );
505    }
506
507    #[test]
508    fn models_url_uses_catalogued_inventory_path_for_native_endpoint() {
509        let def = ProviderDef {
510            base_url: "http://localhost:11434".to_string(),
511            chat_endpoint: "/api/chat".to_string(),
512            healthcheck: Some(crate::llm_config::HealthcheckDef {
513                method: "GET".to_string(),
514                path: Some("/api/tags".to_string()),
515                url: None,
516                body: None,
517            }),
518            ..Default::default()
519        };
520
521        assert!(supports_model_readiness_probe(&def));
522        assert_eq!(
523            models_url(&def, &def.base_url).expect("models url"),
524            "http://127.0.0.1:11434/api/tags"
525        );
526    }
527
528    #[test]
529    fn models_url_rejects_native_endpoint_without_model_inventory() {
530        let def = ProviderDef {
531            base_url: "https://api.example.com/v1".to_string(),
532            chat_endpoint: "/messages".to_string(),
533            healthcheck: None,
534            ..Default::default()
535        };
536
537        assert!(!supports_model_readiness_probe(&def));
538        assert!(models_url(&def, &def.base_url)
539            .expect_err("unsupported native endpoint")
540            .contains("model readiness endpoint"));
541    }
542
543    #[test]
544    fn model_is_served_accepts_exact_ids_or_tag_boundaries() {
545        let models = vec![
546            "qwen3:8b".to_string(),
547            "unsloth/Qwen3.6-35B-A3B-UD-MLX-4bit".to_string(),
548            "gpt-4o".to_string(),
549        ];
550        assert!(model_is_served("qwen3", &models));
551        assert!(model_is_served(
552            "unsloth/Qwen3.6-35B-A3B-UD-MLX-4bit",
553            &models
554        ));
555        assert!(!model_is_served("unsloth/Qwen3.6", &models));
556        assert!(!model_is_served("gpt-4", &models));
557    }
558
559    #[tokio::test]
560    async fn probe_provider_readiness_verifies_served_model() {
561        let (base_url, handle) = spawn_models_stub(
562            200,
563            r#"{"data":[{"id":"unsloth/Qwen3.6-35B-A3B-UD-MLX-4bit"}]}"#,
564        );
565        let result = probe_provider_readiness("mlx", Some("mlx-qwen36-27b"), Some(&base_url)).await;
566        handle.join().expect("stub joins");
567        assert!(result.ok);
568        assert_eq!(result.status, ReadinessStatus::Ok);
569        assert_eq!(
570            result.model.as_deref(),
571            Some("unsloth/Qwen3.6-35B-A3B-UD-MLX-4bit")
572        );
573    }
574
575    #[tokio::test]
576    async fn probe_provider_readiness_verifies_wire_model_for_catalog_key() {
577        let (base_url, handle) = spawn_models_stub(200, r#"{"data":[{"id":"zai-org/GLM-5.2"}]}"#);
578        let result = probe_provider_readiness(
579            "deepinfra",
580            Some("deepinfra/zai-org/GLM-5.2"),
581            Some(&base_url),
582        )
583        .await;
584        handle.join().expect("stub joins");
585        assert!(result.ok, "{}", result.message);
586        assert_eq!(result.status, ReadinessStatus::Ok);
587        assert_eq!(result.model.as_deref(), Some("deepinfra/zai-org/GLM-5.2"));
588        assert!(result.message.contains("served as 'zai-org/GLM-5.2'"));
589    }
590
591    #[tokio::test]
592    async fn probe_provider_readiness_uses_explicit_api_key_override() {
593        let (base_url, handle) = spawn_models_stub_with_expected_header(
594            200,
595            r#"{"data":[{"id":"zai-org/GLM-5.2"}]}"#,
596            Some("authorization: Bearer test-key"),
597        );
598        let result = probe_provider_readiness_with_options(
599            "deepinfra",
600            ProviderReadinessOptions {
601                requested_model: Some("deepinfra/zai-org/GLM-5.2"),
602                base_url_override: Some(&base_url),
603                api_key_override: Some("test-key"),
604            },
605        )
606        .await;
607        handle.join().expect("stub joins");
608        assert!(result.ok, "{}", result.message);
609    }
610
611    #[tokio::test]
612    async fn probe_provider_readiness_reports_missing_model() {
613        let (base_url, handle) = spawn_models_stub(200, r#"{"data":[{"id":"other-model"}]}"#);
614        let result = probe_provider_readiness("mlx", Some("mlx-qwen36-27b"), Some(&base_url)).await;
615        handle.join().expect("stub joins");
616        assert!(!result.ok);
617        assert_eq!(result.status, ReadinessStatus::ModelMissing);
618        assert!(result.message.contains("Currently served: other-model"));
619    }
620
621    fn spawn_models_stub(status: u16, body: &'static str) -> (String, std::thread::JoinHandle<()>) {
622        spawn_models_stub_with_expected_header(status, body, None)
623    }
624
625    fn spawn_models_stub_with_expected_header(
626        status: u16,
627        body: &'static str,
628        expected_header: Option<&'static str>,
629    ) -> (String, std::thread::JoinHandle<()>) {
630        let listener = TcpListener::bind("127.0.0.1:0").expect("bind models stub");
631        let addr = listener.local_addr().expect("stub addr");
632        // Block on `accept()` directly rather than polling. The
633        // earlier 3s wall-clock deadline + 20ms polling sleep was
634        // brittle under nextest's flake-detection profile: another
635        // concurrent test could starve this thread of CPU long
636        // enough for the deadline to elapse before the kernel even
637        // delivered the SYN that the client had already sent.
638        // Blocking accept is deterministic; the test invariably
639        // sends a request, so it returns promptly.
640        let handle = std::thread::spawn(move || {
641            let (mut stream, _) = listener
642                .accept()
643                .unwrap_or_else(|e| panic!("models stub: accept failed: {e}"));
644            let mut buf = vec![0u8; 4096];
645            let n = stream.read(&mut buf).expect("read request");
646            let request = String::from_utf8_lossy(&buf[..n]);
647            assert!(
648                request.starts_with("GET /v1/models HTTP/1.1\r\n")
649                    || request.starts_with("GET /models HTTP/1.1\r\n")
650                    || request.starts_with("GET /api/tags HTTP/1.1\r\n")
651            );
652            if let Some(header) = expected_header {
653                assert!(
654                    request
655                        .lines()
656                        .any(|line| line.eq_ignore_ascii_case(header)),
657                    "expected request header {header:?}, got:\n{request}"
658                );
659            }
660            let response = format!(
661                "HTTP/1.1 {status} OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
662                body.len(),
663                body
664            );
665            stream
666                .write_all(response.as_bytes())
667                .expect("write response");
668        });
669        (format!("http://{addr}"), handle)
670    }
671}