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
35pub struct ProxyRequest {
37 pub headers: HeaderMap,
38 pub body: Bytes,
39 pub query: Option<String>,
40}
41
42#[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 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
186pub 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
228pub async fn proxy_request(request: ProxyRequest, state: &ProxyState) -> ProxyResponse {
230 proxy_request_with_path(request, "/v1/responses", ProxyAuth::OpenAiBearer, state).await
231}
232
233pub 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 tools: crate::config::ToolRuntimeConfig::default(),
323 }
324 }
325
326 fn test_config_no_key() -> Config {
327 Config {
328 openai_api_key: None,
329 ..test_config()
330 }
331 }
332
333 #[test]
334 fn hop_by_hop_detected() {
335 assert!(is_hop_by_hop("connection"));
336 assert!(is_hop_by_hop("Connection"));
337 assert!(is_hop_by_hop("keep-alive"));
338 assert!(is_hop_by_hop("transfer-encoding"));
339 assert!(is_hop_by_hop("proxy-authorization"));
340 }
341
342 #[test]
343 fn non_hop_by_hop_passes() {
344 assert!(!is_hop_by_hop("content-type"));
345 assert!(!is_hop_by_hop("x-custom"));
346 assert!(!is_hop_by_hop("authorization"));
347 }
348
349 #[test]
350 fn request_drop_includes_host_and_content_length() {
351 assert!(is_request_drop("host"));
352 assert!(is_request_drop("content-length"));
353 assert!(is_request_drop("connection"));
354 assert!(!is_request_drop("content-type"));
355 }
356
357 #[test]
358 fn proxy_request_retains_legacy_construction_shape() {
359 let _request = ProxyRequest {
360 headers: HeaderMap::new(),
361 body: Bytes::new(),
362 query: None,
363 };
364 }
365
366 #[test]
367 fn filter_request_headers_strips_hop_by_hop() {
368 let mut headers = HeaderMap::new();
369 headers.insert("content-type", "application/json".parse().unwrap());
370 headers.insert("connection", "keep-alive".parse().unwrap());
371 headers.insert("proxy-authorization", "Basic abc".parse().unwrap());
372 headers.insert("x-custom", "value".parse().unwrap());
373
374 let config = test_config_no_key();
375 let filtered = filter_request_headers(&headers, &config, ProxyAuth::OpenAiBearer);
376
377 assert!(filtered.contains_key("content-type"));
378 assert!(filtered.contains_key("x-custom"));
379 assert!(!filtered.contains_key("connection"));
380 assert!(!filtered.contains_key("proxy-authorization"));
381 }
382
383 #[test]
384 fn filter_request_headers_strips_host_and_content_length() {
385 let mut headers = HeaderMap::new();
386 headers.insert("host", "example.com".parse().unwrap());
387 headers.insert("content-length", "42".parse().unwrap());
388 headers.insert("accept", "*/*".parse().unwrap());
389
390 let config = test_config_no_key();
391 let filtered = filter_request_headers(&headers, &config, ProxyAuth::OpenAiBearer);
392
393 assert!(!filtered.contains_key("host"));
394 assert!(!filtered.contains_key("content-length"));
395 assert!(filtered.contains_key("accept"));
396 }
397
398 #[test]
399 fn auth_injected_when_no_client_auth() {
400 let headers = HeaderMap::new();
401 let config = test_config();
402 let filtered = filter_request_headers(&headers, &config, ProxyAuth::OpenAiBearer);
403
404 assert_eq!(
405 filtered.get("authorization").unwrap().to_str().unwrap(),
406 "Bearer test-key"
407 );
408 }
409
410 #[test]
411 fn client_auth_takes_precedence() {
412 let mut headers = HeaderMap::new();
413 headers.insert("authorization", "Bearer client-token".parse().unwrap());
414
415 let config = test_config();
416 let filtered = filter_request_headers(&headers, &config, ProxyAuth::OpenAiBearer);
417
418 assert_eq!(
419 filtered.get("authorization").unwrap().to_str().unwrap(),
420 "Bearer client-token"
421 );
422 }
423
424 #[test]
425 fn anthropic_auth_preserves_client_api_key() {
426 let mut headers = HeaderMap::new();
427 headers.insert("x-api-key", "client-anthropic-key".parse().unwrap());
428
429 let filtered = filter_request_headers(&headers, &test_config(), ProxyAuth::Anthropic);
430
431 assert_eq!(filtered.get("x-api-key").unwrap(), "client-anthropic-key");
432 assert!(!filtered.contains_key("authorization"));
433 }
434
435 #[test]
436 fn anthropic_auth_uses_configured_key_as_api_key_fallback() {
437 let filtered = filter_request_headers(&HeaderMap::new(), &test_config(), ProxyAuth::Anthropic);
438
439 assert_eq!(filtered.get("x-api-key").unwrap(), "test-key");
440 assert!(!filtered.contains_key("authorization"));
441 }
442
443 #[test]
444 fn no_auth_injected_when_key_empty() {
445 let headers = HeaderMap::new();
446 let config = Config {
447 openai_api_key: Some(" ".to_owned()),
448 ..test_config()
449 };
450 let filtered = filter_request_headers(&headers, &config, ProxyAuth::OpenAiBearer);
451
452 assert!(!filtered.contains_key("authorization"));
453 }
454
455 #[test]
456 fn no_auth_injected_when_key_none() {
457 let headers = HeaderMap::new();
458 let config = test_config_no_key();
459 let filtered = filter_request_headers(&headers, &config, ProxyAuth::OpenAiBearer);
460
461 assert!(!filtered.contains_key("authorization"));
462 }
463
464 #[test]
465 fn filter_response_headers_strips_hop_by_hop() {
466 let mut headers = reqwest::header::HeaderMap::new();
467 headers.insert("content-type", "application/json".parse().unwrap());
468 headers.insert("connection", "keep-alive".parse().unwrap());
469 headers.insert("x-request-id", "abc".parse().unwrap());
470
471 let filtered = filter_response_headers(&headers);
472
473 assert!(filtered.contains_key("content-type"));
474 assert!(filtered.contains_key("x-request-id"));
475 assert!(!filtered.contains_key("connection"));
476 }
477
478 #[test]
479 fn sse_content_type_detected() {
480 let mut headers = reqwest::header::HeaderMap::new();
481 headers.insert("content-type", "text/event-stream; charset=utf-8".parse().unwrap());
482 assert!(is_sse_content_type(&headers));
483 }
484
485 #[test]
486 fn sse_content_type_case_insensitive() {
487 let mut headers = reqwest::header::HeaderMap::new();
488 headers.insert("content-type", "Text/Event-Stream".parse().unwrap());
489 assert!(is_sse_content_type(&headers));
490 }
491
492 #[test]
493 fn non_sse_content_type_rejected() {
494 let mut headers = reqwest::header::HeaderMap::new();
495 headers.insert("content-type", "application/json".parse().unwrap());
496 assert!(!is_sse_content_type(&headers));
497 }
498
499 #[test]
500 fn missing_content_type_not_sse() {
501 let headers = reqwest::header::HeaderMap::new();
502 assert!(!is_sse_content_type(&headers));
503 }
504}