Skip to main content

harn_vm/llm/
model_test.rs

1use std::time::Instant;
2
3use serde::Serialize;
4
5use super::api::{vm_call_llm_full_streaming, LlmCallOptions};
6use crate::value::{VmError, VmValue};
7
8const SMOKE_TEST_MAX_TOKENS: i64 = 32;
9
10#[derive(Clone, Debug, PartialEq, Eq)]
11pub struct ModelSmokeTestOptions {
12    pub model: String,
13    pub provider: Option<String>,
14    pub prompt: String,
15}
16
17#[derive(Clone, Debug, PartialEq, Serialize)]
18pub struct ModelSmokeTestResult {
19    pub model_id: String,
20    pub provider: String,
21    pub latency_ms: u64,
22    #[serde(skip_serializing_if = "Option::is_none")]
23    pub first_token_ms: Option<u64>,
24    pub input_tokens: i64,
25    pub output_tokens: i64,
26    pub estimated_cost_usd: f64,
27}
28
29pub async fn run_model_smoke_test(
30    options: ModelSmokeTestOptions,
31) -> Result<ModelSmokeTestResult, String> {
32    super::provider::register_default_providers();
33
34    let resolved = crate::llm_config::resolve_model_info(&options.model);
35    let model_id = resolved.id;
36    let provider = options
37        .provider
38        .as_deref()
39        .map(str::trim)
40        .filter(|provider| !provider.is_empty())
41        .map(str::to_string)
42        .unwrap_or(resolved.provider);
43    let api_key = super::resolve_api_key(&provider).map_err(vm_error_message)?;
44
45    if let Some(def) = crate::llm_config::provider_config(&provider) {
46        if super::supports_model_readiness_probe(&def) {
47            let readiness = super::readiness::probe_provider_readiness_with_options(
48                &provider,
49                super::readiness::ProviderReadinessOptions {
50                    requested_model: Some(&model_id),
51                    base_url_override: None,
52                    api_key_override: Some(&api_key),
53                },
54            )
55            .await;
56            if readiness_status_blocks_smoke_test(readiness.status) {
57                return Err(readiness.message);
58            }
59        }
60    }
61
62    let opts = LlmCallOptions {
63        provider: provider.clone(),
64        model: model_id.clone(),
65        api_key,
66        messages: vec![serde_json::json!({
67            "role": "user",
68            "content": options.prompt,
69        })],
70        max_tokens: SMOKE_TEST_MAX_TOKENS,
71        // A smoke test wants a bare, deterministic completion: no schema, no
72        // thinking, plain text. Those all match `LlmCallOptions::default()`,
73        // so only the routing + prompt + token cap need spelling out here.
74        ..LlmCallOptions::default()
75    };
76
77    let (delta_tx, mut delta_rx) = tokio::sync::mpsc::unbounded_channel::<String>();
78    let started = Instant::now();
79    let first_delta = tokio::spawn(async move { delta_rx.recv().await.map(|_| started.elapsed()) });
80    let result = vm_call_llm_full_streaming(&opts, delta_tx)
81        .await
82        .map_err(vm_error_message);
83    let latency_ms = duration_ms(started.elapsed());
84    let first_token_ms = first_delta.await.ok().flatten().map(duration_ms);
85    let result = result?;
86    let usage = result.usage();
87
88    Ok(ModelSmokeTestResult {
89        model_id: result.model,
90        provider: result.provider,
91        latency_ms,
92        first_token_ms,
93        input_tokens: usage.input_tokens,
94        output_tokens: usage.output_tokens,
95        estimated_cost_usd: usage.cost_usd.unwrap_or(0.0),
96    })
97}
98
99fn readiness_status_blocks_smoke_test(status: super::readiness::ReadinessStatus) -> bool {
100    matches!(
101        status,
102        super::readiness::ReadinessStatus::ModelMissing
103            | super::readiness::ReadinessStatus::InvalidUrl
104            | super::readiness::ReadinessStatus::ProviderMismatch
105    )
106}
107
108fn duration_ms(duration: std::time::Duration) -> u64 {
109    u64::try_from(duration.as_millis()).unwrap_or(u64::MAX)
110}
111
112fn vm_error_message(error: VmError) -> String {
113    match error {
114        VmError::CategorizedError { message, .. } => message,
115        VmError::Thrown(VmValue::String(message)) => message.to_string(),
116        VmError::Thrown(VmValue::Dict(dict)) => dict
117            .get("message")
118            .map(VmValue::display)
119            .unwrap_or_else(|| VmError::Thrown(VmValue::Dict(dict)).to_string()),
120        other => other.to_string(),
121    }
122}
123
124#[cfg(test)]
125mod tests {
126    use super::{readiness_status_blocks_smoke_test, run_model_smoke_test, ModelSmokeTestOptions};
127    use crate::llm::readiness::ReadinessStatus;
128
129    #[test]
130    fn smoke_test_blocks_provider_mismatch_before_generation() {
131        assert!(readiness_status_blocks_smoke_test(
132            ReadinessStatus::ProviderMismatch
133        ));
134        assert!(readiness_status_blocks_smoke_test(
135            ReadinessStatus::ModelMissing
136        ));
137        assert!(!readiness_status_blocks_smoke_test(ReadinessStatus::Ok));
138    }
139
140    #[tokio::test]
141    async fn mock_provider_smoke_test_reports_timing_tokens_and_cost() {
142        crate::llm::reset_llm_state();
143        let result = run_model_smoke_test(ModelSmokeTestOptions {
144            model: "mock".to_string(),
145            provider: Some("mock".to_string()),
146            prompt: "ping".to_string(),
147        })
148        .await
149        .expect("mock provider smoke test should not require network");
150
151        assert_eq!(result.model_id, "mock");
152        assert_eq!(result.provider, "mock");
153        assert_eq!(result.input_tokens, 4);
154        assert_eq!(result.output_tokens, 30);
155        assert_eq!(result.estimated_cost_usd, 0.0);
156        assert!(result.first_token_ms.is_some());
157    }
158}