Skip to main content

codex_api/
provider.rs

1use codex_client::Request;
2use codex_client::RequestCompression;
3use codex_client::RetryOn;
4use codex_client::RetryPolicy;
5use http::Method;
6use http::header::HeaderMap;
7use std::collections::HashMap;
8use std::time::Duration;
9use url::Url;
10
11/// High-level retry configuration for a provider.
12///
13/// This is converted into a `RetryPolicy` used by `codex-client` to drive
14/// transport-level retries for both unary and streaming calls.
15#[derive(Debug, Clone)]
16pub struct RetryConfig {
17    pub max_attempts: u64,
18    pub base_delay: Duration,
19    pub retry_429: bool,
20    pub retry_5xx: bool,
21    pub retry_transport: bool,
22}
23
24impl RetryConfig {
25    pub fn to_policy(&self) -> RetryPolicy {
26        RetryPolicy {
27            max_attempts: self.max_attempts,
28            base_delay: self.base_delay,
29            retry_on: RetryOn {
30                retry_429: self.retry_429,
31                retry_5xx: self.retry_5xx,
32                retry_transport: self.retry_transport,
33            },
34        }
35    }
36}
37
38/// HTTP endpoint configuration used to talk to a concrete API deployment.
39///
40/// Encapsulates base URL, default headers, query params, retry policy, and
41/// stream idle timeout, plus helper methods for building requests.
42#[derive(Debug, Clone)]
43pub struct Provider {
44    pub name: String,
45    pub base_url: String,
46    pub query_params: Option<HashMap<String, String>>,
47    pub headers: HeaderMap,
48    pub retry: RetryConfig,
49    pub stream_idle_timeout: Duration,
50}
51
52impl Provider {
53    pub fn url_for_path(&self, path: &str) -> String {
54        let base = self.base_url.trim_end_matches('/');
55        let path = path.trim_start_matches('/');
56        let mut url = if path.is_empty() {
57            base.to_string()
58        } else {
59            format!("{base}/{path}")
60        };
61
62        if let Some(params) = &self.query_params
63            && !params.is_empty()
64        {
65            let qs = params
66                .iter()
67                .map(|(k, v)| format!("{k}={v}"))
68                .collect::<Vec<_>>()
69                .join("&");
70            url.push('?');
71            url.push_str(&qs);
72        }
73
74        url
75    }
76
77    pub fn build_request(&self, method: Method, path: &str) -> Request {
78        Request {
79            method,
80            url: self.url_for_path(path),
81            headers: self.headers.clone(),
82            body: None,
83            compression: RequestCompression::None,
84            timeout: None,
85        }
86    }
87
88    pub fn is_azure_responses_endpoint(&self) -> bool {
89        is_azure_responses_provider(&self.name, Some(&self.base_url))
90    }
91
92    pub fn websocket_url_for_path(&self, path: &str) -> Result<Url, url::ParseError> {
93        let mut url = Url::parse(&self.url_for_path(path))?;
94
95        let scheme = match url.scheme() {
96            "http" => "ws",
97            "https" => "wss",
98            "ws" | "wss" => return Ok(url),
99            _ => return Ok(url),
100        };
101        let _ = url.set_scheme(scheme);
102        Ok(url)
103    }
104}
105
106pub fn is_azure_responses_provider(name: &str, base_url: Option<&str>) -> bool {
107    if name.eq_ignore_ascii_case("azure") {
108        true
109    } else if let Some(base_url) = base_url {
110        matches_azure_responses_base_url(base_url)
111    } else {
112        false
113    }
114}
115
116fn matches_azure_responses_base_url(base_url: &str) -> bool {
117    let base_url = base_url.to_ascii_lowercase();
118    const AZURE_MARKERS: [&str; 6] = [
119        "openai.azure.",
120        "cognitiveservices.azure.",
121        "aoai.azure.",
122        "azure-api.",
123        "azurefd.",
124        "windows.net/openai",
125    ];
126    AZURE_MARKERS.iter().any(|marker| base_url.contains(marker))
127}
128
129#[cfg(test)]
130mod tests {
131    use super::*;
132
133    #[test]
134    fn detects_azure_responses_base_urls() {
135        let positive_cases = [
136            "https://foo.openai.azure.com/openai",
137            "https://foo.openai.azure.us/openai/deployments/bar",
138            "https://foo.cognitiveservices.azure.cn/openai",
139            "https://foo.aoai.azure.com/openai",
140            "https://foo.openai.azure-api.net/openai",
141            "https://foo.z01.azurefd.net/",
142        ];
143
144        for base_url in positive_cases {
145            assert!(
146                is_azure_responses_provider("test", Some(base_url)),
147                "expected {base_url} to be detected as Azure"
148            );
149        }
150
151        assert!(is_azure_responses_provider(
152            "Azure",
153            Some("https://example.com")
154        ));
155
156        let negative_cases = [
157            "https://api.openai.com/v1",
158            "https://example.com/openai",
159            "https://myproxy.azurewebsites.net/openai",
160        ];
161
162        for base_url in negative_cases {
163            assert!(
164                !is_azure_responses_provider("test", Some(base_url)),
165                "expected {base_url} not to be detected as Azure"
166            );
167        }
168    }
169}