postrust_proxy/vendored/
proxy.rs1use crate::config::ProxyConfig;
6use crate::health::HealthChecker;
7use crate::ratelimit::{RateLimitKey, RateLimiter};
8use crate::vendored::backend::BackendAppManager;
9use crate::vendored::forwarder::ForwarderClient;
10use crate::vendored::handler::MessageHandler;
11use crate::vendored::hyper_ext::{IncomingBodyExt, ProxyBody};
12use crate::vendored::types::PathName;
13use hyper::body::Incoming;
14use hyper::server::conn::http1;
15use hyper::service::service_fn;
16use hyper::{Request, Response};
17use std::convert::Infallible;
18use std::net::SocketAddr;
19use std::sync::Arc;
20use tokio::net::TcpListener;
21use tokio::sync::RwLock;
22use tokio_util::sync::CancellationToken;
23use tracing::{debug, error, info, warn};
24
25pub struct ProxyService {
27 backend_manager: Arc<BackendAppManager>,
29 forwarder: Arc<ForwarderClient>,
31 rate_limiter: Arc<RateLimiter>,
33 config: Arc<RwLock<ProxyConfig>>,
35}
36
37impl ProxyService {
38 pub fn new(
40 config: Arc<RwLock<ProxyConfig>>,
41 health_checker: Arc<HealthChecker>,
42 rate_limiter: Arc<RateLimiter>,
43 ) -> Self {
44 let backend_manager =
45 Arc::new(BackendAppManager::new().with_health_checker(health_checker));
46 let forwarder = Arc::new(ForwarderClient::default());
47
48 Self {
49 backend_manager,
50 forwarder,
51 rate_limiter,
52 config,
53 }
54 }
55
56 pub async fn load_config(&self) {
58 let config = self.config.read().await;
59
60 for upstream in &config.upstreams {
62 self.backend_manager.register_upstream(upstream.clone());
63 }
64
65 for route in &config.routes {
67 if let Some(upstream) = config.upstreams.iter().find(|u| u.name == route.upstream) {
69 if let Some(upstream_id) = upstream.id {
70 use crate::vendored::types::ServerName;
71 let host = route.match_.host.as_deref().unwrap_or("*");
72 let path = route.match_.path.as_deref().unwrap_or("/");
73 let host = ServerName::new(host);
74 let path = PathName::new(path);
75 self.backend_manager.register_route(host, path, upstream_id);
76 }
77 }
78 }
79
80 info!(
81 "Loaded {} routes and {} upstreams",
82 config.routes.len(),
83 config.upstreams.len()
84 );
85 }
86
87 pub async fn serve_http(
89 self: Arc<Self>,
90 addr: SocketAddr,
91 cancel_token: CancellationToken,
92 ) -> std::io::Result<()> {
93 let listener = TcpListener::bind(addr).await?;
94 info!("HTTP proxy listening on {}", addr);
95
96 loop {
97 tokio::select! {
98 _ = cancel_token.cancelled() => {
99 info!("HTTP proxy stopped");
100 break;
101 }
102 result = listener.accept() => {
103 match result {
104 Ok((stream, client_addr)) => {
105 let service = self.clone();
106 tokio::spawn(async move {
107 let service_fn = service_fn(|req| {
108 let svc = service.clone();
109 async move {
110 svc.handle_request(req, client_addr, "http").await
111 }
112 });
113
114 if let Err(err) = http1::Builder::new()
115 .serve_connection(hyper_util::rt::TokioIo::new(stream), service_fn)
116 .await
117 {
118 debug!("Connection error: {}", err);
119 }
120 });
121 }
122 Err(e) => {
123 error!("Accept error: {}", e);
124 }
125 }
126 }
127 }
128 }
129
130 Ok(())
131 }
132
133 async fn handle_request(
135 &self,
136 request: Request<Incoming>,
137 client_addr: SocketAddr,
138 proto: &str,
139 ) -> Result<Response<ProxyBody>, Infallible> {
140 let uri = request.uri().clone();
141 let method = request.method().clone();
142
143 let host = request
145 .headers()
146 .get(hyper::header::HOST)
147 .and_then(|h| h.to_str().ok())
148 .map(|h| h.split(':').next().unwrap_or(h))
149 .unwrap_or("");
150
151 let path = uri.path();
152
153 debug!("{} {} {} from {}", method, host, path, client_addr);
154
155 let (route, upstream_id) = {
157 let config = self.config.read().await;
158 let matched = config
159 .routes
160 .iter()
161 .filter(|r| r.enabled)
162 .filter(|r| {
163 let route_host = r.match_.host.as_deref().unwrap_or("*");
164 route_host == "*" || route_host == host
165 })
166 .filter(|r| {
167 let route_path = r.match_.path.as_deref().unwrap_or("/");
168 path.starts_with(route_path)
169 })
170 .max_by_key(|r| {
171 let path_len = r.match_.path.as_ref().map(|p| p.len()).unwrap_or(0);
172 (r.priority, path_len)
173 })
174 .cloned();
175
176 match matched {
177 Some(r) => {
178 let upstream_id = config
180 .upstreams
181 .iter()
182 .find(|u| u.name == r.upstream)
183 .and_then(|u| u.id);
184 (Some(r), upstream_id)
185 }
186 None => (None, None),
187 }
188 };
189
190 let route = match route {
191 Some(r) => r,
192 None => {
193 debug!("No route found for {} {}", host, path);
194 return Ok(MessageHandler::not_found());
195 }
196 };
197
198 let upstream_id = match upstream_id {
199 Some(id) => id,
200 None => {
201 warn!(
202 "Upstream '{}' not found for route {}",
203 route.upstream, route.name
204 );
205 return Ok(MessageHandler::service_unavailable("Upstream not found"));
206 }
207 };
208
209 if let Some(ref rate_limit) = route.rate_limit {
211 let key = RateLimitKey::Ip(client_addr.ip());
212 let rps = rate_limit
214 .requests
215 .checked_div(rate_limit.window_secs)
216 .unwrap_or(rate_limit.requests);
217 if !self
218 .rate_limiter
219 .check_with_config(key, rps, rate_limit.requests)
220 {
221 return Ok(MessageHandler::too_many_requests());
222 }
223 }
224
225 let backend = match self.backend_manager.select_backend(upstream_id, None) {
227 Some(b) => b,
228 None => {
229 warn!("No healthy backends for route {}", route.name);
230 return Ok(MessageHandler::service_unavailable(
231 "No healthy backends available",
232 ));
233 }
234 };
235
236 let (parts, body) = request.into_parts();
238 let mut forwarded_request = Request::from_parts(parts, body.boxed_body());
239
240 MessageHandler::add_forwarding_headers(&mut forwarded_request, client_addr, proto);
242
243 if route.strip_path {
245 if let Some(ref prefix) = route.match_.path {
246 MessageHandler::strip_path_prefix(&mut forwarded_request, prefix);
247 }
248 }
249
250 MessageHandler::apply_route_headers(&mut forwarded_request, &route);
252
253 MessageHandler::rewrite_host_header(&mut forwarded_request, &backend.address);
255
256 match self.forwarder.forward(&backend, forwarded_request).await {
258 Ok(response) => {
259 let (parts, body) = response.into_parts();
260 Ok(Response::from_parts(parts, body.boxed_body()))
261 }
262 Err(e) => {
263 error!("Forward error to {}: {}", backend.address, e);
264 Ok(MessageHandler::bad_gateway(&format!(
265 "Backend error: {}",
266 e
267 )))
268 }
269 }
270 }
271}