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