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