harn_vm/llm/
model_test.rs1use 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 ..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
87 Ok(ModelSmokeTestResult {
88 model_id: result.model.clone(),
89 provider: result.provider.clone(),
90 latency_ms,
91 first_token_ms,
92 input_tokens: result.input_tokens,
93 output_tokens: result.output_tokens,
94 estimated_cost_usd: result.priced_cost_usd().unwrap_or(0.0),
95 })
96}
97
98fn readiness_status_blocks_smoke_test(status: super::readiness::ReadinessStatus) -> bool {
99 matches!(
100 status,
101 super::readiness::ReadinessStatus::ModelMissing
102 | super::readiness::ReadinessStatus::InvalidUrl
103 | super::readiness::ReadinessStatus::ProviderMismatch
104 )
105}
106
107fn duration_ms(duration: std::time::Duration) -> u64 {
108 u64::try_from(duration.as_millis()).unwrap_or(u64::MAX)
109}
110
111fn vm_error_message(error: VmError) -> String {
112 match error {
113 VmError::CategorizedError { message, .. } => message,
114 VmError::Thrown(VmValue::String(message)) => message.to_string(),
115 VmError::Thrown(VmValue::Dict(dict)) => dict
116 .get("message")
117 .map(VmValue::display)
118 .unwrap_or_else(|| VmError::Thrown(VmValue::Dict(dict)).to_string()),
119 other => other.to_string(),
120 }
121}
122
123#[cfg(test)]
124mod tests {
125 use super::{readiness_status_blocks_smoke_test, run_model_smoke_test, ModelSmokeTestOptions};
126 use crate::llm::readiness::ReadinessStatus;
127
128 #[test]
129 fn smoke_test_blocks_provider_mismatch_before_generation() {
130 assert!(readiness_status_blocks_smoke_test(
131 ReadinessStatus::ProviderMismatch
132 ));
133 assert!(readiness_status_blocks_smoke_test(
134 ReadinessStatus::ModelMissing
135 ));
136 assert!(!readiness_status_blocks_smoke_test(ReadinessStatus::Ok));
137 }
138
139 #[tokio::test]
140 async fn mock_provider_smoke_test_reports_timing_tokens_and_cost() {
141 crate::llm::reset_llm_state();
142 let result = run_model_smoke_test(ModelSmokeTestOptions {
143 model: "mock".to_string(),
144 provider: Some("mock".to_string()),
145 prompt: "ping".to_string(),
146 })
147 .await
148 .expect("mock provider smoke test should not require network");
149
150 assert_eq!(result.model_id, "mock");
151 assert_eq!(result.provider, "mock");
152 assert_eq!(result.input_tokens, 4);
153 assert_eq!(result.output_tokens, 30);
154 assert_eq!(result.estimated_cost_usd, 0.0);
155 assert!(result.first_token_ms.is_some());
156 }
157}