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#[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#[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}