Skip to main content

agentic_core/
proxy.rs

1use std::pin::Pin;
2use std::time::Duration;
3
4use bytes::Bytes;
5use futures::{Stream, TryStreamExt};
6use http::{HeaderMap, HeaderName, HeaderValue, StatusCode};
7use reqwest::Client;
8use serde_json::Value;
9use tracing::warn;
10
11use crate::config::Config;
12use crate::error::Error;
13
14const HOP_BY_HOP: &[&str] = &[
15    "connection",
16    "keep-alive",
17    "proxy-authenticate",
18    "proxy-authorization",
19    "te",
20    "trailers",
21    "transfer-encoding",
22    "upgrade",
23];
24
25const REQUEST_DROP_EXTRA: &[&str] = &["host", "content-length"];
26
27fn is_hop_by_hop(name: &str) -> bool {
28    HOP_BY_HOP.iter().any(|h| h.eq_ignore_ascii_case(name))
29}
30
31fn is_request_drop(name: &str) -> bool {
32    is_hop_by_hop(name) || REQUEST_DROP_EXTRA.iter().any(|h| h.eq_ignore_ascii_case(name))
33}
34
35/// Raw request data forwarded to the default Responses upstream endpoint.
36pub struct ProxyRequest {
37    pub headers: HeaderMap,
38    pub body: Bytes,
39    pub query: Option<String>,
40}
41
42/// Authentication fallback used when proxying a request to an upstream API.
43#[derive(Clone, Copy, Debug, Eq, PartialEq)]
44pub enum ProxyAuth {
45    OpenAiBearer,
46    Anthropic,
47}
48
49pub enum ProxyBody {
50    Full(Bytes),
51    Stream(Pin<Box<dyn Stream<Item = Result<Bytes, std::io::Error>> + Send>>),
52}
53
54pub struct ProxyResponse {
55    pub status: StatusCode,
56    pub headers: HeaderMap,
57    pub body: ProxyBody,
58}
59
60#[derive(Clone)]
61pub struct ProxyState {
62    pub config: Config,
63    pub stream_client: Client,
64    pub non_stream_client: Client,
65}
66
67impl ProxyState {
68    /// # Errors
69    ///
70    /// Returns an error if the HTTP clients cannot be built.
71    pub fn new(config: Config) -> Result<Self, Error> {
72        let stream_client = Client::builder()
73            .connect_timeout(Duration::from_secs(10))
74            .timeout(Duration::from_secs(900))
75            .pool_max_idle_per_host(0)
76            .redirect(reqwest::redirect::Policy::none())
77            .build()
78            .map_err(Error::HttpClient)?;
79
80        let non_stream_client = Client::builder()
81            .connect_timeout(Duration::from_secs(10))
82            .read_timeout(Duration::from_secs(300))
83            .redirect(reqwest::redirect::Policy::none())
84            .build()
85            .map_err(Error::HttpClient)?;
86
87        Ok(Self {
88            config,
89            stream_client,
90            non_stream_client,
91        })
92    }
93}
94
95fn filter_request_headers(headers: &HeaderMap, config: &Config, auth: ProxyAuth) -> reqwest::header::HeaderMap {
96    let mut out = reqwest::header::HeaderMap::new();
97    for (name, value) in headers {
98        if is_request_drop(name.as_str()) {
99            continue;
100        }
101        if let Ok(n) = reqwest::header::HeaderName::from_bytes(name.as_str().as_bytes()) {
102            if let Ok(v) = reqwest::header::HeaderValue::from_bytes(value.as_bytes()) {
103                out.append(n, v);
104            }
105        }
106    }
107
108    let has_auth = out.contains_key(reqwest::header::AUTHORIZATION);
109    let has_api_key = out.contains_key("x-api-key");
110    if !has_auth && !has_api_key {
111        if let Some(key) = config.openai_api_key.as_deref() {
112            let trimmed = key.trim();
113            if !trimmed.is_empty() {
114                let (name, value) = match auth {
115                    ProxyAuth::OpenAiBearer => (reqwest::header::AUTHORIZATION, format!("Bearer {trimmed}")),
116                    ProxyAuth::Anthropic => (
117                        reqwest::header::HeaderName::from_static("x-api-key"),
118                        trimmed.to_owned(),
119                    ),
120                };
121                if let Ok(v) = reqwest::header::HeaderValue::from_str(&value) {
122                    out.insert(name, v);
123                }
124            }
125        }
126    }
127
128    out
129}
130
131fn filter_response_headers(headers: &reqwest::header::HeaderMap) -> HeaderMap {
132    let mut out = HeaderMap::new();
133    for (name, value) in headers {
134        if is_hop_by_hop(name.as_str()) {
135            continue;
136        }
137        if let Ok(n) = HeaderName::from_bytes(name.as_str().as_bytes()) {
138            if let Ok(v) = HeaderValue::from_bytes(value.as_bytes()) {
139                out.append(n, v);
140            }
141        }
142    }
143    out
144}
145
146fn is_sse_content_type(headers: &reqwest::header::HeaderMap) -> bool {
147    headers
148        .get(reqwest::header::CONTENT_TYPE)
149        .and_then(|v| v.to_str().ok())
150        .is_some_and(|ct| ct.to_ascii_lowercase().starts_with("text/event-stream"))
151}
152
153#[must_use]
154pub fn error_response(status: StatusCode, code: &str, message: &str) -> ProxyResponse {
155    error_response_for_auth(status, code, message, ProxyAuth::OpenAiBearer)
156}
157
158#[must_use]
159pub fn error_response_for_auth(status: StatusCode, code: &str, message: &str, auth: ProxyAuth) -> ProxyResponse {
160    let body = match auth {
161        ProxyAuth::OpenAiBearer => serde_json::json!({
162            "error": {
163                "message": message,
164                "type": "api_error",
165                "param": null,
166                "code": code,
167            }
168        }),
169        ProxyAuth::Anthropic => serde_json::json!({
170            "type": "error",
171            "error": {
172                "type": "api_error",
173                "message": message,
174            }
175        }),
176    };
177    let mut headers = HeaderMap::new();
178    headers.insert("content-type", HeaderValue::from_static("application/json"));
179    ProxyResponse {
180        status,
181        headers,
182        body: ProxyBody::Full(Bytes::from(serde_json::to_vec(&body).unwrap_or_default())),
183    }
184}
185
186/// Proxy a GET request to an arbitrary upstream path.
187///
188/// Applies the same header filtering and auth injection as [`proxy_request`].
189/// Uses the non-streaming client; the response body is returned as a full
190/// [`ProxyBody::Full`] payload.
191pub async fn proxy_get(path: &str, request_headers: &HeaderMap, state: &ProxyState) -> ProxyResponse {
192    let llm_headers = filter_request_headers(request_headers, &state.config, ProxyAuth::OpenAiBearer);
193    let base = state.config.llm_api_base.trim_end_matches('/');
194    let url = format!("{base}/{}", path.trim_start_matches('/'));
195
196    let llm_resp = match state.non_stream_client.get(&url).headers(llm_headers).send().await {
197        Ok(r) => r,
198        Err(e) if e.is_timeout() => {
199            warn!("upstream GET {path} timed out: {e}");
200            return error_response(StatusCode::GATEWAY_TIMEOUT, "upstream_timeout", "upstream timeout");
201        }
202        Err(e) => {
203            warn!("upstream GET {path} failed: {e}");
204            return error_response(StatusCode::BAD_GATEWAY, "upstream_unavailable", "upstream unavailable");
205        }
206    };
207
208    let status = StatusCode::from_u16(llm_resp.status().as_u16()).unwrap_or(StatusCode::BAD_GATEWAY);
209    let response_headers = filter_response_headers(llm_resp.headers());
210
211    match llm_resp.bytes().await {
212        Ok(payload) => ProxyResponse {
213            status,
214            headers: response_headers,
215            body: ProxyBody::Full(payload),
216        },
217        Err(e) => {
218            warn!("failed to read upstream GET {path} body: {e}");
219            error_response(
220                StatusCode::BAD_GATEWAY,
221                "upstream_unavailable",
222                "failed to read upstream response",
223            )
224        }
225    }
226}
227
228/// Proxy a request to the default `/v1/responses` upstream endpoint.
229pub async fn proxy_request(request: ProxyRequest, state: &ProxyState) -> ProxyResponse {
230    proxy_request_with_path(request, "/v1/responses", ProxyAuth::OpenAiBearer, state).await
231}
232
233/// Proxy a raw request to a selected upstream path.
234pub async fn proxy_request_with_path(
235    request: ProxyRequest,
236    path: &str,
237    auth: ProxyAuth,
238    state: &ProxyState,
239) -> ProxyResponse {
240    let is_streaming = serde_json::from_slice::<Value>(&request.body)
241        .ok()
242        .and_then(|v| v.get("stream")?.as_bool())
243        .unwrap_or(false);
244
245    let llm_headers = filter_request_headers(&request.headers, &state.config, auth);
246
247    let base = state.config.llm_api_base.trim_end_matches('/');
248    let mut url = format!("{base}/{}", path.trim_start_matches('/'));
249    if let Some(q) = &request.query {
250        url.push('?');
251        url.push_str(q);
252    }
253
254    let client = if is_streaming {
255        &state.stream_client
256    } else {
257        &state.non_stream_client
258    };
259
260    let llm_resp = match client.post(&url).headers(llm_headers).body(request.body).send().await {
261        Ok(r) => r,
262        Err(e) if e.is_timeout() => {
263            warn!("LLM request timed out: {e}");
264            return error_response_for_auth(StatusCode::GATEWAY_TIMEOUT, "llm_timeout", "LLM timeout", auth);
265        }
266        Err(e) => {
267            warn!("LLM request failed: {e}");
268            return error_response_for_auth(StatusCode::BAD_GATEWAY, "llm_unavailable", "LLM unavailable", auth);
269        }
270    };
271
272    let status = StatusCode::from_u16(llm_resp.status().as_u16()).unwrap_or(StatusCode::BAD_GATEWAY);
273    let mut response_headers = filter_response_headers(llm_resp.headers());
274
275    if is_sse_content_type(llm_resp.headers()) {
276        response_headers.insert("x-accel-buffering", HeaderValue::from_static("no"));
277
278        let byte_stream = llm_resp.bytes_stream().map_err(std::io::Error::other);
279
280        return ProxyResponse {
281            status,
282            headers: response_headers,
283            body: ProxyBody::Stream(Box::pin(byte_stream)),
284        };
285    }
286
287    let payload: Bytes = match llm_resp.bytes().await {
288        Ok(b) => b,
289        Err(e) => {
290            warn!("failed to read LLM response body: {e}");
291            return error_response_for_auth(
292                StatusCode::BAD_GATEWAY,
293                "llm_unavailable",
294                "Failed to read LLM response",
295                auth,
296            );
297        }
298    };
299
300    ProxyResponse {
301        status,
302        headers: response_headers,
303        body: ProxyBody::Full(payload),
304    }
305}
306
307#[cfg(test)]
308mod tests {
309    use super::*;
310    use crate::config::Config;
311
312    fn test_config() -> Config {
313        Config {
314            llm_api_base: "http://localhost:8000".to_owned(),
315            openai_api_key: Some("test-key".to_owned()),
316            llm_ready_timeout_s: 5.0,
317            llm_ready_interval_s: 0.1,
318            skip_llm_ready_check: false,
319            db_url: None,
320            postgres: crate::config::PostgresConfig::default(),
321            sqlite: crate::config::SqliteConfig::default(),
322        }
323    }
324
325    fn test_config_no_key() -> Config {
326        Config {
327            openai_api_key: None,
328            ..test_config()
329        }
330    }
331
332    #[test]
333    fn hop_by_hop_detected() {
334        assert!(is_hop_by_hop("connection"));
335        assert!(is_hop_by_hop("Connection"));
336        assert!(is_hop_by_hop("keep-alive"));
337        assert!(is_hop_by_hop("transfer-encoding"));
338        assert!(is_hop_by_hop("proxy-authorization"));
339    }
340
341    #[test]
342    fn non_hop_by_hop_passes() {
343        assert!(!is_hop_by_hop("content-type"));
344        assert!(!is_hop_by_hop("x-custom"));
345        assert!(!is_hop_by_hop("authorization"));
346    }
347
348    #[test]
349    fn request_drop_includes_host_and_content_length() {
350        assert!(is_request_drop("host"));
351        assert!(is_request_drop("content-length"));
352        assert!(is_request_drop("connection"));
353        assert!(!is_request_drop("content-type"));
354    }
355
356    #[test]
357    fn proxy_request_retains_legacy_construction_shape() {
358        let _request = ProxyRequest {
359            headers: HeaderMap::new(),
360            body: Bytes::new(),
361            query: None,
362        };
363    }
364
365    #[test]
366    fn filter_request_headers_strips_hop_by_hop() {
367        let mut headers = HeaderMap::new();
368        headers.insert("content-type", "application/json".parse().unwrap());
369        headers.insert("connection", "keep-alive".parse().unwrap());
370        headers.insert("proxy-authorization", "Basic abc".parse().unwrap());
371        headers.insert("x-custom", "value".parse().unwrap());
372
373        let config = test_config_no_key();
374        let filtered = filter_request_headers(&headers, &config, ProxyAuth::OpenAiBearer);
375
376        assert!(filtered.contains_key("content-type"));
377        assert!(filtered.contains_key("x-custom"));
378        assert!(!filtered.contains_key("connection"));
379        assert!(!filtered.contains_key("proxy-authorization"));
380    }
381
382    #[test]
383    fn filter_request_headers_strips_host_and_content_length() {
384        let mut headers = HeaderMap::new();
385        headers.insert("host", "example.com".parse().unwrap());
386        headers.insert("content-length", "42".parse().unwrap());
387        headers.insert("accept", "*/*".parse().unwrap());
388
389        let config = test_config_no_key();
390        let filtered = filter_request_headers(&headers, &config, ProxyAuth::OpenAiBearer);
391
392        assert!(!filtered.contains_key("host"));
393        assert!(!filtered.contains_key("content-length"));
394        assert!(filtered.contains_key("accept"));
395    }
396
397    #[test]
398    fn auth_injected_when_no_client_auth() {
399        let headers = HeaderMap::new();
400        let config = test_config();
401        let filtered = filter_request_headers(&headers, &config, ProxyAuth::OpenAiBearer);
402
403        assert_eq!(
404            filtered.get("authorization").unwrap().to_str().unwrap(),
405            "Bearer test-key"
406        );
407    }
408
409    #[test]
410    fn client_auth_takes_precedence() {
411        let mut headers = HeaderMap::new();
412        headers.insert("authorization", "Bearer client-token".parse().unwrap());
413
414        let config = test_config();
415        let filtered = filter_request_headers(&headers, &config, ProxyAuth::OpenAiBearer);
416
417        assert_eq!(
418            filtered.get("authorization").unwrap().to_str().unwrap(),
419            "Bearer client-token"
420        );
421    }
422
423    #[test]
424    fn anthropic_auth_preserves_client_api_key() {
425        let mut headers = HeaderMap::new();
426        headers.insert("x-api-key", "client-anthropic-key".parse().unwrap());
427
428        let filtered = filter_request_headers(&headers, &test_config(), ProxyAuth::Anthropic);
429
430        assert_eq!(filtered.get("x-api-key").unwrap(), "client-anthropic-key");
431        assert!(!filtered.contains_key("authorization"));
432    }
433
434    #[test]
435    fn anthropic_auth_uses_configured_key_as_api_key_fallback() {
436        let filtered = filter_request_headers(&HeaderMap::new(), &test_config(), ProxyAuth::Anthropic);
437
438        assert_eq!(filtered.get("x-api-key").unwrap(), "test-key");
439        assert!(!filtered.contains_key("authorization"));
440    }
441
442    #[test]
443    fn no_auth_injected_when_key_empty() {
444        let headers = HeaderMap::new();
445        let config = Config {
446            openai_api_key: Some("  ".to_owned()),
447            ..test_config()
448        };
449        let filtered = filter_request_headers(&headers, &config, ProxyAuth::OpenAiBearer);
450
451        assert!(!filtered.contains_key("authorization"));
452    }
453
454    #[test]
455    fn no_auth_injected_when_key_none() {
456        let headers = HeaderMap::new();
457        let config = test_config_no_key();
458        let filtered = filter_request_headers(&headers, &config, ProxyAuth::OpenAiBearer);
459
460        assert!(!filtered.contains_key("authorization"));
461    }
462
463    #[test]
464    fn filter_response_headers_strips_hop_by_hop() {
465        let mut headers = reqwest::header::HeaderMap::new();
466        headers.insert("content-type", "application/json".parse().unwrap());
467        headers.insert("connection", "keep-alive".parse().unwrap());
468        headers.insert("x-request-id", "abc".parse().unwrap());
469
470        let filtered = filter_response_headers(&headers);
471
472        assert!(filtered.contains_key("content-type"));
473        assert!(filtered.contains_key("x-request-id"));
474        assert!(!filtered.contains_key("connection"));
475    }
476
477    #[test]
478    fn sse_content_type_detected() {
479        let mut headers = reqwest::header::HeaderMap::new();
480        headers.insert("content-type", "text/event-stream; charset=utf-8".parse().unwrap());
481        assert!(is_sse_content_type(&headers));
482    }
483
484    #[test]
485    fn sse_content_type_case_insensitive() {
486        let mut headers = reqwest::header::HeaderMap::new();
487        headers.insert("content-type", "Text/Event-Stream".parse().unwrap());
488        assert!(is_sse_content_type(&headers));
489    }
490
491    #[test]
492    fn non_sse_content_type_rejected() {
493        let mut headers = reqwest::header::HeaderMap::new();
494        headers.insert("content-type", "application/json".parse().unwrap());
495        assert!(!is_sse_content_type(&headers));
496    }
497
498    #[test]
499    fn missing_content_type_not_sse() {
500        let headers = reqwest::header::HeaderMap::new();
501        assert!(!is_sse_content_type(&headers));
502    }
503}