1use std::time::{Duration, Instant};
4
5use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
6use serde::{Deserialize, Serialize};
7use serde_json::{Value, json};
8
9use super::profile::{StoredModelProfile, redacted_url};
10use super::*;
11pub(super) use crate::clock::system_now_millis_or_zero as now_millis;
12use crate::net::{
13 http::{HttpConfig, QosHttpClientError, QosHttpResponse, send_request_with_qos},
14 qos::{QosPolicy, QosRuntime},
15};
16use crate::retrieval::ReadModelBackendConfig;
17
18#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
20pub struct ModelConnectivityProbeRequest {
21 pub profile_name: Option<String>,
22 pub override_config: Option<ModelProfileSaveRequest>,
23 pub timeout_ms: Option<u64>,
24}
25
26#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
28pub struct ModelDiscoveryRequest {
29 pub profile_name: Option<String>,
30 pub override_config: Option<ModelProfileSaveRequest>,
31 pub timeout_ms: Option<u64>,
32}
33
34#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
36pub struct ModelConnectivityTokenUsage {
37 pub prompt_tokens: u64,
38 pub completion_tokens: u64,
39 pub total_tokens: u64,
40}
41
42#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
44pub struct ModelConnectivityDiagnostics {
45 pub endpoint_reachable: bool,
46 pub auth_valid: bool,
47 pub rate_limited: bool,
48}
49
50#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
52pub struct ModelConnectivityProbeResult {
53 pub ok: bool,
54 pub provider: ModelProviderKind,
55 pub model: String,
56 pub latency_ms: u64,
57 pub checked_at_ms: u64,
58 pub diagnostics: ModelConnectivityDiagnostics,
59 pub token_usage: Option<ModelConnectivityTokenUsage>,
60 pub error_code: Option<String>,
61 pub error_message: Option<String>,
62 pub retryable: bool,
63}
64
65#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
67pub struct ModelDiscoveryEntry {
68 pub model: String,
69 pub context_window: Option<u32>,
70 pub output_limit: Option<u32>,
71 pub capabilities: ModelCapabilities,
72}
73
74#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
76pub struct ModelDiscoveryResult {
77 pub ok: bool,
78 pub provider: ModelProviderKind,
79 pub base_url: String,
80 pub latency_ms: u64,
81 pub checked_at_ms: u64,
82 pub diagnostics: ModelConnectivityDiagnostics,
83 pub models: Vec<String>,
84 pub model_entries: Vec<ModelDiscoveryEntry>,
85 pub error_code: Option<String>,
86 pub error_message: Option<String>,
87 pub retryable: bool,
88}
89
90impl ModelProviderConfigService {
91 pub async fn probe(
92 &self,
93 http: &HttpConfig,
94 retrieval: &ReadModelBackendConfig,
95 request: ModelConnectivityProbeRequest,
96 ) -> Result<ModelConnectivityProbeResult, ModelProviderError> {
97 let qos = QosRuntime::default();
98 let policy = default_qos_policy();
99 self.probe_with_qos(http, &qos, &policy, retrieval, request)
100 .await
101 }
102
103 pub async fn probe_with_qos(
104 &self,
105 http: &HttpConfig,
106 qos: &QosRuntime,
107 policy: &QosPolicy,
108 retrieval: &ReadModelBackendConfig,
109 request: ModelConnectivityProbeRequest,
110 ) -> Result<ModelConnectivityProbeResult, ModelProviderError> {
111 let profile = self
112 .resolve_probe_profile(retrieval, request.profile_name, request.override_config)
113 .await?;
114 let request_timeout = request_timeout_from_ms(request.timeout_ms);
115 let started = Instant::now();
116 let checked_at_ms = now_millis();
117 if profile.provider == ModelProviderKind::Echo {
118 return Ok(ModelConnectivityProbeResult {
119 ok: true,
120 provider: profile.provider,
121 model: profile.model,
122 latency_ms: elapsed_millis(started),
123 checked_at_ms,
124 diagnostics: ok_diagnostics(),
125 token_usage: Some(ModelConnectivityTokenUsage {
126 prompt_tokens: 4,
127 completion_tokens: 2,
128 total_tokens: 6,
129 }),
130 error_code: None,
131 error_message: None,
132 retryable: false,
133 });
134 }
135 if matches!(
136 profile.provider,
137 ModelProviderKind::Maas | ModelProviderKind::Codeagent
138 ) {
139 return Ok(unsupported_probe(profile, started, checked_at_ms));
140 }
141
142 let client = provider_http_client(http, &profile)?;
143 let response =
144 send_probe_request_with_qos(&client, qos, policy, &profile, request_timeout).await;
145 Ok(probe_result_from_http(profile, started, checked_at_ms, response).await)
146 }
147
148 pub async fn discover(
149 &self,
150 http: &HttpConfig,
151 retrieval: &ReadModelBackendConfig,
152 request: ModelDiscoveryRequest,
153 ) -> Result<ModelDiscoveryResult, ModelProviderError> {
154 let qos = QosRuntime::default();
155 let policy = default_qos_policy();
156 self.discover_with_qos(http, &qos, &policy, retrieval, request)
157 .await
158 }
159
160 pub async fn discover_with_qos(
161 &self,
162 http: &HttpConfig,
163 qos: &QosRuntime,
164 policy: &QosPolicy,
165 retrieval: &ReadModelBackendConfig,
166 request: ModelDiscoveryRequest,
167 ) -> Result<ModelDiscoveryResult, ModelProviderError> {
168 let profile = self
169 .resolve_probe_profile(retrieval, request.profile_name, request.override_config)
170 .await?;
171 let request_timeout = request_timeout_from_ms(request.timeout_ms);
172 let started = Instant::now();
173 let checked_at_ms = now_millis();
174 if profile.provider == ModelProviderKind::Echo {
175 return Ok(ModelDiscoveryResult {
176 ok: true,
177 provider: profile.provider,
178 base_url: redacted_url(&profile.base_url),
179 latency_ms: elapsed_millis(started),
180 checked_at_ms,
181 diagnostics: ok_diagnostics(),
182 models: vec![profile.model.clone()],
183 model_entries: vec![ModelDiscoveryEntry {
184 model: profile.model,
185 context_window: None,
186 output_limit: None,
187 capabilities: ModelCapabilities::default(),
188 }],
189 error_code: None,
190 error_message: None,
191 retryable: false,
192 });
193 }
194 if matches!(
195 profile.provider,
196 ModelProviderKind::Maas | ModelProviderKind::Codeagent
197 ) {
198 return Ok(unsupported_discovery(profile, started, checked_at_ms));
199 }
200
201 let client = provider_http_client(http, &profile)?;
202 let response =
203 send_discovery_request_with_qos(&client, qos, policy, &profile, request_timeout).await;
204 Ok(discovery_result_from_http(profile, started, checked_at_ms, response).await)
205 }
206}
207
208fn default_qos_policy() -> QosPolicy {
209 QosPolicy::new(
210 crate::net::qos::DEFAULT_MAX_CONNECTIONS,
211 crate::net::qos::DEFAULT_MAX_IN_FLIGHT_REQUESTS,
212 crate::net::qos::DEFAULT_MAX_QUEUE_DEPTH,
213 )
214 .expect("default QoS policy should validate")
215}
216
217#[cfg(test)]
218pub(super) async fn send_probe_request(
219 client: &reqwest::Client,
220 profile: &StoredModelProfile,
221 request_timeout: Option<Duration>,
222) -> Result<QosHttpResponse, QosHttpClientError> {
223 let qos = QosRuntime::default();
224 let policy = default_qos_policy();
225 send_probe_request_with_qos(client, &qos, &policy, profile, request_timeout).await
226}
227
228pub(super) async fn send_probe_request_with_qos(
229 client: &reqwest::Client,
230 qos: &QosRuntime,
231 policy: &QosPolicy,
232 profile: &StoredModelProfile,
233 request_timeout: Option<Duration>,
234) -> Result<QosHttpResponse, QosHttpClientError> {
235 let request = match profile.provider {
236 ModelProviderKind::Anthropic => client
237 .post(format!(
238 "{}/v1/messages",
239 profile.base_url.trim_end_matches('/')
240 ))
241 .headers(auth_headers(profile))
242 .json(&json!({
243 "model": profile.model,
244 "max_tokens": profile.max_tokens.unwrap_or(16),
245 "messages": [{"role": "user", "content": "relay-knowledge provider probe"}]
246 })),
247 ModelProviderKind::OpenAiCompatible if uses_embedding_probe(profile) => client
248 .post(format!(
249 "{}/embeddings",
250 profile.base_url.trim_end_matches('/')
251 ))
252 .headers(auth_headers(profile))
253 .json(&json!({
254 "model": profile.model,
255 "input": "relay-knowledge provider probe"
256 })),
257 _ => client
258 .post(format!(
259 "{}/chat/completions",
260 profile.base_url.trim_end_matches('/')
261 ))
262 .headers(auth_headers(profile))
263 .json(&json!({
264 "model": profile.model,
265 "temperature": profile.temperature,
266 "top_p": profile.top_p,
267 "max_tokens": profile.max_tokens.unwrap_or(16),
268 "messages": [{"role": "user", "content": "relay-knowledge provider probe"}]
269 })),
270 };
271 send_request_with_qos(qos, policy, apply_request_timeout(request, request_timeout)).await
272}
273
274fn uses_embedding_probe(profile: &StoredModelProfile) -> bool {
275 profile.source == "environment" || is_embedding_model_name(&profile.model)
276}
277
278fn is_embedding_model_name(model: &str) -> bool {
279 let normalized = model.to_ascii_lowercase();
280 normalized.contains("embedding")
281 || normalized.contains("embed")
282 || normalized.starts_with("bge-")
283 || normalized.starts_with("e5-")
284}
285
286#[cfg(test)]
287pub(super) async fn send_discovery_request(
288 client: &reqwest::Client,
289 profile: &StoredModelProfile,
290 request_timeout: Option<Duration>,
291) -> Result<QosHttpResponse, QosHttpClientError> {
292 let qos = QosRuntime::default();
293 let policy = default_qos_policy();
294 send_discovery_request_with_qos(client, &qos, &policy, profile, request_timeout).await
295}
296
297pub(super) async fn send_discovery_request_with_qos(
298 client: &reqwest::Client,
299 qos: &QosRuntime,
300 policy: &QosPolicy,
301 profile: &StoredModelProfile,
302 request_timeout: Option<Duration>,
303) -> Result<QosHttpResponse, QosHttpClientError> {
304 let url = match profile.provider {
305 ModelProviderKind::Anthropic => {
306 format!("{}/v1/models", profile.base_url.trim_end_matches('/'))
307 }
308 _ => format!("{}/models", profile.base_url.trim_end_matches('/')),
309 };
310 send_request_with_qos(
311 qos,
312 policy,
313 apply_request_timeout(
314 client.get(url).headers(auth_headers(profile)),
315 request_timeout,
316 ),
317 )
318 .await
319}
320
321pub(super) fn provider_http_client(
322 http: &HttpConfig,
323 profile: &StoredModelProfile,
324) -> Result<reqwest::Client, ModelProviderError> {
325 crate::net::http::outbound_json_client_with_policy(
326 http,
327 profile.ssl_verify,
328 Some(Duration::from_secs_f64(profile.connect_timeout_seconds)),
329 )
330 .map_err(|error| ModelProviderError::Network(error.to_string()))
331}
332
333fn apply_request_timeout(
334 request: reqwest::RequestBuilder,
335 timeout: Option<Duration>,
336) -> reqwest::RequestBuilder {
337 match timeout {
338 Some(timeout) => request.timeout(timeout),
339 None => request,
340 }
341}
342
343pub(super) fn auth_headers(profile: &StoredModelProfile) -> HeaderMap {
344 let mut headers = HeaderMap::new();
345 match profile.provider {
346 ModelProviderKind::Anthropic => {
347 if let Some(api_key) = &profile.api_key {
348 if let Ok(value) = HeaderValue::from_str(api_key) {
349 headers.insert("x-api-key", value);
350 }
351 }
352 headers.insert("anthropic-version", HeaderValue::from_static("2023-06-01"));
353 }
354 _ => {
355 if let Some(api_key) = &profile.api_key {
356 if let Ok(value) = HeaderValue::from_str(&format!("Bearer {api_key}")) {
357 headers.insert("authorization", value);
358 }
359 }
360 }
361 }
362 for header in &profile.headers {
363 if let (Ok(name), Some(value)) = (
364 HeaderName::from_bytes(header.name.as_bytes()),
365 header.value.as_ref(),
366 ) {
367 if let Ok(value) = HeaderValue::from_str(value) {
368 headers.insert(name, value);
369 }
370 }
371 }
372 headers
373}
374
375pub(super) async fn probe_result_from_http(
376 profile: StoredModelProfile,
377 started: Instant,
378 checked_at_ms: u64,
379 response: Result<QosHttpResponse, QosHttpClientError>,
380) -> ModelConnectivityProbeResult {
381 match response {
382 Ok(response) => {
383 let status = response.status();
384 let token_usage = response
385 .json::<Value>()
386 .await
387 .ok()
388 .and_then(|payload| token_usage(&payload));
389 let ok = status.is_success();
390 ModelConnectivityProbeResult {
391 ok,
392 provider: profile.provider,
393 model: profile.model,
394 latency_ms: elapsed_millis(started),
395 checked_at_ms,
396 diagnostics: diagnostics_from_status(status.as_u16()),
397 token_usage,
398 error_code: (!ok).then(|| status_error_code(status.as_u16()).to_owned()),
399 error_message: (!ok).then(|| format!("provider returned HTTP {status}")),
400 retryable: is_retryable_status(status.as_u16()),
401 }
402 }
403 Err(error) => transport_probe_result(profile, started, checked_at_ms, error),
404 }
405}
406
407pub(super) async fn discovery_result_from_http(
408 profile: StoredModelProfile,
409 started: Instant,
410 checked_at_ms: u64,
411 response: Result<QosHttpResponse, QosHttpClientError>,
412) -> ModelDiscoveryResult {
413 match response {
414 Ok(response) => {
415 let status = response.status();
416 if !status.is_success() {
417 return ModelDiscoveryResult {
418 ok: false,
419 provider: profile.provider,
420 base_url: redacted_url(&profile.base_url),
421 latency_ms: elapsed_millis(started),
422 checked_at_ms,
423 diagnostics: diagnostics_from_status(status.as_u16()),
424 models: Vec::new(),
425 model_entries: Vec::new(),
426 error_code: Some(status_error_code(status.as_u16()).to_owned()),
427 error_message: Some(format!("provider returned HTTP {status}")),
428 retryable: is_retryable_status(status.as_u16()),
429 };
430 }
431 let payload = match response.json::<Value>().await {
432 Ok(payload) => payload,
433 Err(error) => {
434 return ModelDiscoveryResult {
435 ok: false,
436 provider: profile.provider,
437 base_url: redacted_url(&profile.base_url),
438 latency_ms: elapsed_millis(started),
439 checked_at_ms,
440 diagnostics: ModelConnectivityDiagnostics {
441 endpoint_reachable: true,
442 auth_valid: true,
443 rate_limited: false,
444 },
445 models: Vec::new(),
446 model_entries: Vec::new(),
447 error_code: Some("invalid_response".to_owned()),
448 error_message: Some(format!(
449 "provider returned invalid model discovery JSON: {error}"
450 )),
451 retryable: false,
452 };
453 }
454 };
455 let entries = parse_discovery_entries(&payload);
456 let models = entries.iter().map(|entry| entry.model.clone()).collect();
457 ModelDiscoveryResult {
458 ok: true,
459 provider: profile.provider,
460 base_url: redacted_url(&profile.base_url),
461 latency_ms: elapsed_millis(started),
462 checked_at_ms,
463 diagnostics: ok_diagnostics(),
464 models,
465 model_entries: entries,
466 error_code: None,
467 error_message: None,
468 retryable: false,
469 }
470 }
471 Err(error) => ModelDiscoveryResult {
472 ok: false,
473 provider: profile.provider,
474 base_url: redacted_url(&profile.base_url),
475 latency_ms: elapsed_millis(started),
476 checked_at_ms,
477 diagnostics: ModelConnectivityDiagnostics {
478 endpoint_reachable: false,
479 auth_valid: false,
480 rate_limited: false,
481 },
482 models: Vec::new(),
483 model_entries: Vec::new(),
484 error_code: Some(if error.is_timeout() {
485 "network_timeout".to_owned()
486 } else {
487 "network_error".to_owned()
488 }),
489 error_message: Some(error.to_string()),
490 retryable: true,
491 },
492 }
493}
494
495pub(super) fn transport_probe_result(
496 profile: StoredModelProfile,
497 started: Instant,
498 checked_at_ms: u64,
499 error: QosHttpClientError,
500) -> ModelConnectivityProbeResult {
501 ModelConnectivityProbeResult {
502 ok: false,
503 provider: profile.provider,
504 model: profile.model,
505 latency_ms: elapsed_millis(started),
506 checked_at_ms,
507 diagnostics: ModelConnectivityDiagnostics {
508 endpoint_reachable: false,
509 auth_valid: false,
510 rate_limited: false,
511 },
512 token_usage: None,
513 error_code: Some(if error.is_timeout() {
514 "network_timeout".to_owned()
515 } else {
516 "network_error".to_owned()
517 }),
518 error_message: Some(error.to_string()),
519 retryable: true,
520 }
521}
522
523pub(super) fn unsupported_probe(
524 profile: StoredModelProfile,
525 started: Instant,
526 checked_at_ms: u64,
527) -> ModelConnectivityProbeResult {
528 ModelConnectivityProbeResult {
529 ok: false,
530 provider: profile.provider,
531 model: profile.model,
532 latency_ms: elapsed_millis(started),
533 checked_at_ms,
534 diagnostics: ModelConnectivityDiagnostics {
535 endpoint_reachable: false,
536 auth_valid: false,
537 rate_limited: false,
538 },
539 token_usage: None,
540 error_code: Some("unsupported_auth_source".to_owned()),
541 error_message: Some(
542 "this provider requires enterprise auth not configured in relay-knowledge".to_owned(),
543 ),
544 retryable: false,
545 }
546}
547
548pub(super) fn unsupported_discovery(
549 profile: StoredModelProfile,
550 started: Instant,
551 checked_at_ms: u64,
552) -> ModelDiscoveryResult {
553 ModelDiscoveryResult {
554 ok: false,
555 provider: profile.provider,
556 base_url: redacted_url(&profile.base_url),
557 latency_ms: elapsed_millis(started),
558 checked_at_ms,
559 diagnostics: ModelConnectivityDiagnostics {
560 endpoint_reachable: false,
561 auth_valid: false,
562 rate_limited: false,
563 },
564 models: Vec::new(),
565 model_entries: Vec::new(),
566 error_code: Some("unsupported_auth_source".to_owned()),
567 error_message: Some(
568 "this provider requires enterprise auth not configured in relay-knowledge".to_owned(),
569 ),
570 retryable: false,
571 }
572}
573
574pub(super) fn token_usage(payload: &Value) -> Option<ModelConnectivityTokenUsage> {
575 let usage = payload.get("usage")?;
576 Some(ModelConnectivityTokenUsage {
577 prompt_tokens: usage
578 .get("prompt_tokens")
579 .and_then(Value::as_u64)
580 .unwrap_or(0),
581 completion_tokens: usage
582 .get("completion_tokens")
583 .and_then(Value::as_u64)
584 .unwrap_or(0),
585 total_tokens: usage
586 .get("total_tokens")
587 .and_then(Value::as_u64)
588 .unwrap_or(0),
589 })
590}
591
592pub(super) fn parse_discovery_entries(payload: &Value) -> Vec<ModelDiscoveryEntry> {
593 payload
594 .get("data")
595 .and_then(Value::as_array)
596 .into_iter()
597 .flatten()
598 .filter_map(|entry| {
599 let model = entry
600 .get("id")
601 .or_else(|| entry.get("name"))
602 .and_then(Value::as_str)?;
603 Some(ModelDiscoveryEntry {
604 model: model.to_owned(),
605 context_window: entry
606 .get("context_window")
607 .and_then(Value::as_u64)
608 .and_then(|value| u32::try_from(value).ok()),
609 output_limit: entry
610 .get("output_limit")
611 .and_then(Value::as_u64)
612 .and_then(|value| u32::try_from(value).ok()),
613 capabilities: ModelCapabilities::default(),
614 })
615 })
616 .collect()
617}
618
619pub(super) fn diagnostics_from_status(status: u16) -> ModelConnectivityDiagnostics {
620 ModelConnectivityDiagnostics {
621 endpoint_reachable: true,
622 auth_valid: status != 401 && status != 403,
623 rate_limited: status == 429,
624 }
625}
626
627pub(super) fn status_error_code(status: u16) -> &'static str {
628 match status {
629 401 | 403 => "auth_failed",
630 408 | 504 => "network_timeout",
631 429 => "rate_limited",
632 500..=599 => "provider_error",
633 _ => "http_error",
634 }
635}
636
637pub(super) fn is_retryable_status(status: u16) -> bool {
638 matches!(status, 408 | 429 | 500..=599)
639}
640
641pub(super) fn ok_diagnostics() -> ModelConnectivityDiagnostics {
642 ModelConnectivityDiagnostics {
643 endpoint_reachable: true,
644 auth_valid: true,
645 rate_limited: false,
646 }
647}
648
649pub(super) fn request_timeout_from_ms(timeout_ms: Option<u64>) -> Option<Duration> {
650 timeout_ms.map(Duration::from_millis)
651}
652
653pub(super) fn elapsed_millis(started: Instant) -> u64 {
654 u64::try_from(started.elapsed().as_millis()).unwrap_or(u64::MAX)
655}
656
657#[cfg(test)]
658#[path = "connectivity_tests.rs"]
659mod tests;