diff --git a/lib/vey-icap-client/src/reqmod/h1/bidirectional.rs b/lib/vey-icap-client/src/reqmod/h1/bidirectional.rs index 453daccd9..d6441646d 100644 --- a/lib/vey-icap-client/src/reqmod/h1/bidirectional.rs +++ b/lib/vey-icap-client/src/reqmod/h1/bidirectional.rs @@ -66,7 +66,10 @@ impl BidirectionalRecvIcapResponse<'_, I> { } r = self.icap_reader.fill_wait_data() => { return match r { - Ok(true) => self.recv_icap_response().await, + Ok(true) => { + state.clt_req_body_size = Some(body_transfer.body_size()); + self.recv_icap_response().await + } Ok(false) => { state.clt_req_body_size = Some(body_transfer.body_size()); Err(H1ReqmodAdaptationError::IcapServerConnectionClosed) diff --git a/lib/vey-icap-client/src/reqmod/h2/bidirectional.rs b/lib/vey-icap-client/src/reqmod/h2/bidirectional.rs index eb4aad2e7..d2fb6bc58 100644 --- a/lib/vey-icap-client/src/reqmod/h2/bidirectional.rs +++ b/lib/vey-icap-client/src/reqmod/h2/bidirectional.rs @@ -67,7 +67,10 @@ impl BidirectionalRecvIcapResponse<'_, I> { } r = self.icap_reader.fill_wait_data() => { return match r { - Ok(true) => self.recv_icap_response().await, + Ok(true) => { + state.clt_req_body_size = Some(body_transfer.received_size()); + self.recv_icap_response().await + } Ok(false) => { state.clt_req_body_size = Some(body_transfer.received_size()); Err(H2ReqmodAdaptationError::IcapServerConnectionClosed) diff --git a/lib/vey-icap-client/src/reqmod/h2_to_h1/bidirectional.rs b/lib/vey-icap-client/src/reqmod/h2_to_h1/bidirectional.rs index 65c1852bc..a4d632a23 100644 --- a/lib/vey-icap-client/src/reqmod/h2_to_h1/bidirectional.rs +++ b/lib/vey-icap-client/src/reqmod/h2_to_h1/bidirectional.rs @@ -52,7 +52,10 @@ impl BidirectionalRecvIcapResponse<'_, I> { } r = self.icap_reader.fill_wait_data() => { return match r { - Ok(true) => self.recv_icap_response().await, + Ok(true) => { + state.record_clt_body_progress(body_transfer); + self.recv_icap_response().await + } Ok(false) => { state.record_clt_body_progress(body_transfer); Err(H2ToH1ReqmodAdaptationError::IcapServerConnectionClosed) diff --git a/lib/vey-icap-client/src/respmod/h1/bidirectional.rs b/lib/vey-icap-client/src/respmod/h1/bidirectional.rs index 09da170fd..2a1e6f2cf 100644 --- a/lib/vey-icap-client/src/respmod/h1/bidirectional.rs +++ b/lib/vey-icap-client/src/respmod/h1/bidirectional.rs @@ -66,7 +66,10 @@ impl BidirectionalRecvIcapResponse<'_, I> { } r = self.icap_reader.fill_wait_data() => { return match r { - Ok(true) => self.recv_icap_response().await, + Ok(true) => { + state.ups_rsp_body_size = Some(body_transfer.body_size()); + self.recv_icap_response().await + } Ok(false) => { state.ups_rsp_body_size = Some(body_transfer.body_size()); Err(H1RespmodAdaptationError::IcapServerConnectionClosed) diff --git a/lib/vey-icap-client/src/respmod/h1_to_h2/bidirectional.rs b/lib/vey-icap-client/src/respmod/h1_to_h2/bidirectional.rs index ae77d5d0b..53e009a50 100644 --- a/lib/vey-icap-client/src/respmod/h1_to_h2/bidirectional.rs +++ b/lib/vey-icap-client/src/respmod/h1_to_h2/bidirectional.rs @@ -65,7 +65,10 @@ impl BidirectionalRecvIcapResponse<'_, I> { } r = self.icap_reader.fill_wait_data() => { return match r { - Ok(true) => self.recv_icap_response().await, + Ok(true) => { + state.ups_rsp_body_size = Some(body_transfer.body_size()); + self.recv_icap_response().await + } Ok(false) => { state.ups_rsp_body_size = Some(body_transfer.body_size()); Err(H1ToH2RespmodAdaptationError::IcapServerConnectionClosed) diff --git a/lib/vey-icap-client/src/respmod/h2/bidirectional.rs b/lib/vey-icap-client/src/respmod/h2/bidirectional.rs index 96a5db557..14aca15a0 100644 --- a/lib/vey-icap-client/src/respmod/h2/bidirectional.rs +++ b/lib/vey-icap-client/src/respmod/h2/bidirectional.rs @@ -68,7 +68,10 @@ impl BidirectionalRecvIcapResponse<'_, I> { } r = self.icap_reader.fill_wait_data() => { return match r { - Ok(true) => self.recv_icap_response().await, + Ok(true) => { + state.ups_rsp_body_size = Some(body_transfer.received_size()); + self.recv_icap_response().await + } Ok(false) => { state.ups_rsp_body_size = Some(body_transfer.received_size()); Err(H2RespmodAdaptationError::IcapServerConnectionClosed) diff --git a/vey-proxy/src/auth/user.rs b/vey-proxy/src/auth/user.rs index d768f244f..a2e8d94d1 100644 --- a/vey-proxy/src/auth/user.rs +++ b/vey-proxy/src/auth/user.rs @@ -801,14 +801,6 @@ impl TenantContext { pub(crate) fn acquire_request_semaphore(&self) -> Result { self.user.acquire_request_semaphore(&self.forbid_stats) } - - pub(crate) fn check_http_user_agent( - &self, - user_agents: impl IntoIterator>, - ) -> Option { - self.user - .check_http_user_agent(user_agents, &self.forbid_stats) - } } #[derive(Clone)] diff --git a/vey-proxy/src/inspect/http/v1/forward/mod.rs b/vey-proxy/src/inspect/http/v1/forward/mod.rs index d4cc10557..4442a521e 100644 --- a/vey-proxy/src/inspect/http/v1/forward/mod.rs +++ b/vey-proxy/src/inspect/http/v1/forward/mod.rs @@ -443,12 +443,15 @@ impl<'a, SC: ServerConfig> H1ForwardTask<'a, SC> { clt_w, &self.ctx.server_config.limited_copy_config(), ); - (&mut copy_to_clt).await.map_err(|e| match e { - StreamCopyError::ReadFailed(e) => ServerTaskError::InternalAdapterError(anyhow!( - "read http error response from adapter failed: {e:?}" - )), - StreamCopyError::WriteFailed(e) => ServerTaskError::ClientTcpWriteFailed(e), - })?; + if let Err(e) = (&mut copy_to_clt).await { + self.http_notes.clt_rsp_body_size = Some(copy_to_clt.reader().body_size()); + return Err(match e { + StreamCopyError::ReadFailed(e) => ServerTaskError::InternalAdapterError( + anyhow!("read http error response from adapter failed: {e:?}"), + ), + StreamCopyError::WriteFailed(e) => ServerTaskError::ClientTcpWriteFailed(e), + }); + } self.http_notes.clt_rsp_body_size = Some(copy_to_clt.reader().body_size()); recv_body.save_connection().await; } else { @@ -621,6 +624,7 @@ impl<'a, SC: ServerConfig> H1ForwardTask<'a, SC> { let copy_done = clt_to_ups.finished(); let rsp_head = match rsp_head { Some(header) => { + record_progress!(); if !clt_body_reader.finished() { // not all client data read in, drop the client connection self.should_close = true; diff --git a/vey-proxy/src/inspect/http/v2/forward/mod.rs b/vey-proxy/src/inspect/http/v2/forward/mod.rs index 91bfb74e7..a68d8e7b5 100644 --- a/vey-proxy/src/inspect/http/v2/forward/mod.rs +++ b/vey-proxy/src/inspect/http/v2/forward/mod.rs @@ -504,6 +504,7 @@ where match r { Ok(rsp) => { if let Some(final_rsp) = self.check_out_final_response(rsp, clt_send_rsp, &mut ups_recv_rsp)? { + record_progress!(); ups_rsp = Some(final_rsp); break; } 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 5d40a92a1..2e09cb92f 100644 --- a/vey-proxy/src/serve/http_expose/task/forward/task.rs +++ b/vey-proxy/src/serve/http_expose/task/forward/task.rs @@ -443,26 +443,19 @@ impl<'a> HttpExposeForwardTask<'a> { CDR: AsyncRead + Unpin, CDW: AsyncWrite + Unpin, { - let tcp_client_misc_opts; - - if self.task_notes.check_layered_rate_limit().is_err() { - self.reply_too_many_requests(clt_w).await; - return Err(ServerTaskError::ForbiddenByRule( - ServerTaskForbiddenError::RateLimited, - )); - } - - if self.task_notes.acquire_site_request_semaphores().is_err() { - self.reply_too_many_requests(clt_w).await; - return Err(ServerTaskError::ForbiddenByRule( - ServerTaskForbiddenError::FullyLoaded, - )); - } - if let Some(user_ctx) = self.task_notes.user_ctx() { + if user_ctx.check_rate_limit().is_err() { + self.reply_too_many_requests(clt_w).await; + return Err(ServerTaskError::ForbiddenByRule( + ServerTaskForbiddenError::RateLimited, + )); + } let user_ctx = user_ctx.clone(); - - if self.task_notes.acquire_user_request_semaphore().is_err() { + if self + .task_notes + .acquire_user_request_semaphore(&user_ctx) + .is_err() + { self.reply_too_many_requests(clt_w).await; return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::FullyLoaded, @@ -478,13 +471,40 @@ impl<'a> HttpExposeForwardTask<'a> { ) { self.handle_user_ua_acl_action(action, clt_w).await?; } + } + if let Some(site_ctx) = self.task_notes.site_ctx() { + if site_ctx.check_rate_limit().is_err() { + self.reply_too_many_requests(clt_w).await; + return Err(ServerTaskError::ForbiddenByRule( + ServerTaskForbiddenError::RateLimited, + )); + } + let site_ctx = site_ctx.clone(); + if self + .task_notes + .acquire_site_request_semaphores(&site_ctx) + .is_err() + { + self.reply_too_many_requests(clt_w).await; + return Err(ServerTaskError::ForbiddenByRule( + ServerTaskForbiddenError::FullyLoaded, + )); + } + } - tcp_client_misc_opts = user_ctx + let tcp_client_misc_opts = if let Some(user_ctx) = self.task_notes.user_ctx() { + user_ctx .user_config() - .tcp_client_misc_opts(&self.ctx.server_config.tcp_misc_opts); + .tcp_client_misc_opts(&self.ctx.server_config.tcp_misc_opts) + } else if let Some(site_ctx) = self.task_notes.site_ctx() + && let Some(tenant) = site_ctx.tenant_ctx() + { + tenant + .user_config() + .tcp_client_misc_opts(&self.ctx.server_config.tcp_misc_opts) } else { - tcp_client_misc_opts = Cow::Borrowed(&self.ctx.server_config.tcp_misc_opts); - } + Cow::Borrowed(&self.ctx.server_config.tcp_misc_opts) + }; // set client side socket options self.ctx @@ -910,7 +930,12 @@ impl<'a> HttpExposeForwardTask<'a> { .map_err(ServerTaskError::UpstreamWriteFailed)?; self.http_notes.mark_req_send_hdr(); self.http_notes.mark_req_send_all(); - self.http_notes.ups_req_body_size = Some(body.len() as u64); + // Chunked bodies are buffered on the wire, while clt_req_body_size is the + // decoded payload. A fully read body was already counted that way. + self.http_notes.ups_req_body_size = self + .http_notes + .clt_req_body_size + .or(Some(body.len() as u64)); match tokio::time::timeout( self.rsp_hdr_recv_timeout(), @@ -1140,6 +1165,7 @@ impl<'a> HttpExposeForwardTask<'a> { let copy_done = clt_to_ups.finished(); let mut rsp_header = match rsp_header { Some(header) => { + record_progress!(); if !clt_body_reader.finished() { // not all client data read in, drop the client connection self.should_close = true; diff --git a/vey-proxy/src/serve/http_expose/task/untrusted/task.rs b/vey-proxy/src/serve/http_expose/task/untrusted/task.rs index 2b347e580..986967ff4 100644 --- a/vey-proxy/src/serve/http_expose/task/untrusted/task.rs +++ b/vey-proxy/src/serve/http_expose/task/untrusted/task.rs @@ -68,12 +68,22 @@ impl<'a> HttpExposeUntrustedTask<'a> { CDR: AsyncRead + Unpin, CDW: AsyncWrite + Unpin, { - if self.task_notes.check_layered_rate_limit().is_err() - || self.task_notes.acquire_site_request_semaphores().is_err() - { - self.should_close = true; - self.reply_too_many_requests(clt_w).await; - return; + if let Some(site_ctx) = self.task_notes.site_ctx() { + if site_ctx.check_rate_limit().is_err() { + self.should_close = true; + self.reply_too_many_requests(clt_w).await; + return; + } + let site_ctx = site_ctx.clone(); + if self + .task_notes + .acquire_site_request_semaphores(&site_ctx) + .is_err() + { + self.should_close = true; + self.reply_too_many_requests(clt_w).await; + return; + } } let site_io = self.site_ctx.fetch_traffic_stats( 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 84fb4accf..ee8d47969 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 @@ -26,7 +26,6 @@ use vey_io_ext::{ GlobalLimitGroup, LimitedBufReadExt, LimitedReadExt, LimitedWriteExt, StreamCopy, StreamCopyError, }; -use vey_types::acl::AclAction; use vey_types::net::{KeepAliveValue, TcpSockSpeedLimitConfig, UpstreamAddr}; use super::protocol::{HttpClientReader, HttpClientWriter, HttpGuardRequest}; @@ -154,18 +153,6 @@ impl<'a> HttpGuardForwardTask<'a> { self.should_close = true; } - async fn reply_forbidden(&mut self, clt_w: &mut W) - where - W: AsyncWrite + Unpin, - { - let mut rsp = HttpProxyClientResponse::forbidden(self.req.version); - self.enable_custom_header_for_local_reply(&mut rsp); - if rsp.reply_err_to_request(clt_w).await.is_ok() { - self.http_notes.rsp_status = rsp.status(); - } - self.should_close = true; - } - async fn reply_connect_err(&mut self, e: &TcpConnectError, clt_w: &mut W) where W: AsyncWrite + Unpin, @@ -273,36 +260,6 @@ impl<'a> HttpGuardForwardTask<'a> { } } - async fn handle_user_ua_acl_action( - &mut self, - action: AclAction, - clt_w: &mut W, - ) -> ServerTaskResult<()> - where - W: AsyncWrite + Unpin, - { - let forbid = match action { - AclAction::Permit => false, - AclAction::PermitAndLog => { - // TODO log permit - false - } - AclAction::Forbid => true, - AclAction::ForbidAndLog => { - // TODO log forbid - true - } - }; - if forbid { - self.reply_forbidden(clt_w).await; - Err(ServerTaskError::ForbiddenByRule( - ServerTaskForbiddenError::UaBlocked, - )) - } else { - Ok(()) - } - } - fn clt_speed_limit(&self) -> Option { let server = self.ctx.server_config.tcp_sock_speed_limit; let limit = self @@ -382,42 +339,42 @@ impl<'a> HttpGuardForwardTask<'a> { CDR: AsyncRead + Send + Unpin, CDW: AsyncWrite + Send + Unpin, { - if self.task_notes.check_layered_rate_limit().is_err() { - self.reply_too_many_requests(clt_w).await; - return Err(ServerTaskError::ForbiddenByRule( - ServerTaskForbiddenError::RateLimited, - )); - } - - if self.task_notes.acquire_site_request_semaphores().is_err() { - self.reply_too_many_requests(clt_w).await; - return Err(ServerTaskError::ForbiddenByRule( - ServerTaskForbiddenError::FullyLoaded, - )); - } - - let tcp_client_misc_opts = if let Some(tenant) = self.task_notes.tenant_ctx().cloned() { - if let Some(action) = tenant.check_http_user_agent( - self.req - .end_to_end_headers - .get_all(header::USER_AGENT) - .iter() - .map(|v| v.to_str()), - ) { - self.handle_user_ua_acl_action(action, clt_w).await?; + let tcp_client_misc_opts = if let Some(site_ctx) = self.task_notes.site_ctx() { + if site_ctx.check_rate_limit().is_err() { + self.reply_too_many_requests(clt_w).await; + return Err(ServerTaskError::ForbiddenByRule( + ServerTaskForbiddenError::RateLimited, + )); } - - if let Some(audit_handle) = self.ctx.audit_handle.as_ref() { - self.audit_task = tenant - .user() - .audit() - .do_task_audit() - .unwrap_or_else(|| audit_handle.do_task_audit()); + let site_ctx = site_ctx.clone(); + if self + .task_notes + .acquire_site_request_semaphores(&site_ctx) + .is_err() + { + self.reply_too_many_requests(clt_w).await; + return Err(ServerTaskError::ForbiddenByRule( + ServerTaskForbiddenError::FullyLoaded, + )); } - tenant - .user_config() - .tcp_client_misc_opts(&self.ctx.server_config.tcp_misc_opts) + if let Some(tenant) = site_ctx.tenant_ctx() { + if let Some(audit_handle) = self.ctx.audit_handle.as_ref() { + self.audit_task = tenant + .user() + .audit() + .do_task_audit() + .unwrap_or_else(|| audit_handle.do_task_audit()); + } + tenant + .user_config() + .tcp_client_misc_opts(&self.ctx.server_config.tcp_misc_opts) + } else { + if let Some(audit_handle) = self.ctx.audit_handle.as_ref() { + self.audit_task = audit_handle.do_task_audit(); + } + Cow::Borrowed(&self.ctx.server_config.tcp_misc_opts) + } } else { if let Some(audit_handle) = self.ctx.audit_handle.as_ref() { self.audit_task = audit_handle.do_task_audit(); @@ -950,12 +907,15 @@ impl<'a> HttpGuardForwardTask<'a> { let mut body_reader = recv_body.body_reader(); let mut copy_to_clt = StreamCopy::new(&mut body_reader, clt_w, &self.ctx.server_config.tcp_copy); - (&mut copy_to_clt).await.map_err(|e| match e { - StreamCopyError::ReadFailed(e) => ServerTaskError::InternalAdapterError(anyhow!( - "read http error response from adapter failed: {e:?}" - )), - StreamCopyError::WriteFailed(e) => ServerTaskError::ClientTcpWriteFailed(e), - })?; + if let Err(e) = (&mut copy_to_clt).await { + self.http_notes.clt_rsp_body_size = Some(copy_to_clt.reader().body_size()); + return Err(match e { + StreamCopyError::ReadFailed(e) => ServerTaskError::InternalAdapterError( + anyhow!("read http error response from adapter failed: {e:?}"), + ), + StreamCopyError::WriteFailed(e) => ServerTaskError::ClientTcpWriteFailed(e), + }); + } self.http_notes.clt_rsp_body_size = Some(copy_to_clt.reader().body_size()); recv_body.save_connection().await; } else { @@ -1135,7 +1095,12 @@ impl<'a> HttpGuardForwardTask<'a> { .map_err(ServerTaskError::UpstreamWriteFailed)?; self.http_notes.mark_req_send_hdr(); self.http_notes.mark_req_send_all(); - self.http_notes.ups_req_body_size = Some(body.len() as u64); + // Chunked bodies are buffered on the wire, while clt_req_body_size is the + // decoded payload. A fully read body was already counted that way. + self.http_notes.ups_req_body_size = self + .http_notes + .clt_req_body_size + .or(Some(body.len() as u64)); match tokio::time::timeout( self.rsp_hdr_recv_timeout(), @@ -1352,6 +1317,7 @@ impl<'a> HttpGuardForwardTask<'a> { let copy_done = clt_to_ups.finished(); let mut rsp_header = match rsp_header { Some(header) => { + record_progress!(); if !clt_body_reader.finished() { // not all client data read in, drop the client connection self.should_close = true; 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 9d3f4ee98..6ce9492df 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 @@ -9,7 +9,6 @@ use std::time::Duration; use anyhow::anyhow; use bytes::Bytes; -use http::header; use tokio::io::{AsyncBufRead, AsyncRead, AsyncWrite, AsyncWriteExt}; use vey_daemon::server::ServerQuitPolicy; @@ -26,7 +25,6 @@ use vey_io_ext::{ FlexBufReader, IdleInterval, LimitedReader, LimitedWriteExt, OnceBufReader, StreamCopy, StreamCopyConfig, StreamCopyError, }; -use vey_types::acl::AclAction; use vey_types::net::UpstreamAddr; use super::H1TaskContext; @@ -233,49 +231,43 @@ impl HttpGuardWebsocketTask { where CDW: AsyncWrite + Send + Unpin, { - if self.task_notes.check_layered_rate_limit().is_err() { - self.reply_too_many_requests(clt_w).await; - return Err(ServerTaskError::ForbiddenByRule( - ServerTaskForbiddenError::RateLimited, - )); - } - - if self.task_notes.acquire_site_request_semaphores().is_err() { - self.reply_too_many_requests(clt_w).await; - return Err(ServerTaskError::ForbiddenByRule( - ServerTaskForbiddenError::FullyLoaded, - )); - } - - let tenant = self.task_notes.tenant_ctx().cloned(); let mut audit_task = false; - let tcp_client_misc_opts = if let Some(tenant) = &tenant { - if let Some(action) = tenant.check_http_user_agent( - req.end_to_end_headers - .get_all(header::USER_AGENT) - .iter() - .map(|v| v.to_str()), - ) { - match action { - AclAction::Permit | AclAction::PermitAndLog => {} - AclAction::Forbid | AclAction::ForbidAndLog => { - self.reply_forbidden(clt_w).await; - return Err(ServerTaskError::ForbiddenByRule( - ServerTaskForbiddenError::UaBlocked, - )); - } - } + let tcp_client_misc_opts = if let Some(site_ctx) = self.task_notes.site_ctx() { + if site_ctx.check_rate_limit().is_err() { + self.reply_too_many_requests(clt_w).await; + return Err(ServerTaskError::ForbiddenByRule( + ServerTaskForbiddenError::RateLimited, + )); } - if let Some(audit_handle) = self.ctx.audit_handle.as_ref() { - audit_task = tenant - .user() - .audit() - .do_task_audit() - .unwrap_or_else(|| audit_handle.do_task_audit()); + let site_ctx = site_ctx.clone(); + if self + .task_notes + .acquire_site_request_semaphores(&site_ctx) + .is_err() + { + self.reply_too_many_requests(clt_w).await; + return Err(ServerTaskError::ForbiddenByRule( + ServerTaskForbiddenError::FullyLoaded, + )); + } + + if let Some(tenant) = site_ctx.tenant_ctx() { + if let Some(audit_handle) = self.ctx.audit_handle.as_ref() { + audit_task = tenant + .user() + .audit() + .do_task_audit() + .unwrap_or_else(|| audit_handle.do_task_audit()); + } + tenant + .user_config() + .tcp_client_misc_opts(&self.ctx.server_config.tcp_misc_opts) + } else { + if let Some(audit_handle) = self.ctx.audit_handle.as_ref() { + audit_task = audit_handle.do_task_audit(); + } + Cow::Borrowed(&self.ctx.server_config.tcp_misc_opts) } - tenant - .user_config() - .tcp_client_misc_opts(&self.ctx.server_config.tcp_misc_opts) } else { if let Some(audit_handle) = self.ctx.audit_handle.as_ref() { audit_task = audit_handle.do_task_audit(); @@ -742,18 +734,6 @@ impl HttpGuardWebsocketTask { } } - async fn reply_forbidden(&mut self, clt_w: &mut W) - where - W: AsyncWrite + Unpin, - { - self.send_error_response = false; - let mut rsp = HttpProxyClientResponse::forbidden(self.ws_notes.version); - self.enable_custom_header_for_local_reply(&mut rsp); - if rsp.reply_err_to_request(clt_w).await.is_ok() { - self.ws_notes.rsp_status = rsp.status(); - } - } - async fn reply_connect_err(&mut self, e: &TcpConnectError, clt_w: &mut W) where W: AsyncWrite + Unpin, 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 1d517fe7b..56f08ac37 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 @@ -12,7 +12,6 @@ use h2::server::SendResponse; use h2::{RecvStream, SendStream}; use http::{HeaderMap, Response}; use tokio::io::AsyncWriteExt; -use tokio::time::Instant; use vey_h2::{ H2BodyEncodeTransfer, H2StreamBodyTransferError, H2StreamFromChunkedTransfer, @@ -254,7 +253,8 @@ impl H2ForwardTask { no_body: bool, ) -> Result { self.http_notes.retry_new_connection = no_body; - let mut adaptation_state = ReqmodAdaptationRunState::new(Instant::now()); + let mut adaptation_state = + ReqmodAdaptationRunState::new(self.task_notes.task_created_instant()); let mut rsp_header = None; let icap_error = { let ups_w = &mut origin.connection.0; @@ -409,6 +409,12 @@ impl H2ForwardTask { ); let mut idle_interval = self.ctx.idle_wheel.register(); let mut idle_count = 0; + macro_rules! record_req_body { + () => { + self.http_notes.clt_req_body_size = Some(body_transfer.received_size()); + self.http_notes.ups_req_body_size = Some(body_transfer.copied_size()); + }; + } loop { tokio::select! { biased; @@ -419,6 +425,7 @@ impl H2ForwardTask { if let Some(final_hdr) = self.check_out_h1_informational(hdr, clt_send_rsp)? { + record_req_body!(); rsp_header = Some(final_hdr); break (body_transfer.recv_finished(), body_transfer.finished()); } @@ -427,12 +434,14 @@ impl H2ForwardTask { if body_transfer.received_size() == 0 { self.http_notes.retry_new_connection = true; } + record_req_body!(); return Err(H2StreamTransferError::OriginClosed); } Err(e) => { if body_transfer.received_size() == 0 { self.http_notes.retry_new_connection = true; } + record_req_body!(); return Err(H2StreamTransferError::OriginReadFailed(e)); } } @@ -445,7 +454,10 @@ impl H2ForwardTask { self.http_notes.ups_req_body_size = Some(n); break (true, true); } - Err(e) => return Err(e.into()), + Err(e) => { + record_req_body!(); + return Err(e.into()); + } } } n = idle_interval.tick() => { @@ -456,6 +468,7 @@ impl H2ForwardTask { self.ctx.server_config.task_idle_max_count, ) { + record_req_body!(); return Err(H2StreamTransferError::Idle( idle_interval.period(), idle_count, @@ -466,6 +479,7 @@ impl H2ForwardTask { body_transfer.reset_active(); } if self.ctx.server_quit_policy.force_quit() { + record_req_body!(); return Err(H2StreamTransferError::CanceledAsServerQuit); } } @@ -632,7 +646,7 @@ impl H2ForwardTask { { Ok(mut adapter) => { let mut adaptation_state = RespmodAdaptationRunState::new( - Instant::now(), + self.task_notes.task_created_instant(), self.http_notes.dur_rsp_recv_hdr, ); adapter.set_client_addr(self.task_notes.client_addr()); @@ -685,6 +699,9 @@ impl H2ForwardTask { clt_send_rsp, ) .await; + if let Some(dur) = adaptation_state.dur_ups_recv_all { + self.http_notes.dur_rsp_recv_all = dur; + } self.http_notes.ups_rsp_body_size = adaptation_state.ups_rsp_body_size; self.http_notes.clt_rsp_body_size = adaptation_state.clt_rsp_body_size; match r { @@ -752,7 +769,7 @@ impl H2ForwardTask { Ok(_) => break, Err(e) => { self.http_notes.ups_rsp_body_size = - Some(body_transfer.copied_size()); + Some(body_transfer.received_size()); self.http_notes.clt_rsp_body_size = Some(body_transfer.copied_size()); return Err(e.into()); @@ -768,7 +785,7 @@ impl H2ForwardTask { ) { self.http_notes.ups_rsp_body_size = - Some(body_transfer.copied_size()); + Some(body_transfer.received_size()); self.http_notes.clt_rsp_body_size = Some(body_transfer.copied_size()); return Err(H2StreamTransferError::Idle( @@ -782,7 +799,7 @@ impl H2ForwardTask { } if self.ctx.server_quit_policy.force_quit() { self.http_notes.ups_rsp_body_size = - Some(body_transfer.copied_size()); + Some(body_transfer.received_size()); self.http_notes.clt_rsp_body_size = Some(body_transfer.copied_size()); return Err(H2StreamTransferError::CanceledAsServerQuit); @@ -817,7 +834,7 @@ impl H2ForwardTask { Ok(_) => break, Err(e) => { self.http_notes.ups_rsp_body_size = - Some(body_transfer.copied_size()); + Some(body_transfer.received_size()); self.http_notes.clt_rsp_body_size = Some(body_transfer.copied_size()); return Err(e.into()); @@ -833,7 +850,7 @@ impl H2ForwardTask { ) { self.http_notes.ups_rsp_body_size = - Some(body_transfer.copied_size()); + Some(body_transfer.received_size()); self.http_notes.clt_rsp_body_size = Some(body_transfer.copied_size()); return Err(H2StreamTransferError::Idle( @@ -847,7 +864,7 @@ impl H2ForwardTask { } if self.ctx.server_quit_policy.force_quit() { self.http_notes.ups_rsp_body_size = - Some(body_transfer.copied_size()); + Some(body_transfer.received_size()); self.http_notes.clt_rsp_body_size = Some(body_transfer.copied_size()); return Err(H2StreamTransferError::CanceledAsServerQuit); 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 b9d1fed2a..30243ad11 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 @@ -9,7 +9,7 @@ use bytes::Bytes; use h2::client::{ResponseFuture, SendRequest}; use h2::server::SendResponse; use h2::{RecvStream, SendStream, StreamId}; -use http::{HeaderMap, Request, Response, StatusCode, Version, header}; +use http::{HeaderMap, Request, Response, StatusCode, Version}; use tokio::time::Instant; use vey_h2::{H2BodyTransfer, H2ResponseHeaderReceiver, RequestExt}; @@ -18,7 +18,6 @@ use vey_icap_client::reqmod::h2::{ ReqmodRecvHttpResponseBody, }; use vey_icap_client::respmod::h2::{RespmodAdaptationEndState, RespmodAdaptationRunState}; -use vey_types::acl::AclAction; use vey_types::net::UpstreamAddr; use super::{H2StreamTransferError, H2TaskContext, OriginConnection, OriginH2Sender}; @@ -157,47 +156,39 @@ impl H2ForwardTask { } } - fn should_audit(&self) -> bool { - self.task_notes - .tenant_user() - .and_then(|u| u.audit().do_task_audit()) - .unwrap_or_else(|| { - self.ctx - .audit_handle - .as_ref() - .is_some_and(|h| h.do_task_audit()) - }) - } - async fn do_forward( &mut self, clt_body: RecvStream, clt_send_rsp: &mut SendResponse, ) -> Result<(), H2StreamTransferError> { - if self.task_notes.check_layered_rate_limit().is_err() { - self.reply_denied(clt_send_rsp, StatusCode::TOO_MANY_REQUESTS); - return Err(H2StreamTransferError::InternalServerError("rate limited")); - } - if self.task_notes.acquire_site_request_semaphores().is_err() { - self.reply_denied(clt_send_rsp, StatusCode::TOO_MANY_REQUESTS); - return Err(H2StreamTransferError::InternalServerError("fully loaded")); - } - - if let Some(tenant) = self.task_notes.tenant_ctx() - && let Some(action) = tenant.check_http_user_agent( - self.req - .headers() - .get_all(header::USER_AGENT) - .iter() - .filter_map(|v| v.to_str().ok()), - ) - && matches!(action, AclAction::Forbid | AclAction::ForbidAndLog) - { - self.reply_denied(clt_send_rsp, StatusCode::FORBIDDEN); - return Err(H2StreamTransferError::InternalServerError("ua denied")); + if let Some(site_ctx) = self.task_notes.site_ctx() { + if site_ctx.check_rate_limit().is_err() { + self.reply_denied(clt_send_rsp, StatusCode::TOO_MANY_REQUESTS); + return Err(H2StreamTransferError::InternalServerError("rate limited")); + } + let site_ctx = site_ctx.clone(); + if self + .task_notes + .acquire_site_request_semaphores(&site_ctx) + .is_err() + { + self.reply_denied(clt_send_rsp, StatusCode::TOO_MANY_REQUESTS); + return Err(H2StreamTransferError::InternalServerError("fully loaded")); + } + if let Some(tenant) = site_ctx.tenant_ctx() { + if let Some(audit_handle) = self.ctx.audit_handle.as_ref() { + self.audit_task = tenant + .user() + .audit() + .do_task_audit() + .unwrap_or_else(|| audit_handle.do_task_audit()); + } + } else if let Some(audit_handle) = self.ctx.audit_handle.as_ref() { + self.audit_task = audit_handle.do_task_audit(); + } + } else if let Some(audit_handle) = self.ctx.audit_handle.as_ref() { + self.audit_task = audit_handle.do_task_audit(); } - - self.audit_task = self.should_audit(); self.prepare_upstream()?; let origin = if self.req.maybe_grpc() { @@ -265,7 +256,8 @@ impl H2ForwardTask { .await { Ok(mut adapter) => { - let mut adaptation_state = ReqmodAdaptationRunState::new(Instant::now()); + let mut adaptation_state = + ReqmodAdaptationRunState::new(self.task_notes.task_created_instant()); adapter.set_client_addr(self.task_notes.client_addr()); if let Some(username) = self.task_notes.raw_user_name() { adapter.set_client_username(username.clone()); @@ -338,6 +330,15 @@ impl H2ForwardTask { clt_send_rsp, ) .await; + if let Some(dur) = adaptation_state.dur_ups_send_header { + self.http_notes.dur_req_send_hdr = dur; + } + if let Some(dur) = adaptation_state.dur_ups_send_all { + self.http_notes.dur_req_send_all = dur; + } + if let Some(dur) = adaptation_state.dur_ups_recv_header { + self.http_notes.dur_rsp_recv_hdr = dur; + } self.http_notes.clt_req_body_size = adaptation_state.clt_req_body_size; self.http_notes.ups_req_body_size = adaptation_state.ups_req_body_size; match end_state { @@ -502,6 +503,7 @@ impl H2ForwardTask { match r { Ok(rsp) => { if let Some(final_rsp) = self.check_out_final_response(rsp, clt_send_rsp, &mut ups_recv_rsp)? { + record_progress!(); ups_rsp = Some(final_rsp); break; } @@ -630,7 +632,7 @@ impl H2ForwardTask { { Ok(mut adapter) => { let mut adaptation_state = RespmodAdaptationRunState::new( - Instant::now(), + self.task_notes.task_created_instant(), self.http_notes.dur_rsp_recv_hdr, ); adapter.set_client_addr(self.task_notes.client_addr()); @@ -650,6 +652,9 @@ impl H2ForwardTask { clt_send_rsp, ) .await; + if let Some(dur) = adaptation_state.dur_ups_recv_all { + self.http_notes.dur_rsp_recv_all = dur; + } self.http_notes.ups_rsp_body_size = adaptation_state.ups_rsp_body_size; self.http_notes.clt_rsp_body_size = adaptation_state.clt_rsp_body_size; match r { 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 36a14efa2..35289ac69 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 @@ -8,14 +8,13 @@ use std::sync::Arc; use bytes::Bytes; use h2::server::SendResponse; use h2::{RecvStream, SendStream, StreamId}; -use http::{Request, Response, StatusCode, Version, header}; +use http::{Request, Response, StatusCode, Version}; use vey_h2::{H2BodyTransfer, H2ResponseHeaderReceiver, RequestExt}; use vey_icap_client::reqmod::h2::{ H2RequestAdapter, HttpAdapterErrorResponse, ReqmodAdaptationMidState, ReqmodAdaptationRunState, ReqmodRecvHttpResponseBody, }; -use vey_types::acl::AclAction; use vey_types::net::{Host, UpstreamAddr}; use super::{H2StreamTransferError, H2TaskContext, OriginH2Sender}; @@ -128,27 +127,35 @@ impl H2WebsocketTask { clt_r: RecvStream, clt_send_rsp: &mut SendResponse, ) -> Result<(), H2StreamTransferError> { - if self.task_notes.check_layered_rate_limit().is_err() { - self.reply_denied(clt_send_rsp, StatusCode::TOO_MANY_REQUESTS); - return Err(H2StreamTransferError::InternalServerError("rate limited")); - } - if self.task_notes.acquire_site_request_semaphores().is_err() { - self.reply_denied(clt_send_rsp, StatusCode::TOO_MANY_REQUESTS); - return Err(H2StreamTransferError::InternalServerError("fully loaded")); - } - if let Some(tenant) = self.task_notes.tenant_ctx() - && let Some(action) = tenant.check_http_user_agent( - req.headers() - .get_all(header::USER_AGENT) - .iter() - .filter_map(|v| v.to_str().ok()), - ) - && matches!(action, AclAction::Forbid | AclAction::ForbidAndLog) - { - self.reply_denied(clt_send_rsp, StatusCode::FORBIDDEN); - return Err(H2StreamTransferError::InternalServerError("ua denied")); + let mut audit_task = false; + if let Some(site_ctx) = self.task_notes.site_ctx() { + if site_ctx.check_rate_limit().is_err() { + self.reply_denied(clt_send_rsp, StatusCode::TOO_MANY_REQUESTS); + return Err(H2StreamTransferError::InternalServerError("rate limited")); + } + let site_ctx = site_ctx.clone(); + if self + .task_notes + .acquire_site_request_semaphores(&site_ctx) + .is_err() + { + self.reply_denied(clt_send_rsp, StatusCode::TOO_MANY_REQUESTS); + return Err(H2StreamTransferError::InternalServerError("fully loaded")); + } + if let Some(tenant) = site_ctx.tenant_ctx() { + if let Some(audit_handle) = self.ctx.audit_handle.as_ref() { + audit_task = tenant + .user() + .audit() + .do_task_audit() + .unwrap_or_else(|| audit_handle.do_task_audit()); + } + } else if let Some(audit_handle) = self.ctx.audit_handle.as_ref() { + audit_task = audit_handle.do_task_audit(); + } + } else if let Some(audit_handle) = self.ctx.audit_handle.as_ref() { + audit_task = audit_handle.do_task_audit(); } - self.upstream = self .ctx .site_ctx @@ -161,17 +168,6 @@ impl H2WebsocketTask { .checkout_or_connect_h2(&mut self.task_notes, &self.upstream, &request_host) .await?; - let audit_task = self - .task_notes - .tenant_user() - .and_then(|u| u.audit().do_task_audit()) - .unwrap_or_else(|| { - self.ctx - .audit_handle - .as_ref() - .is_some_and(|h| h.do_task_audit()) - }); - if audit_task && let Some(audit_handle) = self.ctx.audit_handle.as_ref() && let Some(reqmod) = audit_handle.icap_reqmod_client() diff --git a/vey-proxy/src/serve/http_proxy/task/connect/task.rs b/vey-proxy/src/serve/http_proxy/task/connect/task.rs index 563966fe2..60e83d032 100644 --- a/vey-proxy/src/serve/http_proxy/task/connect/task.rs +++ b/vey-proxy/src/serve/http_proxy/task/connect/task.rs @@ -258,16 +258,18 @@ impl HttpProxyConnectTask { { let tcp_client_misc_opts; if let Some(user_ctx) = self.task_notes.user_ctx() { - let user_ctx = user_ctx.clone(); - if user_ctx.check_rate_limit().is_err() { self.reply_too_many_requests(clt_w).await; return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::RateLimited, )); } - - if self.task_notes.acquire_user_request_semaphore().is_err() { + let user_ctx = user_ctx.clone(); + if self + .task_notes + .acquire_user_request_semaphore(&user_ctx) + .is_err() + { self.reply_too_many_requests(clt_w).await; return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::FullyLoaded, diff --git a/vey-proxy/src/serve/http_proxy/task/connect_udp/task.rs b/vey-proxy/src/serve/http_proxy/task/connect_udp/task.rs index fd00e9ba7..09c5ddf2f 100644 --- a/vey-proxy/src/serve/http_proxy/task/connect_udp/task.rs +++ b/vey-proxy/src/serve/http_proxy/task/connect_udp/task.rs @@ -288,16 +288,18 @@ impl HttpProxyConnectUdpTask { { let tcp_client_misc_opts; if let Some(user_ctx) = self.task_notes.user_ctx() { - let user_ctx = user_ctx.clone(); - if user_ctx.check_rate_limit().is_err() { self.reply_too_many_requests(clt_w).await; return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::RateLimited, )); } - - if self.task_notes.acquire_user_request_semaphore().is_err() { + let user_ctx = user_ctx.clone(); + if self + .task_notes + .acquire_user_request_semaphore(&user_ctx) + .is_err() + { self.reply_too_many_requests(clt_w).await; return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::FullyLoaded, diff --git a/vey-proxy/src/serve/http_proxy/task/forward/task.rs b/vey-proxy/src/serve/http_proxy/task/forward/task.rs index a36b1ee24..40f2b25e2 100644 --- a/vey-proxy/src/serve/http_proxy/task/forward/task.rs +++ b/vey-proxy/src/serve/http_proxy/task/forward/task.rs @@ -535,16 +535,18 @@ impl<'a> HttpProxyForwardTask<'a> { let mut audit_task = false; if let Some(user_ctx) = self.task_notes.user_ctx() { - let user_ctx = user_ctx.clone(); - if user_ctx.check_rate_limit().is_err() { self.reply_too_many_requests(clt_w).await; return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::RateLimited, )); } - - if self.task_notes.acquire_user_request_semaphore().is_err() { + let user_ctx = user_ctx.clone(); + if self + .task_notes + .acquire_user_request_semaphore(&user_ctx) + .is_err() + { self.reply_too_many_requests(clt_w).await; return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::FullyLoaded, @@ -1054,12 +1056,15 @@ impl<'a> HttpProxyForwardTask<'a> { let mut body_reader = recv_body.body_reader(); let mut copy_to_clt = StreamCopy::new(&mut body_reader, clt_w, &self.ctx.server_config.tcp_copy); - (&mut copy_to_clt).await.map_err(|e| match e { - StreamCopyError::ReadFailed(e) => ServerTaskError::InternalAdapterError(anyhow!( - "read http error response from adapter failed: {e:?}" - )), - StreamCopyError::WriteFailed(e) => ServerTaskError::ClientTcpWriteFailed(e), - })?; + if let Err(e) = (&mut copy_to_clt).await { + self.http_notes.clt_rsp_body_size = Some(copy_to_clt.reader().body_size()); + return Err(match e { + StreamCopyError::ReadFailed(e) => ServerTaskError::InternalAdapterError( + anyhow!("read http error response from adapter failed: {e:?}"), + ), + StreamCopyError::WriteFailed(e) => ServerTaskError::ClientTcpWriteFailed(e), + }); + } self.http_notes.clt_rsp_body_size = Some(copy_to_clt.reader().body_size()); recv_body.save_connection().await; } else { @@ -1236,7 +1241,12 @@ impl<'a> HttpProxyForwardTask<'a> { .map_err(ServerTaskError::UpstreamWriteFailed)?; self.http_notes.mark_req_send_hdr(); self.http_notes.mark_req_send_all(); - self.http_notes.ups_req_body_size = Some(body.len() as u64); + // Chunked bodies are buffered on the wire, while clt_req_body_size is the + // decoded payload. A fully read body was already counted that way. + self.http_notes.ups_req_body_size = self + .http_notes + .clt_req_body_size + .or(Some(body.len() as u64)); match tokio::time::timeout( self.rsp_hdr_recv_timeout(), @@ -1466,6 +1476,7 @@ impl<'a> HttpProxyForwardTask<'a> { let copy_done = clt_to_ups.finished(); let mut rsp_header = match rsp_header { Some(header) => { + record_progress!(); if !clt_body_reader.finished() { // not all client data read in, drop the client connection self.should_close = true; diff --git a/vey-proxy/src/serve/http_proxy/task/ftp/task.rs b/vey-proxy/src/serve/http_proxy/task/ftp/task.rs index faf2ac82f..31f96d80a 100644 --- a/vey-proxy/src/serve/http_proxy/task/ftp/task.rs +++ b/vey-proxy/src/serve/http_proxy/task/ftp/task.rs @@ -354,16 +354,18 @@ impl<'a> FtpOverHttpTask<'a> { let tcp_client_misc_opts; if let Some(user_ctx) = self.task_notes.user_ctx() { - let user_ctx = user_ctx.clone(); - if user_ctx.check_rate_limit().is_err() { self.reply_too_many_requests(clt_w).await; return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::RateLimited, )); } - - if self.task_notes.acquire_user_request_semaphore().is_err() { + let user_ctx = user_ctx.clone(); + if self + .task_notes + .acquire_user_request_semaphore(&user_ctx) + .is_err() + { self.reply_too_many_requests(clt_w).await; return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::FullyLoaded, diff --git a/vey-proxy/src/serve/sni_proxy/task/relay/task.rs b/vey-proxy/src/serve/sni_proxy/task/relay/task.rs index 39b216c44..125dadd44 100644 --- a/vey-proxy/src/serve/sni_proxy/task/relay/task.rs +++ b/vey-proxy/src/serve/sni_proxy/task/relay/task.rs @@ -149,15 +149,17 @@ impl TcpStreamTask { let tcp_client_misc_opts; if let Some(user_ctx) = self.task_notes.user_ctx() { - let user_ctx = user_ctx.clone(); - if user_ctx.check_rate_limit().is_err() { return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::RateLimited, )); } - - if self.task_notes.acquire_user_request_semaphore().is_err() { + let user_ctx = user_ctx.clone(); + if self + .task_notes + .acquire_user_request_semaphore(&user_ctx) + .is_err() + { return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::FullyLoaded, )); diff --git a/vey-proxy/src/serve/socks_proxy/task/tcp_connect/task.rs b/vey-proxy/src/serve/socks_proxy/task/tcp_connect/task.rs index c7f8c4348..2f4e834bc 100644 --- a/vey-proxy/src/serve/socks_proxy/task/tcp_connect/task.rs +++ b/vey-proxy/src/serve/socks_proxy/task/tcp_connect/task.rs @@ -212,16 +212,18 @@ impl SocksProxyTcpConnectTask { let tcp_client_misc_opts; if let Some(user_ctx) = self.task_notes.user_ctx() { - let user_ctx = user_ctx.clone(); - if user_ctx.check_rate_limit().is_err() { self.reply_forbidden(&mut clt_w).await; return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::RateLimited, )); } - - if self.task_notes.acquire_user_request_semaphore().is_err() { + let user_ctx = user_ctx.clone(); + if self + .task_notes + .acquire_user_request_semaphore(&user_ctx) + .is_err() + { self.reply_forbidden(&mut clt_w).await; return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::FullyLoaded, diff --git a/vey-proxy/src/serve/socks_proxy/task/udp_associate/task.rs b/vey-proxy/src/serve/socks_proxy/task/udp_associate/task.rs index 92bbb93d7..0f56f2a51 100644 --- a/vey-proxy/src/serve/socks_proxy/task/udp_associate/task.rs +++ b/vey-proxy/src/serve/socks_proxy/task/udp_associate/task.rs @@ -168,16 +168,18 @@ impl SocksProxyUdpAssociateTask { W: AsyncWrite + Unpin, { if let Some(user_ctx) = self.task_notes.user_ctx() { - let user_ctx = user_ctx.clone(); - if user_ctx.check_rate_limit().is_err() { self.reply_forbidden(&mut clt_tcp_w).await; return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::RateLimited, )); } - - if self.task_notes.acquire_user_request_semaphore().is_err() { + let user_ctx = user_ctx.clone(); + if self + .task_notes + .acquire_user_request_semaphore(&user_ctx) + .is_err() + { self.reply_forbidden(&mut clt_tcp_w).await; return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::FullyLoaded, diff --git a/vey-proxy/src/serve/socks_proxy/task/udp_connect/task.rs b/vey-proxy/src/serve/socks_proxy/task/udp_connect/task.rs index e151ecbea..50f048ef7 100644 --- a/vey-proxy/src/serve/socks_proxy/task/udp_connect/task.rs +++ b/vey-proxy/src/serve/socks_proxy/task/udp_connect/task.rs @@ -219,16 +219,18 @@ impl SocksProxyUdpConnectTask { W: AsyncWrite + Unpin, { if let Some(user_ctx) = self.task_notes.user_ctx() { - let user_ctx = user_ctx.clone(); - if user_ctx.check_rate_limit().is_err() { self.reply_forbidden(&mut clt_tcp_w).await; return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::RateLimited, )); } - - if self.task_notes.acquire_user_request_semaphore().is_err() { + let user_ctx = user_ctx.clone(); + if self + .task_notes + .acquire_user_request_semaphore(&user_ctx) + .is_err() + { self.reply_forbidden(&mut clt_tcp_w).await; return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::FullyLoaded, diff --git a/vey-proxy/src/serve/task.rs b/vey-proxy/src/serve/task.rs index fa2b27f56..dc4a6314c 100644 --- a/vey-proxy/src/serve/task.rs +++ b/vey-proxy/src/serve/task.rs @@ -246,30 +246,20 @@ impl ServerTaskNotes { ) } - pub(crate) fn check_layered_rate_limit(&self) -> Result<(), ()> { - if let Some(site_ctx) = &self.site_ctx { - site_ctx.check_rate_limit()?; - } - if let Some(user_ctx) = &self.user_ctx { - user_ctx.check_rate_limit()?; - } - Ok(()) - } - - /// Tenant then site. No-op when this notes has no site. - pub(crate) fn acquire_site_request_semaphores(&mut self) -> Result<(), ()> { - let Some(site_ctx) = &self.site_ctx else { - return Ok(()); - }; + /// Tenant then site. `site_ctx` is the context the caller already resolved. + pub(crate) fn acquire_site_request_semaphores( + &mut self, + site_ctx: &SiteContext, + ) -> Result<(), ()> { self._site_req_alive_permits = site_ctx.acquire_request_semaphores()?; Ok(()) } - /// No-op when this notes has no user. - pub(crate) fn acquire_user_request_semaphore(&mut self) -> Result<(), ()> { - let Some(user_ctx) = &self.user_ctx else { - return Ok(()); - }; + /// `user_ctx` is the context the caller already resolved. + pub(crate) fn acquire_user_request_semaphore( + &mut self, + user_ctx: &UserContext, + ) -> Result<(), ()> { self._user_req_alive_permit = Some(user_ctx.acquire_request_semaphore()?); Ok(()) } diff --git a/vey-proxy/src/serve/tcp_stream/task.rs b/vey-proxy/src/serve/tcp_stream/task.rs index 227152652..74df3a6dd 100644 --- a/vey-proxy/src/serve/tcp_stream/task.rs +++ b/vey-proxy/src/serve/tcp_stream/task.rs @@ -133,15 +133,17 @@ impl TcpStreamTask { let tcp_client_misc_opts; if let Some(user_ctx) = self.task_notes.user_ctx() { - let user_ctx = user_ctx.clone(); - if user_ctx.check_rate_limit().is_err() { return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::RateLimited, )); } - - if self.task_notes.acquire_user_request_semaphore().is_err() { + let user_ctx = user_ctx.clone(); + if self + .task_notes + .acquire_user_request_semaphore(&user_ctx) + .is_err() + { return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::FullyLoaded, )); diff --git a/vey-proxy/src/serve/tcp_tproxy/task.rs b/vey-proxy/src/serve/tcp_tproxy/task.rs index ba89c1a2e..ffb768006 100644 --- a/vey-proxy/src/serve/tcp_tproxy/task.rs +++ b/vey-proxy/src/serve/tcp_tproxy/task.rs @@ -125,15 +125,17 @@ impl TProxyStreamTask { let tcp_client_misc_opts; if let Some(user_ctx) = self.task_notes.user_ctx() { - let user_ctx = user_ctx.clone(); - if user_ctx.check_rate_limit().is_err() { return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::RateLimited, )); } - - if self.task_notes.acquire_user_request_semaphore().is_err() { + let user_ctx = user_ctx.clone(); + if self + .task_notes + .acquire_user_request_semaphore(&user_ctx) + .is_err() + { return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::FullyLoaded, )); diff --git a/vey-proxy/src/serve/tls_proxy/server.rs b/vey-proxy/src/serve/tls_proxy/server.rs index 049ba83b0..b8669a1cb 100644 --- a/vey-proxy/src/serve/tls_proxy/server.rs +++ b/vey-proxy/src/serve/tls_proxy/server.rs @@ -201,12 +201,18 @@ impl TlsProxyServer { self.config.name(), self.server_stats.share_extra_tags(), ); - let task_notes = - ServerTaskNotes::new(cc_info.clone(), None, Duration::ZERO).with_site_ctx(site_ctx); + let task_notes = ServerTaskNotes::new(cc_info.clone(), None, Duration::ZERO); let ctx = self.get_common_task_context(cc_info); - TlsRelayTask::new(ctx, host, sni_host, self.audit_context(), task_notes) - .into_running_from_tls(stream) - .await; + TlsRelayTask::new( + ctx, + host, + sni_host, + self.audit_context(), + task_notes, + site_ctx, + ) + .into_running_from_tls(stream) + .await; } fn get_proxy_host(&self, sni_host: &Host) -> Option> { diff --git a/vey-proxy/src/serve/tls_proxy/task/accept/task.rs b/vey-proxy/src/serve/tls_proxy/task/accept/task.rs index 9d1bb95ea..835f17755 100644 --- a/vey-proxy/src/serve/tls_proxy/task/accept/task.rs +++ b/vey-proxy/src/serve/tls_proxy/task/accept/task.rs @@ -4,7 +4,6 @@ */ use std::sync::Arc; -use std::time::Duration; use anyhow::anyhow; use bytes::BytesMut; @@ -12,6 +11,7 @@ use log::debug; use openssl::ssl::Ssl; use tokio::io::AsyncReadExt; use tokio::net::TcpStream; +use tokio::time::Instant; use vey_codec::tls::{ ClientHello, ExtensionType, HandshakeCoalescer, Record, RecordHeader, RecordParseError, @@ -36,6 +36,7 @@ pub(crate) struct TlsAcceptTask { ctx: CommonTaskContext, hosts: Arc>>, audit_ctx: AuditContext, + time_accepted: Instant, } impl TlsAcceptTask { @@ -48,6 +49,7 @@ impl TlsAcceptTask { ctx, hosts, audit_ctx, + time_accepted: Instant::now(), } } @@ -113,11 +115,18 @@ impl TlsAcceptTask { } let ssl_stream = self.accept_tls(host, stream, clt_r_buf).await?; - let task_notes = ServerTaskNotes::new(self.ctx.cc_info.clone(), None, Duration::ZERO) - .with_site_ctx(site_ctx); - TlsRelayTask::new(self.ctx, host.clone(), req_host, self.audit_ctx, task_notes) - .into_running(ssl_stream) - .await; + let task_notes = + ServerTaskNotes::new(self.ctx.cc_info.clone(), None, self.time_accepted.elapsed()); + TlsRelayTask::new( + self.ctx, + host.clone(), + req_host, + self.audit_ctx, + task_notes, + site_ctx, + ) + .into_running(ssl_stream) + .await; Ok(()) } diff --git a/vey-proxy/src/serve/tls_proxy/task/relay/task.rs b/vey-proxy/src/serve/tls_proxy/task/relay/task.rs index 5a5cd8d51..9b31d4fc4 100644 --- a/vey-proxy/src/serve/tls_proxy/task/relay/task.rs +++ b/vey-proxy/src/serve/tls_proxy/task/relay/task.rs @@ -28,6 +28,7 @@ use crate::serve::{ ServerStats, ServerTaskError, ServerTaskForbiddenError, ServerTaskNotes, ServerTaskResult, ServerTaskStage, }; +use crate::site::SiteContext; use crate::stat::types::RequestAliveKind; pub(crate) struct TlsRelayTask { @@ -36,6 +37,7 @@ pub(crate) struct TlsRelayTask { req_host: Host, upstream: UpstreamAddr, egress_notes: EgressNotes, + site_ctx: SiteContext, task_notes: ServerTaskNotes, task_stats: Arc, audit_ctx: AuditContext, @@ -49,17 +51,20 @@ impl TlsRelayTask { req_host: Host, audit_ctx: AuditContext, task_notes: ServerTaskNotes, + site_ctx: SiteContext, ) -> Self { let upstream = host .site() .select_upstream(task_notes.client_ip()) .unwrap_or_else(|_| UpstreamAddr::empty()); + let task_notes = task_notes.with_site_ctx(site_ctx.clone()); TlsRelayTask { ctx, host, req_host, upstream, egress_notes: EgressNotes::default(), + site_ctx, task_notes, task_stats: Arc::new(TcpStreamTaskStats::default()), audit_ctx, @@ -83,10 +88,6 @@ impl TlsRelayTask { }) } - pub(crate) fn tenant_ctx(&self) -> Option<&crate::auth::TenantContext> { - self.task_notes.tenant_ctx() - } - /// TLS ingress. Tenant expiry and block delay are checked before the relay starts. pub(crate) async fn into_running_from_tls(mut self, stream: S) where @@ -138,7 +139,7 @@ impl TlsRelayTask { S::R: AsyncRead + Send + Sync + Unpin + 'static, S::W: AsyncWrite + Send + Sync + Unpin + 'static, { - if let Some(tenant) = self.tenant_ctx() { + if let Some(tenant) = self.site_ctx.tenant_ctx() { if tenant.is_expired() { return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::UserBlocked, @@ -163,23 +164,23 @@ impl TlsRelayTask { S::R: AsyncRead + Send + Sync + Unpin + 'static, S::W: AsyncWrite + Send + Sync + Unpin + 'static, { - if self.task_notes.check_layered_rate_limit().is_err() { + if self.site_ctx.check_rate_limit().is_err() { return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::RateLimited, )); } - - if self.task_notes.site_ctx().is_none() { - return Err(ServerTaskError::InternalServerError("no site context")); - } - if self.task_notes.acquire_site_request_semaphores().is_err() { + if self + .task_notes + .acquire_site_request_semaphores(&self.site_ctx) + .is_err() + { return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::FullyLoaded, )); } - - let tcp_client_misc_opts = if let Some(user) = self.task_notes.tenant_user() { - user.config() + let tcp_client_misc_opts = if let Some(tenant) = self.site_ctx.tenant_ctx() { + tenant + .user_config() .tcp_client_misc_opts(&self.ctx.server_config.tcp_misc_opts) } else { Cow::Borrowed(&self.ctx.server_config.tcp_misc_opts) @@ -271,7 +272,7 @@ impl TlsRelayTask { if let Some(audit_handle) = self.audit_ctx.check_take_handle() { let audit_task = self - .task_notes + .site_ctx .tenant_user() .map(|user| { let audit = user.audit(); @@ -336,7 +337,7 @@ impl TlsRelayTask { .site() .tcp_sock_speed_limit() .shrink_as_smaller(&limit_config); - if let Some(user) = self.task_notes.tenant_user() { + if let Some(user) = self.site_ctx.tenant_user() { limit_config = user .config() .tcp_sock_speed_limit @@ -357,7 +358,7 @@ impl TlsRelayTask { wrapper_stats, ); - if let Some(user) = self.task_notes.tenant_user() { + if let Some(user) = self.site_ctx.tenant_user() { if let Some(limiter) = user.tcp_all_upload_speed_limit() { clt_r.add_global_limiter(limiter.clone()); } diff --git a/vey-proxy/src/serve/tls_stream/task.rs b/vey-proxy/src/serve/tls_stream/task.rs index 54b6c5de0..fb8455cdc 100644 --- a/vey-proxy/src/serve/tls_stream/task.rs +++ b/vey-proxy/src/serve/tls_stream/task.rs @@ -126,15 +126,17 @@ impl TlsStreamTask { let tcp_client_misc_opts; if let Some(user_ctx) = self.task_notes.user_ctx() { - let user_ctx = user_ctx.clone(); - if user_ctx.check_rate_limit().is_err() { return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::RateLimited, )); } - - if self.task_notes.acquire_user_request_semaphore().is_err() { + let user_ctx = user_ctx.clone(); + if self + .task_notes + .acquire_user_request_semaphore(&user_ctx) + .is_err() + { return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::FullyLoaded, )); diff --git a/vey-proxy/src/serve/udp_stream/task.rs b/vey-proxy/src/serve/udp_stream/task.rs index 19d7d2386..560a4a80b 100644 --- a/vey-proxy/src/serve/udp_stream/task.rs +++ b/vey-proxy/src/serve/udp_stream/task.rs @@ -138,15 +138,17 @@ impl UdpStreamTask { clt_w: AcceptedUdpPacketSender, ) -> ServerTaskResult<()> { if let Some(user_ctx) = self.task_notes.user_ctx() { - let user_ctx = user_ctx.clone(); - if user_ctx.check_rate_limit().is_err() { return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::RateLimited, )); } - - if self.task_notes.acquire_user_request_semaphore().is_err() { + let user_ctx = user_ctx.clone(); + if self + .task_notes + .acquire_user_request_semaphore(&user_ctx) + .is_err() + { return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::FullyLoaded, )); diff --git a/vey-proxy/src/serve/udp_tproxy/task.rs b/vey-proxy/src/serve/udp_tproxy/task.rs index 96cf3c00d..85dac012a 100644 --- a/vey-proxy/src/serve/udp_tproxy/task.rs +++ b/vey-proxy/src/serve/udp_tproxy/task.rs @@ -137,15 +137,17 @@ impl TProxyStreamTask { clt_w: AcceptedUdpPacketSender, ) -> ServerTaskResult<()> { if let Some(user_ctx) = self.task_notes.user_ctx() { - let user_ctx = user_ctx.clone(); - if user_ctx.check_rate_limit().is_err() { return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::RateLimited, )); } - - if self.task_notes.acquire_user_request_semaphore().is_err() { + let user_ctx = user_ctx.clone(); + if self + .task_notes + .acquire_user_request_semaphore(&user_ctx) + .is_err() + { return Err(ServerTaskError::ForbiddenByRule( ServerTaskForbiddenError::FullyLoaded, ));