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};
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/// Request for profile-aware model connectivity checks.
19#[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/// Request for profile-aware model discovery.
27#[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/// Token counts reported by providers that include usage metadata.
35#[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/// Provider connectivity diagnostics safe for Web display.
43#[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/// Result of a provider probe request.
51#[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/// Discovered provider model with optional metadata.
66#[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/// Result of a model discovery request.
75#[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;