Skip to main content

postrust_proxy/vendored/
proxy.rs

1//! Vendored proxy service from rpxy-lib: proxy/proxy_main.rs
2//!
3//! This module provides the main proxy service that handles incoming connections.
4
5use 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
25/// Proxy service that handles incoming HTTP connections.
26pub struct ProxyService {
27    /// Backend manager
28    backend_manager: Arc<BackendAppManager>,
29    /// Forwarder client
30    forwarder: Arc<ForwarderClient>,
31    /// Rate limiter
32    rate_limiter: Arc<RateLimiter>,
33    /// Configuration
34    config: Arc<RwLock<ProxyConfig>>,
35}
36
37impl ProxyService {
38    /// Create a new proxy service.
39    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    /// Load configuration into the backend manager.
57    pub async fn load_config(&self) {
58        let config = self.config.read().await;
59
60        // Register upstreams
61        for upstream in &config.upstreams {
62            self.backend_manager.register_upstream(upstream.clone());
63        }
64
65        // Register routes
66        for route in &config.routes {
67            // Find upstream by name
68            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    /// Start the HTTP proxy server.
88    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    /// Handle a single HTTP request.
134    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        // Extract host from request
144        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        // Find matching route and upstream
156        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                    // Find upstream ID
179                    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        // Rate limiting
210        if let Some(ref rate_limit) = route.rate_limit {
211            let key = RateLimitKey::Ip(client_addr.ip());
212            // Convert requests/window to rps
213            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        // Select backend
226        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        // Build forwarded request
237        let (parts, body) = request.into_parts();
238        let mut forwarded_request = Request::from_parts(parts, body.boxed_body());
239
240        // Add forwarding headers
241        MessageHandler::add_forwarding_headers(&mut forwarded_request, client_addr, proto);
242
243        // Strip path prefix if configured
244        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        // Apply route headers
251        MessageHandler::apply_route_headers(&mut forwarded_request, &route);
252
253        // Rewrite host header
254        MessageHandler::rewrite_host_header(&mut forwarded_request, &backend.address);
255
256        // Forward request
257        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}