From 3dbb6f2488d10277342dc7d47ddcdb29de9e2f48 Mon Sep 17 00:00:00 2001 From: Zhang Jingqiang Date: Mon, 28 Sep 2026 12:15:08 +0800 Subject: [PATCH 1/3] vey-proxy: initial code fix for http-expose --- vey-proxy/src/serve/http_expose/host.rs | 6 +- vey-proxy/src/serve/http_expose/mod.rs | 6 +- vey-proxy/src/serve/http_expose/server.rs | 239 +++++++++--------- .../serve/http_expose/task/pipeline/writer.rs | 10 +- 4 files changed, 125 insertions(+), 136 deletions(-) diff --git a/vey-proxy/src/serve/http_expose/host.rs b/vey-proxy/src/serve/http_expose/host.rs index 81fa2b683..2c38b5c27 100644 --- a/vey-proxy/src/serve/http_expose/host.rs +++ b/vey-proxy/src/serve/http_expose/host.rs @@ -12,13 +12,13 @@ use vey_types::net::{AlpnProtocol, OpensslServerConfig, OpensslTicketKey, Rollin use crate::site::{Site, SiteEgress}; -pub(crate) struct HttpHost { +pub(crate) struct HttpExposeHost { site: Arc, egress: Arc, tls_server: Option, } -impl HttpHost { +impl HttpExposeHost { pub(super) fn try_build( site: Arc, ticketer: Option>>, @@ -36,7 +36,7 @@ impl HttpHost { }; let egress = Arc::new(SiteEgress::new(site.config())); - Ok(HttpHost { + Ok(HttpExposeHost { site, egress, tls_server, diff --git a/vey-proxy/src/serve/http_expose/mod.rs b/vey-proxy/src/serve/http_expose/mod.rs index c6847fab8..bac452e3c 100644 --- a/vey-proxy/src/serve/http_expose/mod.rs +++ b/vey-proxy/src/serve/http_expose/mod.rs @@ -7,10 +7,10 @@ mod stats; use stats::{HttpExposeServerStats, HttpForwardTaskAliveGuard, HttpUntrustedTaskAliveGuard}; +mod host; +use host::HttpExposeHost; + mod task; mod server; pub(super) use server::HttpExposeServer; - -mod host; -pub(crate) use host::HttpHost; diff --git a/vey-proxy/src/serve/http_expose/server.rs b/vey-proxy/src/serve/http_expose/server.rs index 7d0093596..be9ba5682 100644 --- a/vey-proxy/src/serve/http_expose/server.rs +++ b/vey-proxy/src/serve/http_expose/server.rs @@ -43,7 +43,7 @@ use super::task::{ CommonTaskContext, HttpExposePipelineReaderTask, HttpExposePipelineStats, HttpExposePipelineWriterTask, }; -use super::{HttpExposeServerStats, HttpHost}; +use super::{HttpExposeHost, HttpExposeServerStats}; use crate::auth::UserGroup; use crate::config::server::http_expose::HttpExposeServerConfig; use crate::config::server::{AnyServerConfig, ServerConfig}; @@ -62,7 +62,7 @@ pub(crate) struct HttpExposeServer { ingress_net_filter: Option, reload_sender: broadcast::Sender>, task_logger: Option, - hosts: ArcSwap>>, + hosts: ArcSwap>>, escaper: ArcSwap, user_group: ArcSwapOption, @@ -76,10 +76,14 @@ impl HttpExposeServer { config: Arc, server_stats: Arc, listen_stats: Arc, - hosts: HostMatch>, tls_rolling_ticketer: Option>>, version: usize, ) -> anyhow::Result { + let group = crate::site::get_or_insert_default(&config.site_group); + let hosts = group.sites_by_host().try_build_arc(|site| { + HttpExposeHost::try_build(Arc::clone(site), tls_rolling_ticketer.clone()) + })?; + let reload_sender = ServerReloadCommand::new_sender(); let global_tls_server = match &config.global_tls_server { @@ -144,53 +148,36 @@ impl HttpExposeServer { } else { None }; - let hosts = build_hosts(&config.site_group, tls_rolling_ticketer.clone())?; + + let server = + HttpExposeServer::new(config, server_stats, listen_stats, tls_rolling_ticketer, 1)?; + Ok(Arc::new(server)) + } + + fn prepare_reload(&self, config: HttpExposeServerConfig) -> anyhow::Result { + let config = Arc::new(config); + let server_stats = Arc::clone(&self.server_stats); + let listen_stats = Arc::clone(&self.listen_stats); + + let tls_rolling_ticketer = if self.config.tls_ticketer.eq(&config.tls_ticketer) { + self.tls_rolling_ticketer.clone() + } else if let Some(c) = &config.tls_ticketer { + let ticketer = c + .build_and_spawn_updater() + .context("failed to create tls rolling ticketer")?; + Some(ticketer) + } else { + None + }; let server = HttpExposeServer::new( config, server_stats, listen_stats, - hosts, tls_rolling_ticketer, - 1, + self.reload_version + 1, )?; - Ok(Arc::new(server)) - } - - fn prepare_reload(&self, config: AnyServerConfig) -> anyhow::Result { - if let AnyServerConfig::HttpExpose(config) = config { - let config = Arc::new(config); - let server_stats = Arc::clone(&self.server_stats); - let listen_stats = Arc::clone(&self.listen_stats); - - let tls_rolling_ticketer = if self.config.tls_ticketer.eq(&config.tls_ticketer) { - self.tls_rolling_ticketer.clone() - } else if let Some(c) = &config.tls_ticketer { - let ticketer = c - .build_and_spawn_updater() - .context("failed to create tls rolling ticketer")?; - Some(ticketer) - } else { - None - }; - let hosts = build_hosts(&config.site_group, tls_rolling_ticketer.clone())?; - - let server = HttpExposeServer::new( - config, - server_stats, - listen_stats, - hosts, - tls_rolling_ticketer, - self.reload_version + 1, - )?; - Ok(server) - } else { - Err(anyhow!( - "config type mismatch: expect {}, actual {}", - self.config.r#type(), - config.r#type() - )) - } + Ok(server) } fn get_common_task_context( @@ -255,16 +242,10 @@ impl HttpExposeServer { async fn run_tls_tcp_task(&self, mut stream: TcpStream, cc_info: ClientConnectionInfo) { const TLS_MAX_CLIENT_HELLO_SIZE: u32 = 1 << 16; - let hosts = self.hosts.load(); let mut clt_r_buf = BytesMut::with_capacity(2048); let host = match tokio::time::timeout( self.config.client_hello_recv_timeout, - read_sni_host( - &mut stream, - &mut clt_r_buf, - TLS_MAX_CLIENT_HELLO_SIZE, - &hosts, - ), + self.read_sni_host(&mut stream, &mut clt_r_buf, TLS_MAX_CLIENT_HELLO_SIZE), ) .await { @@ -290,6 +271,7 @@ impl HttpExposeServer { }; let Some(tls_config) = host + .as_ref() .and_then(|h| h.tls_server()) .or(self.global_tls_server.as_ref()) else { @@ -328,82 +310,73 @@ impl HttpExposeServer { } } } -} -async fn read_sni_host<'a>( - clt_r: &mut TcpStream, - clt_r_buf: &mut BytesMut, - max_client_hello_size: u32, - hosts: &'a HostMatch>, -) -> anyhow::Result>> { - let max_hello_size = max_client_hello_size as usize; - let max_buf_size = max_hello_size - .saturating_mul(RecordHeader::SIZE + 1) - .saturating_add(1 << 14); - let mut handshake_coalescer = HandshakeCoalescer::new(max_client_hello_size); - let mut record_offset = 0; - loop { - let mut record = match Record::parse(&clt_r_buf[record_offset..]) { - Ok(r) => r, - Err(RecordParseError::NeedMoreData(_)) => { - if clt_r_buf.len() >= max_buf_size { - return Err(anyhow!("tls client hello message too large")); + async fn read_sni_host( + &self, + clt_r: &mut TcpStream, + clt_r_buf: &mut BytesMut, + max_client_hello_size: u32, + ) -> anyhow::Result>> { + let max_hello_size = max_client_hello_size as usize; + let max_buf_size = max_hello_size + .saturating_mul(RecordHeader::SIZE + 1) + .saturating_add(1 << 14); + let mut handshake_coalescer = HandshakeCoalescer::new(max_client_hello_size); + let mut record_offset = 0; + loop { + let mut record = match Record::parse(&clt_r_buf[record_offset..]) { + Ok(r) => r, + Err(RecordParseError::NeedMoreData(_)) => { + if clt_r_buf.len() >= max_buf_size { + return Err(anyhow!("tls client hello message too large")); + } + match clt_r.read_buf(clt_r_buf).await { + Ok(0) => return Err(anyhow!("connection closed by client")), + Ok(_) => continue, + Err(e) => return Err(anyhow!("client read error: {e}")), + } } - match clt_r.read_buf(clt_r_buf).await { - Ok(0) => return Err(anyhow!("connection closed by client")), - Ok(_) => continue, - Err(e) => return Err(anyhow!("client read error: {e}")), + Err(_) => return Err(anyhow!("invalid tls client hello request")), + }; + record_offset += record.encoded_len(); + + match record.consume_handshake(&mut handshake_coalescer) { + Ok(Some(handshake_msg)) => { + let ch = handshake_msg + .parse_client_hello() + .map_err(|_| anyhow!("invalid tls client hello request"))?; + return self.host_from_client_hello(ch); } - } - Err(_) => return Err(anyhow!("invalid tls client hello request")), - }; - record_offset += record.encoded_len(); - - match record.consume_handshake(&mut handshake_coalescer) { - Ok(Some(handshake_msg)) => { - let ch = handshake_msg - .parse_client_hello() - .map_err(|_| anyhow!("invalid tls client hello request"))?; - return Ok(host_from_client_hello(ch, hosts)); - } - Ok(None) => match handshake_coalescer.parse_client_hello() { - Ok(Some(ch)) => return Ok(host_from_client_hello(ch, hosts)), - Ok(None) => { - if !record.consume_done() { - return Err(anyhow!("partial fragmented tls client hello request")); + Ok(None) => match handshake_coalescer.parse_client_hello() { + Ok(Some(ch)) => return self.host_from_client_hello(ch), + Ok(None) => { + if !record.consume_done() { + return Err(anyhow!("partial fragmented tls client hello request")); + } } - } - Err(_) => return Err(anyhow!("invalid fragmented tls client hello request")), - }, - Err(_) => return Err(anyhow!("invalid tls client hello request")), + Err(_) => return Err(anyhow!("invalid fragmented tls client hello request")), + }, + Err(_) => return Err(anyhow!("invalid tls client hello request")), + } } } -} -fn host_from_client_hello<'a>( - ch: ClientHello<'_>, - hosts: &'a HostMatch>, -) -> Option<&'a Arc> { - match ch.get_ext(ExtensionType::ServerName) { - Ok(Some(data)) => match TlsServerName::from_extension_value(data) { - Ok(sni) => hosts.get(&Host::from(sni)), - Err(_) => hosts.get_default(), - }, - Ok(None) => hosts.get_default(), - Err(_) => hosts.get_default(), + fn host_from_client_hello( + &self, + ch: ClientHello<'_>, + ) -> anyhow::Result>> { + let hosts = self.hosts.load(); + match ch.get_ext(ExtensionType::ServerName) { + Ok(Some(data)) => match TlsServerName::from_extension_value(data) { + Ok(sni) => Ok(hosts.get(&Host::from(sni)).cloned()), + Err(e) => Err(anyhow!("invalid server name extension value: {e}")), + }, + Ok(None) => Ok(hosts.get_default().cloned()), + Err(e) => Err(anyhow!("error getting server name tls extension: {e}")), + } } } -fn build_hosts( - site_group: &NodeName, - ticketer: Option>>, -) -> anyhow::Result>> { - let group = crate::site::get_or_insert_default(site_group); - group - .sites_by_host() - .try_build_arc(|site| HttpHost::try_build(Arc::clone(site), ticketer.clone())) -} - impl ServerInternal for HttpExposeServer { fn _clone_config(&self) -> AnyServerConfig { AnyServerConfig::HttpExpose(self.config.as_ref().clone()) @@ -435,10 +408,10 @@ impl ServerInternal for HttpExposeServer { } fn _update_site_group_in_place(&self) -> anyhow::Result<()> { - if self.config.site_group.is_empty() { - return Ok(()); - } - let hosts = build_hosts(&self.config.site_group, self.tls_rolling_ticketer.clone())?; + let group = crate::site::get_or_insert_default(&self.config.site_group); + let hosts = group.sites_by_host().try_build_arc(|site| { + HttpExposeHost::try_build(Arc::clone(site), self.tls_rolling_ticketer.clone()) + })?; self.hosts.store(Arc::new(hosts)); Ok(()) } @@ -452,9 +425,17 @@ impl ServerInternal for HttpExposeServer { config: AnyServerConfig, _registry: &mut ServerRegistry, ) -> anyhow::Result { - let mut server = self.prepare_reload(config)?; - server.reload_sender = self.reload_sender.clone(); - Ok(Arc::new(server)) + if let AnyServerConfig::HttpExpose(config) = config { + let mut server = self.prepare_reload(config)?; + server.reload_sender = self.reload_sender.clone(); + Ok(Arc::new(server)) + } else { + Err(anyhow!( + "config type mismatch: expect {}, actual {}", + self.config.r#type(), + config.r#type() + )) + } } fn _reload_with_new_notifier( @@ -462,8 +443,16 @@ impl ServerInternal for HttpExposeServer { config: AnyServerConfig, _registry: &mut ServerRegistry, ) -> anyhow::Result { - let server = self.prepare_reload(config)?; - Ok(Arc::new(server)) + if let AnyServerConfig::HttpExpose(config) = config { + let server = self.prepare_reload(config)?; + Ok(Arc::new(server)) + } else { + Err(anyhow!( + "config type mismatch: expect {}, actual {}", + self.config.r#type(), + config.r#type() + )) + } } fn _start_runtime(&self, server: ArcServer) -> anyhow::Result<()> { diff --git a/vey-proxy/src/serve/http_expose/task/pipeline/writer.rs b/vey-proxy/src/serve/http_expose/task/pipeline/writer.rs index aed978eff..0d74bd1de 100644 --- a/vey-proxy/src/serve/http_expose/task/pipeline/writer.rs +++ b/vey-proxy/src/serve/http_expose/task/pipeline/writer.rs @@ -25,7 +25,7 @@ use super::{ use crate::auth::{UserContext, UserGroup, UserRequestStats}; use crate::config::server::ServerConfig; use crate::module::http_forward::{BoxHttpForwardContext, HttpProxyClientResponse}; -use crate::serve::http_expose::HttpHost; +use crate::serve::http_expose::HttpExposeHost; use crate::serve::{ServerStats, ServerTaskNotes}; use crate::site::{Site, SiteContext, SiteHttpConnGuard}; @@ -206,7 +206,7 @@ where } } - pub(crate) async fn into_running(mut self, hosts: Arc>>) { + pub(crate) async fn into_running(mut self, hosts: Arc>>) { loop { let res = match self.task_queue.recv().await { Some(Ok((req, pipeline_task))) => { @@ -241,7 +241,7 @@ where async fn check_run( &mut self, req: HttpExposeRequest, - hosts: &HostMatch>, + hosts: &HostMatch>, ) -> LoopAction { let Some(host) = hosts.get(req.upstream.host()) else { self.req_count.invalid += 1; @@ -338,12 +338,12 @@ where .with_site_ctx(site_ctx.clone()); // check in final escaper so we can use route escapers - let upstream = task_notes.site_upstream_addr().clone(); + let upstream = task_notes.site_upstream_addr(); let _ = self .forward_context .check_in_final_escaper( &task_notes, - &upstream, + upstream, site_ctx.site().tls_client().is_some(), ) .await; From 19d0e0990e1239963a8d8a6fb23bf68f2d040585 Mon Sep 17 00:00:00 2001 From: Zhang Jingqiang Date: Mon, 28 Sep 2026 12:41:58 +0800 Subject: [PATCH 2/3] vey-proxy: select site upstream in the forward task Choose the upstream and check in the final escaper before reusing a connection, and drop the cached origin address on task notes. Co-authored-by: Cursor --- .../serve/http_expose/task/forward/task.rs | 41 ++++++++-- .../serve/http_expose/task/pipeline/writer.rs | 11 --- .../serve/http_guard/task/h1/forward/task.rs | 33 ++++++-- .../http_guard/task/h1/pipeline/writer.rs | 6 -- .../http_guard/task/h1/websocket/task.rs | 9 ++- .../src/serve/http_guard/task/h2/context.rs | 77 ++++++++++++------- .../serve/http_guard/task/h2/forward/h1.rs | 27 ++++--- .../serve/http_guard/task/h2/forward/task.rs | 18 ++++- .../http_guard/task/h2/websocket/task.rs | 11 ++- vey-proxy/src/serve/task.rs | 48 +----------- 10 files changed, 156 insertions(+), 125 deletions(-) diff --git a/vey-proxy/src/serve/http_expose/task/forward/task.rs b/vey-proxy/src/serve/http_expose/task/forward/task.rs index 071eec2b6..4037a1de5 100644 --- a/vey-proxy/src/serve/http_expose/task/forward/task.rs +++ b/vey-proxy/src/serve/http_expose/task/forward/task.rs @@ -84,7 +84,6 @@ impl<'a> HttpExposeForwardTask<'a> { uri_log_max_chars, ); let max_idle_count = task_notes.task_max_idle_count(ctx.server_config.task_idle_max_count); - let upstream = task_notes.site_upstream_addr().clone(); HttpExposeForwardTask { ctx: Arc::clone(ctx), site_ctx, @@ -100,7 +99,7 @@ impl<'a> HttpExposeForwardTask<'a> { max_idle_count, _alive_guard: None, alive_reuse_notes: None, - upstream, + upstream: UpstreamAddr::empty(), } } @@ -245,6 +244,36 @@ impl<'a> HttpExposeForwardTask<'a> { } } + async fn prepare_upstream( + &mut self, + fwd_ctx: &mut BoxHttpForwardContext, + clt_w: &mut HttpClientWriter, + ) -> ServerTaskResult<()> + where + CDW: AsyncWrite + Unpin, + { + let upstream = match self.site().select_upstream(self.ctx.client_ip()) { + Ok(upstream) => upstream, + Err(_) => { + let e = TcpConnectError::InternalServerError("failed to select site upstream"); + self.reply_connect_err(&e, clt_w).await; + return Err(e.into()); + } + }; + self.upstream = upstream; + + if let Some(user_ctx) = self.task_notes.user_ctx() { + let action = user_ctx.check_upstream(&self.upstream); + self.handle_user_upstream_acl_action(action, clt_w).await?; + } + + // check in final escaper so we can use route escapers + let _ = fwd_ctx + .check_in_final_escaper(&self.task_notes, &self.upstream, self.origin_tls()) + .await; + Ok(()) + } + fn pre_start(&mut self) { self._alive_guard = Some(self.ctx.server_stats.add_forward_task()); @@ -434,9 +463,6 @@ impl<'a> HttpExposeForwardTask<'a> { )); } - let action = user_ctx.check_upstream(&self.upstream); - self.handle_user_upstream_acl_action(action, clt_w).await?; - if let Some(action) = user_ctx.check_http_user_agent( self.req .end_to_end_headers @@ -464,6 +490,8 @@ impl<'a> HttpExposeForwardTask<'a> { self.setup_clt_limit_and_stats(clt_r, clt_w); + self.prepare_upstream(fwd_ctx, clt_w).await?; + let keepalive = self.site().h1_keepalive_config(); if keepalive.is_enabled() && let Some(mut connection) = self @@ -659,9 +687,6 @@ impl<'a> HttpExposeForwardTask<'a> { &self, fwd_ctx: &mut BoxHttpForwardContext, ) -> Result<(BoxHttpForwardConnection, HttpAliveReuseNotes), TcpConnectError> { - self.task_notes - .site_upstream() - .map_err(|_| TcpConnectError::InternalServerError("failed to select site upstream"))?; let mut audit_ctx = AuditContext::default(); if let Some(tls_client) = self.site().tls_client() { let task_conf = TlsConnectTaskConf { diff --git a/vey-proxy/src/serve/http_expose/task/pipeline/writer.rs b/vey-proxy/src/serve/http_expose/task/pipeline/writer.rs index 0d74bd1de..6b11eb96d 100644 --- a/vey-proxy/src/serve/http_expose/task/pipeline/writer.rs +++ b/vey-proxy/src/serve/http_expose/task/pipeline/writer.rs @@ -337,17 +337,6 @@ where ) .with_site_ctx(site_ctx.clone()); - // check in final escaper so we can use route escapers - let upstream = task_notes.site_upstream_addr(); - let _ = self - .forward_context - .check_in_final_escaper( - &task_notes, - upstream, - site_ctx.site().tls_client().is_some(), - ) - .await; - match self .run_forward(&mut stream_w, req, site_ctx, task_notes) .await diff --git a/vey-proxy/src/serve/http_guard/task/h1/forward/task.rs b/vey-proxy/src/serve/http_guard/task/h1/forward/task.rs index 853b62d21..f3cccc1f1 100644 --- a/vey-proxy/src/serve/http_guard/task/h1/forward/task.rs +++ b/vey-proxy/src/serve/http_guard/task/h1/forward/task.rs @@ -87,7 +87,6 @@ impl<'a> HttpGuardForwardTask<'a> { uri_log_max_chars, ); let max_idle_count = task_notes.task_max_idle_count(ctx.server_config.task_idle_max_count); - let upstream = task_notes.site_upstream_addr().clone(); HttpGuardForwardTask { ctx: Arc::clone(ctx), site_ctx, @@ -105,7 +104,7 @@ impl<'a> HttpGuardForwardTask<'a> { _alive_guard: None, alive_reuse_notes: None, origin_session_auth, - upstream, + upstream: UpstreamAddr::empty(), } } @@ -433,6 +432,8 @@ impl<'a> HttpGuardForwardTask<'a> { self.setup_clt_limit_and_stats(clt_r, clt_w); + self.prepare_upstream(fwd_ctx, clt_w).await?; + let keepalive = self.site().h1_keepalive_config(); if keepalive.is_enabled() && let Some(mut connection) = self @@ -641,13 +642,35 @@ impl<'a> HttpGuardForwardTask<'a> { } } + async fn prepare_upstream( + &mut self, + fwd_ctx: &mut BoxHttpForwardContext, + clt_w: &mut HttpClientWriter, + ) -> ServerTaskResult<()> + where + CDW: AsyncWrite + Unpin, + { + let upstream = match self.site().select_upstream(self.ctx.client_ip()) { + Ok(upstream) => upstream, + Err(_) => { + let e = TcpConnectError::InternalServerError("failed to select site upstream"); + self.reply_connect_err(&e, clt_w).await; + return Err(e.into()); + } + }; + self.upstream = upstream; + + // check in final escaper so we can use route escapers + let _ = fwd_ctx + .check_in_final_escaper(&self.task_notes, &self.upstream, self.origin_tls()) + .await; + Ok(()) + } + async fn make_new_connection( &self, fwd_ctx: &mut BoxHttpForwardContext, ) -> Result<(BoxHttpForwardConnection, HttpAliveReuseNotes), TcpConnectError> { - self.task_notes - .site_upstream() - .map_err(|_| TcpConnectError::InternalServerError("failed to select site upstream"))?; let mut audit_ctx = AuditContext::new(self.ctx.audit_handle.clone()); if let Some(tls_client) = self.site().tls_client() { let task_conf = TlsConnectTaskConf { diff --git a/vey-proxy/src/serve/http_guard/task/h1/pipeline/writer.rs b/vey-proxy/src/serve/http_guard/task/h1/pipeline/writer.rs index 606748694..5f0090ae8 100644 --- a/vey-proxy/src/serve/http_guard/task/h1/pipeline/writer.rs +++ b/vey-proxy/src/serve/http_guard/task/h1/pipeline/writer.rs @@ -242,12 +242,6 @@ where LoopAction::Break } None => { - let site = site_ctx.site(); - let upstream = task_notes.site_upstream_addr().clone(); - let _ = self - .forward_context - .check_in_final_escaper(&task_notes, &upstream, site.tls_client().is_some()) - .await; match self .run_forward(&mut stream_w, req, site_ctx, task_notes) .await diff --git a/vey-proxy/src/serve/http_guard/task/h1/websocket/task.rs b/vey-proxy/src/serve/http_guard/task/h1/websocket/task.rs index a6c7e8b01..171e49c88 100644 --- a/vey-proxy/src/serve/http_guard/task/h1/websocket/task.rs +++ b/vey-proxy/src/serve/http_guard/task/h1/websocket/task.rs @@ -79,7 +79,6 @@ impl HttpGuardWebsocketTask { let ws_notes = WebSocketTaskNotes::new(req.inner.version, req.inner.uri.clone(), uri_log_max_chars); let max_idle_count = task_notes.task_max_idle_count(ctx.server_config.task_idle_max_count); - let upstream = task_notes.site_upstream_addr().clone(); HttpGuardWebsocketTask { ctx: Arc::clone(ctx), site_ctx, @@ -91,7 +90,7 @@ impl HttpGuardWebsocketTask { ups_r_leftover: None, send_error_response: true, _alive_guard: None, - upstream, + upstream: UpstreamAddr::empty(), } } @@ -409,8 +408,10 @@ impl HttpGuardWebsocketTask { &mut self, req: &HttpProxyClientRequest, ) -> Result { - self.task_notes - .site_upstream() + self.upstream = self + .site_ctx + .site() + .select_upstream(self.ctx.client_ip()) .map_err(|_| TcpConnectError::InternalServerError("failed to select site upstream"))?; let mut audit_ctx = AuditContext::new(self.ctx.audit_handle.clone()); let task_stats: ArcTcpConnectionTaskRemoteStats = self.task_stats.clone(); diff --git a/vey-proxy/src/serve/http_guard/task/h2/context.rs b/vey-proxy/src/serve/http_guard/task/h2/context.rs index ec03f727d..09e26b275 100644 --- a/vey-proxy/src/serve/http_guard/task/h2/context.rs +++ b/vey-proxy/src/serve/http_guard/task/h2/context.rs @@ -19,7 +19,7 @@ use uuid::Uuid; use vey_daemon::stat::remote::ArcTcpConnectionTaskRemoteStats; use vey_daemon::stat::task::TcpStreamTaskStats; use vey_h2::RequestExt; -use vey_types::net::{AlpnProtocol, ForwardedValue, Host, HttpForwardedHeaderType}; +use vey_types::net::{AlpnProtocol, ForwardedValue, Host, HttpForwardedHeaderType, UpstreamAddr}; use super::{CommonTaskContext, H2StreamTransferError}; use crate::audit::AuditContext; @@ -158,33 +158,40 @@ impl H2TaskContext { pub(super) async fn checkout_or_connect( &self, task_notes: &mut ServerTaskNotes, + upstream: &UpstreamAddr, request_host: &Host, ) -> Result { - if let Some(origin) = self.checkout_h2(task_notes).await { + if let Some(origin) = self.checkout_h2(task_notes, upstream).await { return Ok(OriginConnection::H2(origin)); } - if let Some(origin) = self.checkout_h1(task_notes).await { + if let Some(origin) = self.checkout_h1(task_notes, upstream).await { return Ok(OriginConnection::H1(origin)); } task_notes.stage = ServerTaskStage::Connecting; - self.connect_origin(task_notes, request_host).await + self.connect_origin(task_notes, upstream, request_host) + .await } pub(super) async fn checkout_or_connect_h2( &self, task_notes: &mut ServerTaskNotes, + upstream: &UpstreamAddr, request_host: &Host, ) -> Result { - if let Some(origin) = self.checkout_h2(task_notes).await { + if let Some(origin) = self.checkout_h2(task_notes, upstream).await { return Ok(origin); } task_notes.stage = ServerTaskStage::Connecting; - self.connect_origin_h2(task_notes, request_host).await + self.connect_origin_h2(task_notes, upstream, request_host) + .await } - async fn checkout_h2(&self, task_notes: &ServerTaskNotes) -> Option { + async fn checkout_h2( + &self, + task_notes: &ServerTaskNotes, + upstream: &UpstreamAddr, + ) -> Option { let open_timeout = self.server_config.h2.upstream_stream_open_timeout; - let peer = task_notes.site_upstream_peer(); let (sender, egress_notes) = self .site_ctx .site() @@ -192,7 +199,7 @@ impl H2TaskContext { .checkout( task_notes.worker_id(), self.escaper.name(), - peer, + upstream.socket_addr(), open_timeout, ) .await?; @@ -203,16 +210,23 @@ impl H2TaskContext { }) } - async fn checkout_h1(&self, task_notes: &ServerTaskNotes) -> Option { + async fn checkout_h1( + &self, + task_notes: &ServerTaskNotes, + upstream: &UpstreamAddr, + ) -> Option { let site = self.site_ctx.site(); let keepalive = site.h1_keepalive_config(); if !keepalive.is_enabled() { return None; } let pool = site.http1_pool()?; - let peer = task_notes.site_upstream_peer(); let (connection, reuse_notes, egress_notes) = pool - .get(task_notes.worker_id(), self.escaper.name(), peer) + .get( + task_notes.worker_id(), + self.escaper.name(), + upstream.socket_addr(), + ) .await?; let task_stats: ArcHttpForwardTaskRemoteStats = Arc::new(NilHttpForwardTaskRemoteStats); let connection = reuse_notes.escaper.prepare_reused_http_forward_connection( @@ -232,12 +246,10 @@ impl H2TaskContext { async fn connect_origin( &self, task_notes: &ServerTaskNotes, + upstream: &UpstreamAddr, request_host: &Host, ) -> Result { let site = self.site_ctx.site(); - let upstream = task_notes - .site_upstream() - .map_err(|e| H2StreamTransferError::OriginConnectFailed(anyhow!("{e}")))?; let mut egress_notes = EgressNotes::default(); let mut audit_ctx = AuditContext::new(self.audit_handle.clone()); let task_stats: ArcTcpConnectionTaskRemoteStats = Arc::new(TcpStreamTaskStats::default()); @@ -268,12 +280,18 @@ impl H2TaskContext { } wrap_escaper.tls_connection_with_task_stats(stream, task_notes, task_stats) } else { - self.setup_origin_tcp(task_notes, &mut egress_notes, &mut audit_ctx, task_stats) - .await? + self.setup_origin_tcp( + task_notes, + upstream, + &mut egress_notes, + &mut audit_ctx, + task_stats, + ) + .await? }; Ok(OriginConnection::H2( - self.finish_h2_origin(stream, egress_notes, task_notes) + self.finish_h2_origin(stream, egress_notes, task_notes, upstream) .await?, )) } @@ -281,12 +299,10 @@ impl H2TaskContext { async fn connect_origin_h2( &self, task_notes: &ServerTaskNotes, + upstream: &UpstreamAddr, request_host: &Host, ) -> Result { let site = self.site_ctx.site(); - let upstream = task_notes - .site_upstream() - .map_err(|e| H2StreamTransferError::OriginConnectFailed(anyhow!("{e}")))?; let mut egress_notes = EgressNotes::default(); let mut audit_ctx = AuditContext::new(self.audit_handle.clone()); let task_stats: ArcTcpConnectionTaskRemoteStats = Arc::new(TcpStreamTaskStats::default()); @@ -309,24 +325,28 @@ impl H2TaskContext { .await .map_err(|e| H2StreamTransferError::OriginConnectFailed(anyhow!("{e}")))? } else { - self.setup_origin_tcp(task_notes, &mut egress_notes, &mut audit_ctx, task_stats) - .await? + self.setup_origin_tcp( + task_notes, + upstream, + &mut egress_notes, + &mut audit_ctx, + task_stats, + ) + .await? }; - self.finish_h2_origin(stream, egress_notes, task_notes) + self.finish_h2_origin(stream, egress_notes, task_notes, upstream) .await } async fn setup_origin_tcp( &self, task_notes: &ServerTaskNotes, + upstream: &UpstreamAddr, egress_notes: &mut EgressNotes, audit_ctx: &mut AuditContext, task_stats: ArcTcpConnectionTaskRemoteStats, ) -> Result { - let upstream = task_notes - .site_upstream() - .map_err(H2StreamTransferError::OriginConnectFailed)?; let task_conf = TcpConnectTaskConf { upstream }; self.escaper .tcp_setup_connection(&task_conf, egress_notes, task_notes, task_stats, audit_ctx) @@ -339,12 +359,13 @@ impl H2TaskContext { stream: TcpConnection, egress_notes: EgressNotes, task_notes: &ServerTaskNotes, + upstream: &UpstreamAddr, ) -> Result { let (sender, conn_state) = self.handshake_h2(stream).await?; self.site_ctx.site().http2_pool().insert( task_notes.worker_id(), self.escaper.name().clone(), - task_notes.site_upstream_peer(), + upstream.socket_addr(), sender.clone(), conn_state, egress_notes.clone(), diff --git a/vey-proxy/src/serve/http_guard/task/h2/forward/h1.rs b/vey-proxy/src/serve/http_guard/task/h2/forward/h1.rs index bcdbf6485..188c65efb 100644 --- a/vey-proxy/src/serve/http_guard/task/h2/forward/h1.rs +++ b/vey-proxy/src/serve/http_guard/task/h2/forward/h1.rs @@ -147,8 +147,10 @@ impl H2ForwardTask { self.http_notes.retry_new_connection = false; self.egress_notes = origin.egress_notes.clone(); self.task_notes.stage = ServerTaskStage::Connected; - let upstream = self.task_notes.site_upstream_addr().clone(); - origin.connection.0.prepare_new(&self.task_notes, &upstream); + origin + .connection + .0 + .prepare_new(&self.task_notes, &self.upstream); } fn poll_idle_h1_origin( @@ -193,13 +195,18 @@ impl H2ForwardTask { let task_stats: ArcHttpForwardTaskRemoteStats = Arc::new(NilHttpForwardTaskRemoteStats); let site = self.ctx.site_ctx.site(); let request_host = self.req.host(); - let upstream = self - .task_notes - .site_upstream() - .map_err(|e| H2StreamTransferError::OriginConnectFailed(anyhow!("{e}")))?; + let _ = fwd_ctx + .check_in_final_escaper( + &self.task_notes, + &self.upstream, + site.tls_client().is_some(), + ) + .await; let (connection, reuse_notes) = if let Some(tls_client) = site.tls_client() { let task_conf = TlsConnectTaskConf { - tcp: TcpConnectTaskConf { upstream }, + tcp: TcpConnectTaskConf { + upstream: &self.upstream, + }, tls_config: tls_client, tls_name: site.tls_name_or(&request_host), alpn_protocols: None, @@ -213,7 +220,9 @@ impl H2ForwardTask { ) .await } else { - let task_conf = TcpConnectTaskConf { upstream }; + let task_conf = TcpConnectTaskConf { + upstream: &self.upstream, + }; fwd_ctx .new_prepared_http_connection( &task_conf, @@ -852,7 +861,7 @@ impl H2ForwardTask { pool.save( self.task_notes.worker_id(), self.ctx.escaper.name().clone(), - self.task_notes.site_upstream_peer(), + self.upstream.socket_addr(), origin.connection, origin.reuse_notes, origin.egress_notes, diff --git a/vey-proxy/src/serve/http_guard/task/h2/forward/task.rs b/vey-proxy/src/serve/http_guard/task/h2/forward/task.rs index ba32323cb..ad7057ac7 100644 --- a/vey-proxy/src/serve/http_guard/task/h2/forward/task.rs +++ b/vey-proxy/src/serve/http_guard/task/h2/forward/task.rs @@ -62,7 +62,6 @@ impl H2ForwardTask { ); let task_notes = ServerTaskNotes::new(ctx.cc_info.clone(), None, Default::default()) .with_site_ctx(ctx.site_ctx_for_request()); - let upstream = task_notes.site_upstream_addr().clone(); let allow_continue = req.expect_100_continue(); H2ForwardTask { ctx, @@ -75,7 +74,7 @@ impl H2ForwardTask { send_error_response: true, allow_continue, audit_task: false, - upstream, + upstream: UpstreamAddr::empty(), _alive_guard: None, } } @@ -200,18 +199,19 @@ impl H2ForwardTask { } self.audit_task = self.should_audit(); + self.prepare_upstream()?; let origin = if self.req.maybe_grpc() { let request_host = self.req.host(); OriginConnection::H2( self.ctx - .checkout_or_connect_h2(&mut self.task_notes, &request_host) + .checkout_or_connect_h2(&mut self.task_notes, &self.upstream, &request_host) .await?, ) } else { let request_host = self.req.host(); self.ctx - .checkout_or_connect(&mut self.task_notes, &request_host) + .checkout_or_connect(&mut self.task_notes, &self.upstream, &request_host) .await? }; match origin { @@ -228,6 +228,16 @@ impl H2ForwardTask { Ok(()) } + fn prepare_upstream(&mut self) -> Result<(), H2StreamTransferError> { + self.upstream = self + .ctx + .site_ctx + .site() + .select_upstream(self.ctx.client_ip()) + .map_err(H2StreamTransferError::OriginConnectFailed)?; + Ok(()) + } + pub(super) fn mark_relaying(&mut self) { self.task_notes.mark_relaying(); self.task_notes diff --git a/vey-proxy/src/serve/http_guard/task/h2/websocket/task.rs b/vey-proxy/src/serve/http_guard/task/h2/websocket/task.rs index 1ac788d2f..bc1e3fdc6 100644 --- a/vey-proxy/src/serve/http_guard/task/h2/websocket/task.rs +++ b/vey-proxy/src/serve/http_guard/task/h2/websocket/task.rs @@ -52,7 +52,6 @@ impl H2WebsocketTask { let ws_notes = WebSocketTaskNotes::new(req.version(), req.uri().clone(), uri_log_max_chars); let task_notes = ServerTaskNotes::new(ctx.cc_info.clone(), None, Default::default()) .with_site_ctx(ctx.site_ctx_for_request()); - let upstream = task_notes.site_upstream_addr().clone(); H2WebsocketTask { ctx, clt_stream_id, @@ -65,7 +64,7 @@ impl H2WebsocketTask { ups_rd_bytes: 0, ups_wr_bytes: 0, send_error_response: true, - upstream, + upstream: UpstreamAddr::empty(), _alive_guard: None, } } @@ -138,10 +137,16 @@ impl H2WebsocketTask { return Err(H2StreamTransferError::InternalServerError("fully loaded")); } + self.upstream = self + .ctx + .site_ctx + .site() + .select_upstream(self.ctx.client_ip()) + .map_err(H2StreamTransferError::OriginConnectFailed)?; let request_host = req.host(); let origin = self .ctx - .checkout_or_connect_h2(&mut self.task_notes, &request_host) + .checkout_or_connect_h2(&mut self.task_notes, &self.upstream, &request_host) .await?; self.egress_notes = origin.egress_notes; self.task_notes.stage = ServerTaskStage::Connected; diff --git a/vey-proxy/src/serve/task.rs b/vey-proxy/src/serve/task.rs index 7533832c2..fa2b27f56 100644 --- a/vey-proxy/src/serve/task.rs +++ b/vey-proxy/src/serve/task.rs @@ -5,7 +5,7 @@ */ use std::net::{IpAddr, SocketAddr}; -use std::sync::{Arc, OnceLock}; +use std::sync::Arc; use std::time::Duration; use arc_swap::ArcSwapOption; @@ -17,7 +17,6 @@ use uuid::Uuid; use vey_daemon::server::ClientConnectionInfo; use vey_types::limit::GaugeSemaphorePermit; use vey_types::metrics::{MetricTagMap, NodeName}; -use vey_types::net::UpstreamAddr; use vey_types::resolve::ResolveRedirection; use crate::auth::{ @@ -74,7 +73,6 @@ pub(crate) struct ServerTaskNotes { _user_req_alive_permit: Option, _site_req_alive_permits: SiteRequestPermits, _req_alive_guard: Option, - origin: OnceLock, } impl ServerTaskNotes { @@ -108,7 +106,6 @@ impl ServerTaskNotes { _user_req_alive_permit: None, _site_req_alive_permits: SiteRequestPermits::default(), _req_alive_guard: None, - origin: OnceLock::new(), } } @@ -127,44 +124,6 @@ impl ServerTaskNotes { self.cc_info.client_ip() } - /// Upstream chosen for this task. The first call selects; later calls reuse it. - pub(crate) fn site_upstream(&self) -> anyhow::Result<&UpstreamAddr> { - let cached = self.cached_upstream(); - if let Some(error) = &cached.error { - Err(anyhow::anyhow!("{error}")) - } else { - Ok(&cached.addr) - } - } - - pub(crate) fn site_upstream_addr(&self) -> &UpstreamAddr { - &self.cached_upstream().addr - } - - pub(crate) fn site_upstream_peer(&self) -> Option { - self.site_upstream() - .ok() - .and_then(UpstreamAddr::socket_addr) - } - - fn cached_upstream(&self) -> &CachedUpstream { - let ip = self.client_ip(); - let site = self.site_ctx.as_ref().map(|ctx| Arc::clone(ctx.site())); - self.origin.get_or_init(|| match site { - Some(site) => match site.select_upstream(ip) { - Ok(addr) => CachedUpstream { addr, error: None }, - Err(e) => CachedUpstream { - addr: UpstreamAddr::empty(), - error: Some(e.to_string()), - }, - }, - None => CachedUpstream { - addr: UpstreamAddr::empty(), - error: Some("no site context".to_string()), - }, - }) - } - #[inline] pub(crate) fn server_addr(&self) -> SocketAddr { self.cc_info.server_addr() @@ -430,11 +389,6 @@ impl ServerTaskNotes { } } -struct CachedUpstream { - addr: UpstreamAddr, - error: Option, -} - fn layered_task_idle_count( tenant: Option, origin: Option, From 96cd156cbab4446657a10e0d0eeabe072eccae6c Mon Sep 17 00:00:00 2001 From: Zhang Jingqiang Date: Mon, 28 Sep 2026 13:54:42 +0800 Subject: [PATCH 3/3] vey-proxy: remember the route forward upstream for keepalive Record last_upstream when a route escaper does not reuse the idle connection, so the next check-in can keep that connection. Co-authored-by: Cursor --- vey-proxy/src/module/http_forward/context/route.rs | 1 + 1 file changed, 1 insertion(+) diff --git a/vey-proxy/src/module/http_forward/context/route.rs b/vey-proxy/src/module/http_forward/context/route.rs index 180b50bbf..2e16b519e 100644 --- a/vey-proxy/src/module/http_forward/context/route.rs +++ b/vey-proxy/src/module/http_forward/context/route.rs @@ -75,6 +75,7 @@ impl HttpForwardContext for RouteHttpForwardContext { } } + self.last_upstream.clone_from(upstream); self.escaper._update_egress_path(task_notes); self.run_local_update = true; if let Some(next_escaper) = self