Skip to main content

postrust_proxy/vendored/
handler.rs

1//! Vendored message handler from rpxy-lib: message_handler/*.rs
2//!
3//! This module handles request/response manipulation:
4//! - X-Forwarded-* header handling
5//! - Host header rewriting
6//! - Request parsing
7
8use crate::config::Route;
9use crate::vendored::hyper_ext::{empty_body, string_body, ProxyBody};
10use hyper::header::{HeaderName, HeaderValue, HOST};
11use hyper::{Request, Response, StatusCode};
12use std::net::SocketAddr;
13
14/// Standard X-Forwarded headers.
15pub mod headers {
16    use hyper::header::HeaderName;
17
18    pub static X_FORWARDED_FOR: HeaderName = HeaderName::from_static("x-forwarded-for");
19    pub static X_FORWARDED_PROTO: HeaderName = HeaderName::from_static("x-forwarded-proto");
20    pub static X_FORWARDED_HOST: HeaderName = HeaderName::from_static("x-forwarded-host");
21    pub static X_REAL_IP: HeaderName = HeaderName::from_static("x-real-ip");
22}
23
24/// Message handler for request/response manipulation.
25pub struct MessageHandler;
26
27impl MessageHandler {
28    /// Add X-Forwarded-* headers to a request.
29    pub fn add_forwarding_headers(
30        request: &mut Request<ProxyBody>,
31        client_addr: SocketAddr,
32        proto: &str,
33    ) {
34        let headers = request.headers_mut();
35        let client_ip = client_addr.ip().to_string();
36
37        // X-Forwarded-For: append client IP
38        if let Some(existing) = headers.get(&headers::X_FORWARDED_FOR) {
39            let mut new_value = existing.to_str().unwrap_or("").to_string();
40            new_value.push_str(", ");
41            new_value.push_str(&client_ip);
42            if let Ok(value) = HeaderValue::from_str(&new_value) {
43                headers.insert(headers::X_FORWARDED_FOR.clone(), value);
44            }
45        } else if let Ok(value) = HeaderValue::from_str(&client_ip) {
46            headers.insert(headers::X_FORWARDED_FOR.clone(), value);
47        }
48
49        // X-Real-IP: set if not present
50        if !headers.contains_key(&headers::X_REAL_IP) {
51            if let Ok(value) = HeaderValue::from_str(&client_ip) {
52                headers.insert(headers::X_REAL_IP.clone(), value);
53            }
54        }
55
56        // X-Forwarded-Proto
57        if let Ok(value) = HeaderValue::from_str(proto) {
58            headers.insert(headers::X_FORWARDED_PROTO.clone(), value);
59        }
60
61        // X-Forwarded-Host: preserve original Host
62        if let Some(host) = headers.get(HOST).cloned() {
63            headers.insert(headers::X_FORWARDED_HOST.clone(), host);
64        }
65    }
66
67    /// Rewrite the Host header for upstream.
68    pub fn rewrite_host_header(request: &mut Request<ProxyBody>, upstream_host: &str) {
69        if let Ok(value) = HeaderValue::from_str(upstream_host) {
70            request.headers_mut().insert(HOST, value);
71        }
72    }
73
74    /// Strip path prefix from request URI.
75    pub fn strip_path_prefix(request: &mut Request<ProxyBody>, prefix: &str) {
76        let uri = request.uri();
77        let path = uri.path();
78
79        if let Some(new_path) = path.strip_prefix(prefix) {
80            let new_path = if new_path.is_empty() || !new_path.starts_with('/') {
81                format!("/{}", new_path.trim_start_matches('/'))
82            } else {
83                new_path.to_string()
84            };
85
86            let new_uri = if let Some(query) = uri.query() {
87                format!("{}?{}", new_path, query)
88            } else {
89                new_path
90            };
91
92            if let Ok(new_uri) = new_uri.parse() {
93                *request.uri_mut() = new_uri;
94            }
95        }
96    }
97
98    /// Apply route-specific header modifications.
99    pub fn apply_route_headers(request: &mut Request<ProxyBody>, route: &Route) {
100        let headers = request.headers_mut();
101
102        // Add headers
103        for (name, value) in &route.add_headers {
104            if let (Ok(name), Ok(value)) = (
105                HeaderName::from_bytes(name.as_bytes()),
106                HeaderValue::from_str(value),
107            ) {
108                headers.insert(name, value);
109            }
110        }
111
112        // Remove headers
113        for name in &route.remove_headers {
114            if let Ok(name) = HeaderName::from_bytes(name.as_bytes()) {
115                headers.remove(name);
116            }
117        }
118    }
119
120    /// Create a synthetic error response.
121    pub fn error_response(status: StatusCode, message: &str) -> Response<ProxyBody> {
122        Response::builder()
123            .status(status)
124            .header("content-type", "text/plain; charset=utf-8")
125            .body(string_body(message.to_string()))
126            .unwrap_or_else(|_| {
127                Response::builder()
128                    .status(StatusCode::INTERNAL_SERVER_ERROR)
129                    .body(empty_body())
130                    .unwrap()
131            })
132    }
133
134    /// Create a 502 Bad Gateway response.
135    pub fn bad_gateway(message: &str) -> Response<ProxyBody> {
136        Self::error_response(StatusCode::BAD_GATEWAY, message)
137    }
138
139    /// Create a 503 Service Unavailable response.
140    pub fn service_unavailable(message: &str) -> Response<ProxyBody> {
141        Self::error_response(StatusCode::SERVICE_UNAVAILABLE, message)
142    }
143
144    /// Create a 504 Gateway Timeout response.
145    pub fn gateway_timeout() -> Response<ProxyBody> {
146        Self::error_response(StatusCode::GATEWAY_TIMEOUT, "Gateway Timeout")
147    }
148
149    /// Create a 429 Too Many Requests response.
150    pub fn too_many_requests() -> Response<ProxyBody> {
151        Self::error_response(StatusCode::TOO_MANY_REQUESTS, "Too Many Requests")
152    }
153
154    /// Create a 404 Not Found response.
155    pub fn not_found() -> Response<ProxyBody> {
156        Self::error_response(StatusCode::NOT_FOUND, "Not Found")
157    }
158}
159
160#[cfg(test)]
161mod tests {
162    use super::*;
163    use crate::vendored::hyper_ext::empty_body;
164
165    #[test]
166    fn test_strip_path_prefix() {
167        let mut request = Request::builder()
168            .uri("/api/v1/users")
169            .body(empty_body())
170            .unwrap();
171
172        MessageHandler::strip_path_prefix(&mut request, "/api/v1");
173
174        assert_eq!(request.uri().path(), "/users");
175    }
176
177    #[test]
178    fn test_strip_path_prefix_root() {
179        let mut request = Request::builder().uri("/api").body(empty_body()).unwrap();
180
181        MessageHandler::strip_path_prefix(&mut request, "/api");
182
183        assert_eq!(request.uri().path(), "/");
184    }
185
186    #[test]
187    fn test_forwarding_headers() {
188        let mut request = Request::builder()
189            .uri("/test")
190            .header("host", "example.com")
191            .body(empty_body())
192            .unwrap();
193
194        let client_addr: SocketAddr = "192.168.1.100:12345".parse().unwrap();
195        MessageHandler::add_forwarding_headers(&mut request, client_addr, "https");
196
197        assert_eq!(
198            request.headers().get("x-forwarded-for").unwrap(),
199            "192.168.1.100"
200        );
201        assert_eq!(request.headers().get("x-forwarded-proto").unwrap(), "https");
202        assert_eq!(
203            request.headers().get("x-forwarded-host").unwrap(),
204            "example.com"
205        );
206    }
207}