Skip to main content

relay_knowledge/model_provider/
connectivity.rs

1//! Owns provider HTTP requests, protocol responses, and connectivity diagnostics.
2
3use 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/// Request for profile-aware model connectivity checks.
18#[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/// Request for profile-aware model discovery.
26#[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/// Token counts reported by providers that include usage metadata.
34#[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/// Provider connectivity diagnostics safe for Web display.
42#[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/// Result of a provider probe request.
50#[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/// Discovered provider model with optional metadata.
65#[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/// Result of a model discovery request.
74#[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;