postrust_proxy/vendored/
handler.rs1use 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
14pub 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
24pub struct MessageHandler;
26
27impl MessageHandler {
28 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 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 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 if let Ok(value) = HeaderValue::from_str(proto) {
58 headers.insert(headers::X_FORWARDED_PROTO.clone(), value);
59 }
60
61 if let Some(host) = headers.get(HOST).cloned() {
63 headers.insert(headers::X_FORWARDED_HOST.clone(), host);
64 }
65 }
66
67 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 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 pub fn apply_route_headers(request: &mut Request<ProxyBody>, route: &Route) {
100 let headers = request.headers_mut();
101
102 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 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 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 pub fn bad_gateway(message: &str) -> Response<ProxyBody> {
136 Self::error_response(StatusCode::BAD_GATEWAY, message)
137 }
138
139 pub fn service_unavailable(message: &str) -> Response<ProxyBody> {
141 Self::error_response(StatusCode::SERVICE_UNAVAILABLE, message)
142 }
143
144 pub fn gateway_timeout() -> Response<ProxyBody> {
146 Self::error_response(StatusCode::GATEWAY_TIMEOUT, "Gateway Timeout")
147 }
148
149 pub fn too_many_requests() -> Response<ProxyBody> {
151 Self::error_response(StatusCode::TOO_MANY_REQUESTS, "Too Many Requests")
152 }
153
154 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}