diff --git a/Cargo.lock b/Cargo.lock index 0ad3e20..96c11e4 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -383,6 +383,12 @@ dependencies = [ "tracing", ] +[[package]] +name = "hashbrown" +version = "0.14.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e5274423e17b7c9fc20b6e7e208532f9b19825d82dfd615708b70edd83df41f1" + [[package]] name = "hashbrown" version = "0.15.2" @@ -495,8 +501,10 @@ dependencies = [ "log", "magnus", "rb-sys", + "socket2", "tokio", "tokio-stream", + "tokio-util", ] [[package]] @@ -506,7 +514,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8c9c992b02b5b4c94ea26e32fe5bccb7aa7d9f390ab5c1221ff895bc7ea8b652" dependencies = [ "equivalent", - "hashbrown", + "hashbrown 0.15.2", ] [[package]] @@ -931,6 +939,8 @@ dependencies = [ "bytes", "futures-core", "futures-sink", + "futures-util", + "hashbrown 0.14.5", "pin-project-lite", "tokio", ] diff --git a/LICENSE.txt b/LICENSE.txt index e04955b..4e9303c 100644 --- a/LICENSE.txt +++ b/LICENSE.txt @@ -19,3 +19,12 @@ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. + +Third-party components +---------------------- + +ext/hyper_ruby/src/syslog/framing.rs is a port of Vector 0.48.0's +lib/codecs/src/decoding/framing/octet_counting.rs +(https://github.com/vectordotdev/vector) and is licensed under the Mozilla +Public License 2.0 (https://www.mozilla.org/en-US/MPL/2.0/), not the MIT +licence above. diff --git a/README.md b/README.md index eaf8434..a70ba1b 100644 --- a/README.md +++ b/README.md @@ -10,10 +10,23 @@ After checking out the repo, run `bin/setup` to install dependencies. Then, run To install this gem onto your local machine, run `bundle exec rake install`. To release a new version, update the version number in `version.rb`, and then run `bundle exec rake release`, which will create a git tag for the version, push git commits and the created tag, and push the `.gem` file to [rubygems.org](https://rubygems.org). +## Syslog listeners + +Alongside HTTP, the server can accept syslog over a stream socket (Unix socket +or TCP, optionally behind a PROXY v2 header) and over UDP. Each complete message +is yielded to the same `Server#run_worker` block as requests, as a +`HyperRuby::SyslogMessage` the block answers with an admission verdict. See the +configuration keys documented at the top of `lib/hyper_ruby.rb`. + ## License The gem is available as open source under the terms of the [MIT License](https://opensource.org/licenses/MIT). +`ext/hyper_ruby/src/syslog/framing.rs` is an exception: it is a port of +Vector 0.48.0's `lib/codecs/src/decoding/framing/octet_counting.rs` and is +licensed under the [Mozilla Public License 2.0](https://www.mozilla.org/en-US/MPL/2.0/), +which is why the gem's metadata lists both licences. + ## Code of Conduct Everyone interacting in the HyperRuby project's codebases, issue trackers, chat rooms and mailing lists is expected to follow the [code of conduct](https://github.com/[USERNAME]/hyper_ruby/blob/master/CODE_OF_CONDUCT.md). diff --git a/ext/hyper_ruby/Cargo.toml b/ext/hyper_ruby/Cargo.toml index 2844520..4dae9cf 100644 --- a/ext/hyper_ruby/Cargo.toml +++ b/ext/hyper_ruby/Cargo.toml @@ -3,7 +3,7 @@ name = "hyper_ruby" version = "0.1.0" edition = "2021" authors = ["alistairjevans "] -license = "MIT" +license = "MIT AND MPL-2.0" publish = false [lib] @@ -26,3 +26,5 @@ async-stream = "0.3.5" env_logger = "0.11" log = "0.4" form_urlencoded = "1.2.1" +tokio-util = { version = "0.7", features = ["codec", "rt"] } +socket2 = "0.5" diff --git a/ext/hyper_ruby/src/lib.rs b/ext/hyper_ruby/src/lib.rs index 7a89261..cd82792 100644 --- a/ext/hyper_ruby/src/lib.rs +++ b/ext/hyper_ruby/src/lib.rs @@ -2,6 +2,7 @@ mod request; mod response; mod gvl_helpers; mod grpc; +mod syslog; use hyper_util::server::graceful::GracefulShutdown; use request::{Request, GrpcRequest}; @@ -10,7 +11,7 @@ use gvl_helpers::nogvl; use magnus::block::block_proc; use magnus::typed_data::Obj; -use magnus::{function, method, prelude::*, Error as MagnusError, IntoValue, Ruby, Value, RString}; +use magnus::{function, method, prelude::*, Error as MagnusError, IntoValue, RHash, Ruby, Value, RString}; use bytes::Bytes; use tokio::io::{AsyncRead, AsyncWrite}; @@ -48,6 +49,10 @@ use tokio::sync::broadcast; static LOGGER_INIT: Once = Once::new(); +// How long stop() waits for the syslog listeners to finish delivering messages +// they have already framed. +const SYSLOG_DRAIN_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(10); + #[global_allocator] static GLOBAL: Jemalloc = Jemalloc; @@ -83,6 +88,7 @@ struct ServerConfig { channel_capacity: usize, send_timeout: u64, max_connection_age: Option, + syslog: syslog::SyslogConfig, } impl ServerConfig { @@ -95,6 +101,7 @@ impl ServerConfig { channel_capacity: 5000, // Default capacity for worker channel send_timeout: 1000, // Default 1 second timeout for send backpressure max_connection_age: None, // No limit by default + syslog: syslog::SyslogConfig::new(), } } } @@ -105,6 +112,13 @@ struct RequestWithCompletion { response_tx: oneshot::Sender>, } +// A unit of work for a Ruby worker thread; HTTP requests and syslog messages +// have their own channels but share the worker threads. +enum WorkItem { + Http(RequestWithCompletion), + Syslog(syslog::SyslogDelivery), +} + #[magnus::wrap(class = "HyperRuby::Server")] struct Server { server_handle: Arc>>>, @@ -117,6 +131,10 @@ struct Server { // (dev, ino) of the Unix socket file this server bound, so stop() only removes // the file if a replacement server hasn't taken over the path in the meantime. socket_ident: RefCell>, + syslog_work_rx: RefCell>>, + syslog_work_tx: RefCell>>>, + syslog_counters: Arc, + syslog_runtime: RefCell>, } impl Server { @@ -131,6 +149,10 @@ impl Server { shutdown: RefCell::new(None), total_connections: Arc::new(AtomicU64::new(0)), socket_ident: RefCell::new(None), + syslog_work_rx: RefCell::new(None), + syslog_work_tx: RefCell::new(None), + syslog_counters: Arc::new(syslog::SyslogCounters::default()), + syslog_runtime: RefCell::new(None), } } @@ -168,6 +190,8 @@ impl Server { server_config.max_connection_age = Some(u64::try_convert(max_connection_age)?); } + server_config.syslog.apply(&config)?; + // Initialize logging if not already initialized LOGGER_INIT.call_once(|| { let mut builder = env_logger::Builder::from_env(env_logger::Env::default()); @@ -190,26 +214,36 @@ impl Server { // Method that Ruby worker threads will call with a block pub fn run_worker(&self) -> Result<(), MagnusError> { let block = block_proc().unwrap(); - + // Check if we have a work_rx channel, error out if not let work_rx = self.work_rx.borrow().as_ref().ok_or_else(|| { MagnusError::new(magnus::exception::runtime_error(), "Server must be started before running workers") })?.clone(); - + + let syslog_rx = self.syslog_work_rx.borrow().as_ref().ok_or_else(|| { + MagnusError::new(magnus::exception::runtime_error(), "Server must be started before running workers") + })?.clone(); + + let (syslog_enabled, work_ratio) = { + let config = self.config.borrow(); + (config.syslog.enabled(), config.syslog.work_ratio) + }; + let mut syslog_streak = 0; + loop { - // try getting the next request without yielding the GVL, if there's nothing, wait for one - let work_request = match work_rx.try_recv() { - Ok(work_request) => Ok(work_request), - Err(crossbeam_channel::TryRecvError::Empty) => { - nogvl(|| work_rx.recv()) - }, - Err(crossbeam_channel::TryRecvError::Disconnected) => { - break; - } + let next = if syslog_enabled { + next_work_item(&work_rx, &syslog_rx, &mut syslog_streak, work_ratio) + } else { + next_http_work_item(&work_rx) + }; + + let work_item = match next { + Some(work_item) => work_item, + None => break, }; - match work_request { - Ok(work_request) => { + match work_item { + WorkItem::Http(work_request) => { let hyper_request = work_request.request; debug!("Processing request:"); @@ -258,13 +292,16 @@ impl Server { Err(e) => error!("Failed to send response back to client: {:?} - response dropped", e), } } - Err(_) => { - // Channel closed, exit thread - break; + WorkItem::Syslog(delivery) => { + let (message, result_tx) = delivery.into_parts(); + let result = call_block_with_syslog_message(&block, message); + if result_tx.send(result).is_err() { + debug!("Syslog listener stopped waiting for an admission result"); + } } } } - + Ok(()) } @@ -279,7 +316,13 @@ impl Server { *self.work_rx.borrow_mut() = Some(work_rx); let work_tx = Arc::new(work_tx); *self.work_tx.borrow_mut() = Some(work_tx.clone()); - + + // Syslog messages queue separately but are drained by the same workers. + let (syslog_work_tx, syslog_work_rx) = crossbeam_channel::bounded(config.channel_capacity); + *self.syslog_work_rx.borrow_mut() = Some(syslog_work_rx); + let syslog_work_tx = Arc::new(syslog_work_tx); + *self.syslog_work_tx.borrow_mut() = Some(syslog_work_tx.clone()); + let (shutdown_tx, shutdown_rx) = broadcast::channel(1); *self.shutdown.borrow_mut() = Some(shutdown_tx.clone()); @@ -300,6 +343,9 @@ impl Server { *self.runtime.borrow_mut() = Some(rt.clone()); + let syslog_config = config.syslog.clone(); + let syslog_counters = self.syslog_counters.clone(); + rt.block_on(async move { // Instead of spawning a task, we'll run the server setup inline first to catch binding errors // Setup listener and http server components @@ -456,6 +502,19 @@ impl Server { } }); + if syslog_config.enabled() { + match syslog::start(&syslog_config, syslog_counters, syslog_work_tx) { + Ok(syslog_runtime) => *self.syslog_runtime.borrow_mut() = Some(syslog_runtime), + Err(e) => { + error!("Failed to start syslog listeners: {}", e); + return Err(MagnusError::new( + magnus::exception::runtime_error(), + format!("Failed to start syslog listeners: {}", e) + )); + } + } + } + let mut handle = self.server_handle.lock().await; *handle = Some(server_task); @@ -465,22 +524,51 @@ impl Server { Ok(()) } + // True while the configured syslog listeners are bound and accepting. + pub fn syslog_listening(&self) -> bool { + self.syslog_runtime.borrow().as_ref().map(|runtime| runtime.listening()).unwrap_or(false) + } + + pub fn syslog_stats(&self) -> Result { + self.syslog_counters.to_hash() + } + pub fn stop(&self) -> Result<(), MagnusError> { - if let Some(rt) = self.runtime.borrow().as_ref() { + // Take owned handles first: no RefCell borrow may be held while the GVL + // is released, or a concurrent Ruby call panics on borrow_mut. + let runtime = self.runtime.borrow().clone(); + let syslog_runtime = self.syslog_runtime.borrow_mut().take(); + + if let Some(rt) = runtime.as_ref() { if let Some(shutdown) = self.shutdown.borrow().as_ref() { let _ = shutdown.send(()); } + if let Some(syslog_runtime) = syslog_runtime.as_ref() { + syslog_runtime.stop_accepting(); + } + rt.block_on(async { let mut handle = self.server_handle.lock().await; if let Some(task) = handle.take() { task.await.unwrap_or_else(|e| warn!("Server task failed: {:?}", e)); } }); + + if let Some(syslog_runtime) = syslog_runtime.as_ref() { + // Release the GVL so worker threads can accept the messages the + // listeners have already framed. + nogvl(|| rt.block_on(syslog_runtime.drain(SYSLOG_DRAIN_TIMEOUT))); + } + } + + if let Some(syslog_runtime) = syslog_runtime.as_ref() { + syslog_runtime.remove_socket_file(); } // Drop the channel and runtime self.work_tx.borrow_mut().take(); + self.syslog_work_tx.borrow_mut().take(); self.runtime.borrow_mut().take(); self.shutdown.borrow_mut().take(); @@ -505,6 +593,146 @@ impl Server { } } +// Take the next request for a Ruby worker when no syslog listener is running. +fn next_http_work_item( + work_rx: &crossbeam_channel::Receiver, +) -> Option { + // try getting the next request without yielding the GVL, if there's nothing, wait for one + match work_rx.try_recv() { + Ok(request) => Some(WorkItem::Http(request)), + Err(crossbeam_channel::TryRecvError::Empty) => { + nogvl(|| work_rx.recv()).ok().map(WorkItem::Http) + }, + Err(crossbeam_channel::TryRecvError::Disconnected) => None, + } +} + +// Take the next piece of work for a Ruby worker, taking whatever is already +// queued so we only release the GVL when both channels are empty. Syslog +// messages go first until work_ratio of them have run ahead of a waiting +// request, then a request goes and the count starts again; the blocking select +// picks uniformly between the two. +fn next_work_item( + work_rx: &crossbeam_channel::Receiver, + syslog_rx: &crossbeam_channel::Receiver, + syslog_streak: &mut u32, + work_ratio: u32, +) -> Option { + let syslog_first = *syslog_streak < work_ratio; + + if syslog_first { + match try_syslog_work_item(syslog_rx) { + TryWork::Empty => (), + outcome => return count_work_item(outcome, syslog_streak), + } + } + + match try_http_work_item(work_rx) { + TryWork::Empty => (), + outcome => return count_work_item(outcome, syslog_streak), + } + + if !syslog_first { + match try_syslog_work_item(syslog_rx) { + TryWork::Empty => (), + outcome => return count_work_item(outcome, syslog_streak), + } + } + + let selected = nogvl(|| { + crossbeam_channel::select! { + recv(work_rx) -> request => request.ok().map(WorkItem::Http), + recv(syslog_rx) -> delivery => delivery.ok().map(WorkItem::Syslog), + } + }); + + match selected { + Some(work_item) => count_work_item(TryWork::Found(work_item), syslog_streak), + None => None, + } +} + +fn count_work_item(outcome: TryWork, syslog_streak: &mut u32) -> Option { + match &outcome { + TryWork::Found(WorkItem::Syslog(_)) => *syslog_streak += 1, + TryWork::Found(WorkItem::Http(_)) => *syslog_streak = 0, + _ => (), + } + outcome.into_work_item() +} + +enum TryWork { + Found(WorkItem), + Empty, + Closed, +} + +impl TryWork { + fn into_work_item(self) -> Option { + match self { + TryWork::Found(work_item) => Some(work_item), + // A closed channel stops the worker, as it always has. + TryWork::Empty | TryWork::Closed => None, + } + } +} + +fn try_http_work_item(work_rx: &crossbeam_channel::Receiver) -> TryWork { + match work_rx.try_recv() { + Ok(request) => TryWork::Found(WorkItem::Http(request)), + Err(crossbeam_channel::TryRecvError::Empty) => TryWork::Empty, + Err(crossbeam_channel::TryRecvError::Disconnected) => TryWork::Closed, + } +} + +fn try_syslog_work_item(syslog_rx: &crossbeam_channel::Receiver) -> TryWork { + match syslog_rx.try_recv() { + Ok(delivery) => TryWork::Found(WorkItem::Syslog(delivery)), + Err(crossbeam_channel::TryRecvError::Empty) => TryWork::Empty, + Err(crossbeam_channel::TryRecvError::Disconnected) => TryWork::Closed, + } +} + +static SYSLOG_RESPONSE_WARNING: Once = Once::new(); + +// Hand one syslog message to the worker block; anything other than a truthy +// result means the message was not admitted. +fn call_block_with_syslog_message( + block: &magnus::block::Proc, + message: syslog::SyslogMessage, +) -> syslog::HandlerResult { + let message_id = message.message_id(); + let transport = message.transport_name(); + + match block.call::<_, Value>([message.into_value()]) { + Ok(result) => { + // A response is an answer to a request, not a verdict on a syslog + // message, and must not be read as one. + if Obj::::try_convert(result).is_ok() + || Obj::::try_convert(result).is_ok() + { + SYSLOG_RESPONSE_WARNING.call_once(|| { + error!("Block returned a response for a syslog message - a syslog message is answered with an admission verdict, so the message is being refused"); + }); + return syslog::HandlerResult::Failed; + } + + if result.to_bool() { + syslog::HandlerResult::Accepted + } else { + syslog::HandlerResult::Refused + } + } + Err(e) => { + error!( + "Block call failed with error: {:?} - refusing {} syslog message {}", + e, transport, message_id + ); + syslog::HandlerResult::Failed + } + } +} + async fn handle_request( req: HyperRequest, work_tx: Arc>, @@ -694,6 +922,8 @@ fn init(ruby: &Ruby) -> Result<(), MagnusError> { server_class.define_method("stop", method!(Server::stop, 0))?; server_class.define_method("run_worker", method!(Server::run_worker, 0))?; server_class.define_method("total_connections", method!(Server::total_connections, 0))?; + server_class.define_method("syslog_listening?", method!(Server::syslog_listening, 0))?; + server_class.define_method("syslog_stats", method!(Server::syslog_stats, 0))?; let response_class = module.define_class("Response", ruby.class_object())?; response_class.define_singleton_method("new", function!(Response::new, 3))?; @@ -708,6 +938,15 @@ fn init(ruby: &Ruby) -> Result<(), MagnusError> { grpc_response_class.define_method("headers", method!(GrpcResponse::headers, 0))?; grpc_response_class.define_method("body", method!(GrpcResponse::body, 0))?; + let syslog_message_class = module.define_class("SyslogMessage", ruby.class_object())?; + syslog_message_class.define_method("message", method!(syslog::SyslogMessage::message, 0))?; + syslog_message_class.define_method("peer_ip", method!(syslog::SyslogMessage::peer_ip, 0))?; + syslog_message_class.define_method("transport", method!(syslog::SyslogMessage::transport, 0))?; + syslog_message_class.define_method("received_at_ns", method!(syslog::SyslogMessage::received_at_ns, 0))?; + syslog_message_class.define_method("message_id", method!(syslog::SyslogMessage::message_id, 0))?; + syslog_message_class.define_method("attempt", method!(syslog::SyslogMessage::attempt, 0))?; + syslog_message_class.define_method("inspect", method!(syslog::SyslogMessage::inspect, 0))?; + let request_class = module.define_class("Request", ruby.class_object())?; request_class.define_method("http_method", method!(Request::method, 0))?; request_class.define_method("path", method!(Request::path, 0))?; diff --git a/ext/hyper_ruby/src/syslog.rs b/ext/hyper_ruby/src/syslog.rs new file mode 100644 index 0000000..0da60c9 --- /dev/null +++ b/ext/hyper_ruby/src/syslog.rs @@ -0,0 +1,959 @@ +// Syslog transport listeners: a stream listener (Unix socket or TCP) that frames +// messages with RFC 6587 octet counting and a newline fallback, and a datagram +// listener where one datagram is one message. Complete messages are handed to +// the same Ruby worker threads the HTTP path uses. + +mod framing; +mod proxy; + +use std::io; +use std::net::SocketAddr; +use std::os::unix::fs::MetadataExt; +use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; +use std::sync::Arc; +use std::time::{Duration, SystemTime, UNIX_EPOCH}; + +use bytes::{Bytes, BytesMut}; +use crossbeam_channel::TrySendError; +use log::{debug, info, warn}; +use magnus::{ + value::{qnil, ReprValue}, + DataTypeFunctions, Error as MagnusError, RHash, RString, Symbol, TryConvert, TypedData, Value, +}; +use socket2::{Domain, Protocol, Socket, Type}; +use tokio::io::{AsyncRead, AsyncReadExt}; +use tokio::net::{TcpListener, UdpSocket, UnixListener}; +use tokio::sync::{mpsc, oneshot, OwnedSemaphorePermit, Semaphore}; +use tokio::time::{sleep, timeout}; +use tokio_util::sync::CancellationToken; +use tokio_util::task::TaskTracker; + +use framing::{Framer, RejectReason}; +use proxy::{ProxyError, ProxyHeader}; + +/// Largest PROXY v2 header (including TLVs) we will read. +const MAX_PROXY_HEADER_BYTES: usize = 1024; +/// How long to wait for space on the Ruby worker queue before trying again. +const WORKER_QUEUE_RETRY: Duration = Duration::from_millis(10); +/// Backoff bounds used when a handler refuses a message. +const MIN_REFUSAL_DELAY: Duration = Duration::from_millis(5); +const MAX_REFUSAL_DELAY: Duration = Duration::from_millis(500); +/// How long to pause after a failed accept or receive. +const SOCKET_ERROR_DELAY: Duration = Duration::from_millis(10); +/// Read buffer growth increment for stream connections. +const STREAM_READ_CHUNK: usize = 8 * 1024; + +/// Message identities are unique for the life of the process, so a handler can +/// recognise the retries of a message it has already seen. +static NEXT_MESSAGE_ID: AtomicU64 = AtomicU64::new(1); + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum Transport { + Stream, + Datagram, +} + +impl Transport { + pub(crate) fn name(self) -> &'static str { + match self { + Transport::Stream => "stream", + Transport::Datagram => "datagram", + } + } + + pub(crate) fn symbol(self) -> Symbol { + Symbol::new(self.name()) + } +} + +/// What a Ruby handler made of one message. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum HandlerResult { + Accepted, + Refused, + Failed, +} + +#[derive(Clone)] +pub(crate) struct SyslogConfig { + pub stream_path: Option, + pub stream_bind: Option, + pub udp_bind: Option, + pub proxy_protocol: bool, + pub proxy_header_timeout: u64, + pub max_frame_bytes: usize, + pub max_pending: usize, + pub max_pending_per_connection: usize, + pub max_connections: u64, + pub idle_timeout: u64, + pub udp_max_datagram_bytes: usize, + pub udp_so_rcvbuf: Option, + pub work_ratio: u32, +} + +impl SyslogConfig { + pub(crate) fn new() -> Self { + Self { + stream_path: None, + stream_bind: None, + udp_bind: None, + proxy_protocol: false, + proxy_header_timeout: 5000, + max_frame_bytes: 102400, + max_pending: 1000, + max_pending_per_connection: 64, + max_connections: 10000, + idle_timeout: 0, + udp_max_datagram_bytes: 65536, + udp_so_rcvbuf: None, + work_ratio: 4, + } + } + + pub(crate) fn enabled(&self) -> bool { + self.stream_path.is_some() || self.stream_bind.is_some() || self.udp_bind.is_some() + } + + /// Read the syslog keys out of the server configuration hash. + pub(crate) fn apply(&mut self, config: &RHash) -> Result<(), MagnusError> { + if let Some(value) = config.get(Symbol::new("syslog_stream_path")) { + self.stream_path = Some(String::try_convert(value)?); + } + if let Some(value) = config.get(Symbol::new("syslog_stream_bind")) { + self.stream_bind = Some(String::try_convert(value)?); + } + if let Some(value) = config.get(Symbol::new("syslog_udp_bind")) { + self.udp_bind = Some(String::try_convert(value)?); + } + if let Some(value) = config.get(Symbol::new("syslog_proxy_protocol")) { + self.proxy_protocol = bool::try_convert(value)?; + } + if let Some(value) = config.get(Symbol::new("syslog_proxy_header_timeout")) { + self.proxy_header_timeout = u64::try_convert(value)?; + } + if let Some(value) = config.get(Symbol::new("syslog_max_frame_bytes")) { + self.max_frame_bytes = usize::try_convert(value)?; + } + if let Some(value) = config.get(Symbol::new("syslog_max_pending")) { + self.max_pending = usize::try_convert(value)?.clamp(1, Semaphore::MAX_PERMITS); + } + if let Some(value) = config.get(Symbol::new("syslog_max_pending_per_connection")) { + self.max_pending_per_connection = usize::try_convert(value)?.max(1); + } + if let Some(value) = config.get(Symbol::new("syslog_max_connections")) { + self.max_connections = u64::try_convert(value)?; + } + if let Some(value) = config.get(Symbol::new("syslog_idle_timeout_ms")) { + self.idle_timeout = u64::try_convert(value)?; + } + if let Some(value) = config.get(Symbol::new("syslog_udp_max_datagram_bytes")) { + self.udp_max_datagram_bytes = usize::try_convert(value)?.max(1); + } + if let Some(value) = config.get(Symbol::new("syslog_udp_so_rcvbuf")) { + self.udp_so_rcvbuf = Some(usize::try_convert(value)?); + } + if let Some(value) = config.get(Symbol::new("syslog_work_ratio")) { + self.work_ratio = u32::try_convert(value)?; + } + Ok(()) + } +} + +#[derive(Default)] +pub(crate) struct SyslogCounters { + messages_delivered: AtomicU64, + deliveries_refused: AtomicU64, + handler_errors: AtomicU64, + rejected_oversize: AtomicU64, + rejected_invalid_utf8: AtomicU64, + rejected_invalid_length: AtomicU64, + udp_truncated: AtomicU64, + udp_dropped: AtomicU64, + connections_opened: AtomicU64, + connections_closed: AtomicU64, + connections_refused: AtomicU64, + proxy_header_errors: AtomicU64, + proxy_read_errors: AtomicU64, + abandoned_at_shutdown: AtomicU64, + /// Messages framed but not yet through a handler. + pending: AtomicU64, +} + +impl SyslogCounters { + fn record_rejection(&self, reason: RejectReason) { + let counter = match reason { + RejectReason::Oversize => &self.rejected_oversize, + RejectReason::InvalidUtf8 => &self.rejected_invalid_utf8, + RejectReason::InvalidLength => &self.rejected_invalid_length, + }; + counter.fetch_add(1, Ordering::Relaxed); + } + + fn abandon_pending(&self) { + self.abandoned_at_shutdown.fetch_add(1, Ordering::Relaxed); + self.release_pending(); + } + + // Saturating, because a drain timeout clears the gauge while its messages + // may still be finishing. + fn release_pending(&self) { + let _ = self.pending.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |pending| { + Some(pending.saturating_sub(1)) + }); + } + + pub(crate) fn to_hash(&self) -> Result { + let rejected = RHash::new(); + rejected.aset( + Symbol::new(RejectReason::Oversize.as_str()), + self.rejected_oversize.load(Ordering::Relaxed), + )?; + rejected.aset( + Symbol::new(RejectReason::InvalidUtf8.as_str()), + self.rejected_invalid_utf8.load(Ordering::Relaxed), + )?; + rejected.aset( + Symbol::new(RejectReason::InvalidLength.as_str()), + self.rejected_invalid_length.load(Ordering::Relaxed), + )?; + + let stats = RHash::new(); + let counters: [(&str, &AtomicU64); 12] = [ + ("messages_delivered", &self.messages_delivered), + ("deliveries_refused", &self.deliveries_refused), + ("handler_errors", &self.handler_errors), + ("udp_truncated", &self.udp_truncated), + ("udp_dropped", &self.udp_dropped), + ("connections_opened", &self.connections_opened), + ("connections_closed", &self.connections_closed), + ("connections_refused", &self.connections_refused), + ("proxy_header_errors", &self.proxy_header_errors), + ("proxy_read_errors", &self.proxy_read_errors), + ("abandoned_at_shutdown", &self.abandoned_at_shutdown), + ("pending", &self.pending), + ]; + for (name, counter) in counters { + stats.aset(Symbol::new(name), counter.load(Ordering::Relaxed))?; + } + stats.aset(Symbol::new("frames_rejected"), rejected)?; + Ok(stats) + } +} + +/// A complete message on its way to a Ruby worker thread. +pub(crate) struct SyslogDelivery { + message: SyslogMessage, + result_tx: oneshot::Sender, +} + +impl SyslogDelivery { + /// Split into the object handed to the worker block and the channel its + /// verdict goes back on. + pub(crate) fn into_parts(self) -> (SyslogMessage, oneshot::Sender) { + (self.message, self.result_tx) + } +} + +/// One syslog message, as the worker block sees it. Freed as soon as it is +/// collected, and its payload is reported to the GC, as for a request. +#[derive(TypedData)] +#[magnus(class = "HyperRuby::SyslogMessage", free_immediately, size)] +pub(crate) struct SyslogMessage { + bytes: Bytes, + /// Stream frames are validated UTF-8; datagrams stay binary. + utf8: bool, + /// Absent when the transport names no peer, such as a Unix socket + /// connection without a PROXY header. + peer: Option>, + transport: Transport, + received_at_ns: u64, + message_id: u64, + attempt: u32, +} + +impl DataTypeFunctions for SyslogMessage { + fn size(&self) -> usize { + std::mem::size_of_val(self) + self.bytes.len() + } +} + +impl SyslogMessage { + pub(crate) fn message(&self) -> RString { + if self.utf8 { + // Stream frames are validated while they are framed, so this needs + // no second pass over the bytes. + unsafe { RString::new(std::str::from_utf8_unchecked(&self.bytes)) } + } else { + RString::from_slice(&self.bytes) + } + } + + pub(crate) fn transport_name(&self) -> &'static str { + self.transport.name() + } + + pub(crate) fn peer_ip(&self) -> Value { + match &self.peer { + Some(peer) => RString::new(peer).as_value(), + None => qnil().as_value(), + } + } + + pub(crate) fn transport(&self) -> Symbol { + self.transport.symbol() + } + + pub(crate) fn received_at_ns(&self) -> u64 { + self.received_at_ns + } + + pub(crate) fn message_id(&self) -> u64 { + self.message_id + } + + pub(crate) fn attempt(&self) -> u32 { + self.attempt + } + + pub(crate) fn inspect(&self) -> RString { + RString::new(&format!( + "#", + self.transport.name(), + self.bytes.len(), + self.message_id, + self.attempt + )) + } +} + +enum DeliveryOutcome { + Accepted, + Refused, + /// No worker can answer any more; the listener should give up. + Unavailable, +} + +/// Everything a listener needs to hand a message to the Ruby workers. +#[derive(Clone)] +struct Dispatcher { + work_tx: Arc>, + counters: Arc, + credits: Arc, + connections: Arc, + tracker: TaskTracker, +} + +impl Dispatcher { + async fn deliver(&self, pending: &PendingMessage, attempt: u32) -> DeliveryOutcome { + let (result_tx, result_rx) = oneshot::channel(); + let mut delivery = SyslogDelivery { + message: SyslogMessage { + bytes: pending.message.clone(), + utf8: pending.utf8, + peer: pending.peer.clone(), + transport: pending.transport, + received_at_ns: pending.received_at_ns, + message_id: pending.message_id, + attempt, + }, + result_tx, + }; + + loop { + match self.work_tx.try_send(delivery) { + Ok(()) => break, + Err(TrySendError::Full(returned)) => { + delivery = returned; + sleep(WORKER_QUEUE_RETRY).await; + } + Err(TrySendError::Disconnected(_)) => return DeliveryOutcome::Unavailable, + } + } + + match result_rx.await { + Ok(HandlerResult::Accepted) => { + self.counters + .messages_delivered + .fetch_add(1, Ordering::Relaxed); + DeliveryOutcome::Accepted + } + Ok(HandlerResult::Refused) => { + self.counters + .deliveries_refused + .fetch_add(1, Ordering::Relaxed); + DeliveryOutcome::Refused + } + Ok(HandlerResult::Failed) => { + self.counters.handler_errors.fetch_add(1, Ordering::Relaxed); + DeliveryOutcome::Refused + } + Err(_) => DeliveryOutcome::Unavailable, + } + } +} + +/// Handle on the running listeners, used to drain them at shutdown. +pub(crate) struct SyslogRuntime { + tracker: TaskTracker, + token: CancellationToken, + counters: Arc, + listening: Arc, + socket_path: Option, + // (dev, ino) of the Unix socket file we bound, so we only unlink our own. + socket_ident: Option<(u64, u64)>, +} + +impl SyslogRuntime { + pub(crate) fn listening(&self) -> bool { + self.listening.load(Ordering::Relaxed) + } + + /// Stop accepting connections and reading from the open ones. + pub(crate) fn stop_accepting(&self) { + self.listening.store(false, Ordering::Relaxed); + self.token.cancel(); + } + + /// Wait for the listeners to finish delivering messages they have framed. + pub(crate) async fn drain(&self, limit: Duration) { + self.stop_accepting(); + if timeout(limit, self.tracker.wait()).await.is_err() { + let abandoned = self.counters.pending.swap(0, Ordering::Relaxed); + self.counters + .abandoned_at_shutdown + .fetch_add(abandoned, Ordering::Relaxed); + warn!( + "Timed out draining syslog listeners, abandoning {} undelivered messages", + abandoned + ); + } + } + + pub(crate) fn remove_socket_file(&self) { + remove_socket_file(self.socket_path.as_deref(), self.socket_ident); + } +} + +/// Unlink a bound Unix socket file, unless a replacement server has taken the +/// path over in the meantime. +fn remove_socket_file(path: Option<&str>, ident: Option<(u64, u64)>) { + let (Some(path), Some((dev, ino))) = (path, ident) else { + return; + }; + match std::fs::symlink_metadata(path) { + Ok(meta) if (meta.dev(), meta.ino()) == (dev, ino) => { + std::fs::remove_file(path) + .unwrap_or_else(|e| warn!("Failed to remove syslog socket file: {:?}", e)); + } + Ok(_) => info!( + "Syslog socket file {} was replaced by another server; leaving it in place", + path + ), + Err(_) => debug!("Syslog socket file {} already removed", path), + } +} + +/// Bind the configured listeners and start serving. Must be called from within +/// the Tokio runtime. +pub(crate) fn start( + config: &SyslogConfig, + counters: Arc, + work_tx: Arc>, +) -> io::Result { + let mut socket = (None, None); + let started = start_listeners(config, counters, work_tx, &mut socket); + if started.is_err() { + // Leave no socket file behind for listeners that never came up. + remove_socket_file(socket.0.as_deref(), socket.1); + } + started +} + +fn start_listeners( + config: &SyslogConfig, + counters: Arc, + work_tx: Arc>, + socket: &mut (Option, Option<(u64, u64)>), +) -> io::Result { + let tracker = TaskTracker::new(); + let token = CancellationToken::new(); + let dispatcher = Dispatcher { + work_tx, + counters: counters.clone(), + credits: Arc::new(Semaphore::new(config.max_pending)), + connections: Arc::new(AtomicU64::new(0)), + tracker: tracker.clone(), + }; + + if let Some(path) = &config.stream_path { + let (listener, ident) = bind_unix_listener(path)?; + *socket = (Some(path.clone()), ident); + spawn_stream_listener( + StreamListener::Unix(listener), + config.clone(), + dispatcher.clone(), + &token, + &tracker, + ); + } + + if let Some(address) = &config.stream_bind { + let address: SocketAddr = address + .parse() + .map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, format!("{}", e)))?; + let listener = TcpListener::from_std(bind_std_tcp(address)?)?; + spawn_stream_listener( + StreamListener::Tcp(listener), + config.clone(), + dispatcher.clone(), + &token, + &tracker, + ); + } + + if let Some(address) = &config.udp_bind { + let address: SocketAddr = address + .parse() + .map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, format!("{}", e)))?; + let datagrams = UdpSocket::from_std(bind_std_udp(address, config)?)?; + tracker.spawn(serve_datagrams( + datagrams, + config.clone(), + dispatcher.clone(), + token.child_token(), + )); + } + + tracker.close(); + + Ok(SyslogRuntime { + tracker, + token, + counters, + listening: Arc::new(AtomicBool::new(true)), + socket_path: socket.0.clone(), + socket_ident: socket.1, + }) +} + +fn bind_unix_listener(path: &str) -> io::Result<(UnixListener, Option<(u64, u64)>)> { + // Bind a unique temporary path and rename over the target, so the path + // always points at a live socket even when a replacement server takes over + // a path an older, still draining server bound. + static SOCKET_TMP_SEQ: AtomicU64 = AtomicU64::new(0); + let tmp_path = format!( + "{}.{}.{}.tmp", + path, + std::process::id(), + SOCKET_TMP_SEQ.fetch_add(1, Ordering::Relaxed) + ); + + let listener = UnixListener::bind(&tmp_path)?; + let ident = std::fs::symlink_metadata(&tmp_path) + .ok() + .map(|meta| (meta.dev(), meta.ino())); + + if let Err(e) = std::fs::rename(&tmp_path, path) { + let _ = std::fs::remove_file(&tmp_path); + return Err(e); + } + + Ok((listener, ident)) +} + +fn bind_std_tcp(address: SocketAddr) -> io::Result { + let socket = Socket::new(Domain::for_address(address), Type::STREAM, Some(Protocol::TCP))?; + socket.set_reuse_address(true)?; + socket.set_nonblocking(true)?; + socket.bind(&address.into())?; + socket.listen(1024)?; + Ok(socket.into()) +} + +fn bind_std_udp(address: SocketAddr, config: &SyslogConfig) -> io::Result { + let socket = Socket::new(Domain::for_address(address), Type::DGRAM, Some(Protocol::UDP))?; + // Several processes can bind the same port; the kernel hands each datagram + // to one of them. + socket.set_reuse_address(true)?; + socket.set_reuse_port(true)?; + if let Some(size) = config.udp_so_rcvbuf { + socket.set_recv_buffer_size(size)?; + } + socket.set_nonblocking(true)?; + socket.bind(&address.into())?; + Ok(socket.into()) +} + +enum StreamListener { + Unix(UnixListener), + Tcp(TcpListener), +} + +impl StreamListener { + async fn accept(&self) -> io::Result<(Box, Option>)> { + match self { + StreamListener::Unix(listener) => { + let (stream, _) = listener.accept().await?; + Ok((Box::new(stream), None)) + } + StreamListener::Tcp(listener) => { + let (stream, address) = listener.accept().await?; + Ok((Box::new(stream), Some(Arc::from(address.ip().to_string())))) + } + } + } +} + +trait AsyncReadStream: AsyncRead + Unpin + Send {} +impl AsyncReadStream for T {} + +fn spawn_stream_listener( + listener: StreamListener, + config: SyslogConfig, + dispatcher: Dispatcher, + token: &CancellationToken, + tracker: &TaskTracker, +) { + let token = token.clone(); + let connections = tracker.clone(); + + tracker.spawn(async move { + loop { + tokio::select! { + accepted = listener.accept() => { + let (stream, peer) = match accepted { + Ok(accepted) => accepted, + Err(e) => { + // Transient accept failures (descriptor limits) must + // not turn into a hot loop. + warn!("Failed to accept syslog connection: {:?}", e); + sleep(SOCKET_ERROR_DELAY).await; + continue; + } + }; + + if dispatcher.connections.fetch_add(1, Ordering::AcqRel) >= config.max_connections { + dispatcher.connections.fetch_sub(1, Ordering::AcqRel); + dispatcher.counters.connections_refused.fetch_add(1, Ordering::Relaxed); + debug!("Refusing syslog connection: connection limit reached"); + continue; + } + + dispatcher.counters.connections_opened.fetch_add(1, Ordering::Relaxed); + // A child token, so a connection accepted in the same round + // as the shutdown still sees it. + connections.spawn(serve_stream( + stream, + peer, + config.clone(), + dispatcher.clone(), + token.child_token(), + )); + }, + _ = token.cancelled() => { + debug!("Syslog stream listener shutting down"); + break; + } + } + } + }); +} + +/// A framed message waiting for a Ruby worker. The permit is held until +/// delivery finishes, which is what bounds undelivered messages overall. +struct PendingMessage { + message: Bytes, + utf8: bool, + peer: Option>, + transport: Transport, + received_at_ns: u64, + message_id: u64, + _permit: OwnedSemaphorePermit, +} + +async fn serve_stream( + mut stream: Box, + peer: Option>, + config: SyslogConfig, + dispatcher: Dispatcher, + token: CancellationToken, +) { + let mut buffer = BytesMut::with_capacity(STREAM_READ_CHUNK); + + let peer = if config.proxy_protocol { + match read_proxy_header(&mut stream, &mut buffer, &config).await { + Ok(Some(source)) => Some(Arc::from(source.to_string())), + Ok(None) => peer, + Err(failure) => { + failure.record(&dispatcher.counters); + debug!( + "Rejecting syslog connection: PROXY header {}", + failure.as_str() + ); + close_stream(&dispatcher); + return; + } + } + } else { + peer + }; + + // Reading and delivery are separate so that a slow handler stops the reads + // through the bounded channel rather than blocking the reactor. + let (pending_tx, pending_rx) = + mpsc::channel::(config.max_pending_per_connection); + let delivery = tokio::spawn(deliver_stream_messages( + pending_rx, + dispatcher.clone(), + token.clone(), + )); + + let mut framer = Framer::new(config.max_frame_bytes); + let idle_timeout = Duration::from_millis(config.idle_timeout); + let mut eof = false; + + 'read: loop { + // Emit everything the buffer already holds before asking for more. + loop { + let decoded = if eof { + framer.decode_eof(&mut buffer) + } else { + framer.decode(&mut buffer) + }; + + match decoded { + Ok(Some(message)) => { + let received_at_ns = now_nanos(); + let Ok(permit) = dispatcher.credits.clone().acquire_owned().await else { + break 'read; + }; + let pending = PendingMessage { + message, + utf8: true, + peer: peer.clone(), + transport: Transport::Stream, + received_at_ns, + message_id: NEXT_MESSAGE_ID.fetch_add(1, Ordering::Relaxed), + _permit: permit, + }; + dispatcher.counters.pending.fetch_add(1, Ordering::Relaxed); + if pending_tx.send(pending).await.is_err() { + dispatcher.counters.abandon_pending(); + break 'read; + } + } + Ok(None) => break, + Err(error) => { + dispatcher.counters.record_rejection(error.reason); + if error.fatal { + debug!("Closing syslog connection: {} frame", error.reason.as_str()); + break 'read; + } + } + } + } + + if eof { + break; + } + + tokio::select! { + read = stream.read_buf(&mut buffer) => { + match read { + Ok(0) => eof = true, + Ok(_) => (), + Err(e) => { + debug!("Syslog connection read failed: {:?}", e); + break; + } + } + }, + _ = sleep(idle_timeout), if config.idle_timeout > 0 => { + debug!("Closing idle syslog connection"); + break; + }, + _ = token.cancelled() => { + debug!("Syslog connection stopping reads for shutdown"); + break; + } + } + } + + // Dropping the sender lets the delivery task finish what it already has. + drop(pending_tx); + let _ = delivery.await; + close_stream(&dispatcher); +} + +fn close_stream(dispatcher: &Dispatcher) { + dispatcher.connections.fetch_sub(1, Ordering::AcqRel); + dispatcher + .counters + .connections_closed + .fetch_add(1, Ordering::Relaxed); +} + +async fn deliver_stream_messages( + mut pending_rx: mpsc::Receiver, + dispatcher: Dispatcher, + token: CancellationToken, +) { + let mut workers_gone = false; + + while let Some(pending) = pending_rx.recv().await { + if workers_gone { + dispatcher.counters.abandon_pending(); + continue; + } + + let mut attempt = 1; + let mut delay = MIN_REFUSAL_DELAY; + + loop { + match dispatcher.deliver(&pending, attempt).await { + DeliveryOutcome::Accepted => { + dispatcher.counters.release_pending(); + break; + } + DeliveryOutcome::Refused => { + // Hold the message and try again; reads stall behind the + // bounded pending channel while we wait. Shutdown ends the + // retries rather than holding the drain open. + if token.is_cancelled() { + dispatcher.counters.abandon_pending(); + break; + } + sleep(delay).await; + delay = (delay * 2).min(MAX_REFUSAL_DELAY); + attempt += 1; + } + DeliveryOutcome::Unavailable => { + dispatcher.counters.abandon_pending(); + workers_gone = true; + break; + } + } + } + } +} + +/// Read and consume a PROXY v2 header, returning the source address it carries. +async fn read_proxy_header( + stream: &mut Box, + buffer: &mut BytesMut, + config: &SyslogConfig, +) -> Result, ProxyFailure> { + let deadline = Duration::from_millis(config.proxy_header_timeout); + timeout(deadline, async { + loop { + match proxy::parse(buffer, MAX_PROXY_HEADER_BYTES) { + Ok(ProxyHeader::Complete { source, length }) => { + let _ = buffer.split_to(length); + return Ok(source); + } + Ok(ProxyHeader::Incomplete) => match stream.read_buf(buffer).await { + Ok(0) => return Err(ProxyFailure::Read("eof")), + Ok(_) => continue, + Err(_) => return Err(ProxyFailure::Read("io")), + }, + Err(error) => return Err(ProxyFailure::Header(error)), + } + } + }) + .await + .unwrap_or(Err(ProxyFailure::Read("timeout"))) +} + +/// Why a connection gave up before its PROXY header was complete. +enum ProxyFailure { + Header(ProxyError), + Read(&'static str), +} + +impl ProxyFailure { + fn as_str(&self) -> &'static str { + match self { + ProxyFailure::Header(error) => error.as_str(), + ProxyFailure::Read(reason) => reason, + } + } + + fn record(&self, counters: &SyslogCounters) { + match self { + ProxyFailure::Header(_) => counters.proxy_header_errors.fetch_add(1, Ordering::Relaxed), + ProxyFailure::Read(_) => counters.proxy_read_errors.fetch_add(1, Ordering::Relaxed), + }; + } +} + +async fn serve_datagrams( + socket: UdpSocket, + config: SyslogConfig, + dispatcher: Dispatcher, + token: CancellationToken, +) { + let mut buffer = vec![0u8; config.udp_max_datagram_bytes]; + + loop { + tokio::select! { + received = socket.recv_from(&mut buffer) => { + let (length, peer) = match received { + Ok(received) => received, + Err(e) => { + warn!("Syslog datagram receive failed: {:?}", e); + sleep(SOCKET_ERROR_DELAY).await; + continue; + } + }; + let received_at_ns = now_nanos(); + + // A datagram that fills the buffer was almost certainly cut + // short by it, and a truncated prefix must not pass as a message. + if length == buffer.len() { + dispatcher.counters.udp_truncated.fetch_add(1, Ordering::Relaxed); + continue; + } + + let Ok(permit) = dispatcher.credits.clone().try_acquire_owned() else { + dispatcher.counters.udp_dropped.fetch_add(1, Ordering::Relaxed); + continue; + }; + + let pending = PendingMessage { + message: Bytes::copy_from_slice(&buffer[..length]), + utf8: false, + peer: Some(Arc::from(peer.ip().to_string())), + transport: Transport::Datagram, + received_at_ns, + message_id: NEXT_MESSAGE_ID.fetch_add(1, Ordering::Relaxed), + _permit: permit, + }; + dispatcher.counters.pending.fetch_add(1, Ordering::Relaxed); + + let dispatcher = dispatcher.clone(); + dispatcher.tracker.clone().spawn(async move { + // A datagram sender cannot be asked to slow down, so a + // refusal drops the message. + match dispatcher.deliver(&pending, 1).await { + DeliveryOutcome::Accepted => (), + DeliveryOutcome::Refused | DeliveryOutcome::Unavailable => { + dispatcher.counters.udp_dropped.fetch_add(1, Ordering::Relaxed); + } + } + dispatcher.counters.release_pending(); + }); + }, + _ = token.cancelled() => { + debug!("Syslog datagram listener shutting down"); + break; + } + } + } +} + +fn now_nanos() -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|elapsed| elapsed.as_nanos() as u64) + .unwrap_or(0) +} diff --git a/ext/hyper_ruby/src/syslog/framing.rs b/ext/hyper_ruby/src/syslog/framing.rs new file mode 100644 index 0000000..b8df0ec --- /dev/null +++ b/ext/hyper_ruby/src/syslog/framing.rs @@ -0,0 +1,508 @@ +// RFC 6587 framing: octet counting (` `) with a non-transparent +// newline fallback. A buffer whose first byte is a non-zero digit is treated as +// octet counted, everything else is newline delimited. +// +// This file (including its test data) is a port of Vector 0.48.0's +// lib/codecs/src/decoding/framing/octet_counting.rs and is licensed under the +// Mozilla Public License 2.0, not the MIT licence covering the rest of this +// gem. See https://github.com/vectordotdev/vector and +// https://www.mozilla.org/en-US/MPL/2.0/. Deliberate divergences from that +// source are marked below. + +use bytes::{Buf, Bytes, BytesMut}; +use tokio_util::codec::{Decoder, LinesCodec, LinesCodecError}; + +/// Why a frame was rejected. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum RejectReason { + Oversize, + InvalidUtf8, + InvalidLength, +} + +impl RejectReason { + pub(crate) fn as_str(self) -> &'static str { + match self { + RejectReason::Oversize => "oversize", + RejectReason::InvalidUtf8 => "invalid_utf8", + RejectReason::InvalidLength => "invalid_length", + } + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) struct FrameError { + pub reason: RejectReason, + /// A fatal error leaves the stream out of sync: the frame boundary can no + /// longer be established, so the connection must close. A non-fatal error + /// skips the offending frame only. + pub fatal: bool, +} + +impl FrameError { + fn fatal(reason: RejectReason) -> Self { + Self { reason, fatal: true } + } + + fn recoverable(reason: RejectReason) -> Self { + Self { reason, fatal: false } + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum State { + NotDiscarding, + Discarding(usize), + DiscardingToEol, +} + +pub(crate) struct Framer { + other: LinesCodec, + octet_decoding: Option, +} + +impl Framer { + pub(crate) fn new(max_length: usize) -> Self { + Self { + other: LinesCodec::new_with_max_length(max_length), + octet_decoding: None, + } + } + + /// Decode the next complete frame, if the buffer holds one. + pub(crate) fn decode(&mut self, src: &mut BytesMut) -> Result, FrameError> { + match self.checked_decode(src) { + Some(result) => result, + None => self + .other + .decode(src) + .map(|line| line.map(Bytes::from)) + .map_err(line_error), + } + } + + /// Decode at end of stream, which releases a trailing unterminated line. + pub(crate) fn decode_eof(&mut self, src: &mut BytesMut) -> Result, FrameError> { + match self.checked_decode(src) { + Some(result) => result, + None => self + .other + .decode_eof(src) + .map(|line| line.map(Bytes::from)) + .map_err(line_error), + } + } + + /// `None` if this buffer is not octet counting encoded. + fn checked_decode(&mut self, src: &mut BytesMut) -> Option, FrameError>> { + // Divergence from the ported source: only a decoder with no state in + // hand starts a new octet counted frame. The source re-enters + // `NotDiscarding` whenever the buffer happens to start with a digit, + // so a read boundary that lands on a digit inside an oversize body + // loses the countdown and hands the rest of that body over as + // messages. + if self.octet_decoding.is_none() { + if let Some(&first_byte) = src.first() { + if (b'1'..=b'9').contains(&first_byte) { + self.octet_decoding = Some(State::NotDiscarding); + } + } + } + + self.octet_decoding + .map(|state| self.octet_decode(state, src)) + } + + fn octet_decode( + &mut self, + state: State, + src: &mut BytesMut, + ) -> Result, FrameError> { + // Encoding scheme: an ASCII decimal length, a space, then that many + // bytes of message. + let space_pos = src.iter().position(|&b| b == b' '); + let newline_pos = src.iter().position(|&b| b == b'\n'); + + match (state, newline_pos, space_pos) { + (State::Discarding(chars), _, _) if src.len() >= chars => { + // Enough bytes buffered to finish discarding the oversize frame. + src.advance(chars); + self.octet_decoding = None; + Err(FrameError::fatal(RejectReason::Oversize)) + } + + (State::Discarding(chars), _, _) => { + // Not enough bytes yet; discard what we have and carry the + // remainder forward. Divergence from the ported source, which + // subtracts these the other way around: with overflow checks + // off, as in its release builds, the count wraps to a huge + // number and the frame is never reported, so this countdown is + // stricter than that source's behaviour. + self.octet_decoding = Some(State::Discarding(chars - src.len())); + src.advance(src.len()); + Ok(None) + } + + (State::DiscardingToEol, Some(offset), _) => { + src.advance(offset + 1); + self.octet_decoding = None; + Err(FrameError::fatal(RejectReason::Oversize)) + } + + (State::DiscardingToEol, None, _) => { + // No newline to sync on yet; discard the whole buffer. + src.advance(src.len()); + Ok(None) + } + + (State::NotDiscarding, _, Some(space_pos)) if space_pos < self.other.max_length() => { + let len: usize = match std::str::from_utf8(&src[..space_pos]) + .map_err(|_| ()) + .and_then(|num| num.parse().map_err(|_| ())) + { + Ok(len) => len, + Err(_) => { + // Not a sensible number; step past it so we cannot loop + // on the same bytes forever. + src.advance(space_pos + 1); + self.octet_decoding = None; + return Err(FrameError::fatal(RejectReason::InvalidLength)); + } + }; + + let from = space_pos + 1; + let to = from + len; + + if len > self.other.max_length() { + // Discard the declared length before reporting the error, so + // the message body cannot be mistaken for further frames. + self.octet_decoding = Some(State::Discarding(len)); + src.advance(space_pos + 1); + Ok(None) + } else if let Some(msg) = src.get(from..to) { + let bytes = match std::str::from_utf8(msg) { + Ok(_) => Bytes::copy_from_slice(msg), + Err(_) => { + src.advance(to); + self.octet_decoding = None; + return Err(FrameError::fatal(RejectReason::InvalidUtf8)); + } + }; + + src.advance(to); + self.octet_decoding = None; + Ok(Some(bytes)) + } else { + // Acceptable length, but the message is not all here yet. + Ok(None) + } + } + + (State::NotDiscarding, Some(newline_pos), _) => { + // Beyond the maximum length; advance to the newline. + src.advance(newline_pos + 1); + Err(FrameError::fatal(RejectReason::Oversize)) + } + + (State::NotDiscarding, None, _) if src.len() < self.other.max_length() => Ok(None), + + (State::NotDiscarding, None, _) => { + // More data than we will handle and nothing to sync on. + self.octet_decoding = Some(State::DiscardingToEol); + src.advance(src.len()); + Ok(None) + } + } + } +} + +fn line_error(error: LinesCodecError) -> FrameError { + match error { + LinesCodecError::MaxLineLengthExceeded => FrameError::recoverable(RejectReason::Oversize), + // The line codec only fails with an IO error for non-UTF-8 input. + LinesCodecError::Io(_) => FrameError::fatal(RejectReason::InvalidUtf8), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use bytes::BufMut; + + fn decode_all(framer: &mut Framer, buffer: &mut BytesMut) -> Vec> { + let mut out = Vec::new(); + loop { + match framer.decode(buffer) { + Ok(Some(frame)) => out.push(Ok(frame)), + Ok(None) => return out, + Err(error) => { + out.push(Err(error)); + if error.fatal { + return out; + } + } + } + } + } + + #[test] + fn non_octet_decode_works_with_multiple_frames() { + let mut decoder = Framer::new(128); + let mut buffer = BytesMut::with_capacity(16); + + buffer.put(&b"<57>Mar 25 21:47:46 host quaerat[2444]: There were "[..]); + assert_eq!(Ok(None), decoder.decode(&mut buffer)); + + buffer.put(&b"8 penguins in the shop.\n"[..]); + assert_eq!( + Ok(Some(Bytes::from( + "<57>Mar 25 21:47:46 host quaerat[2444]: There were 8 penguins in the shop." + ))), + decoder.decode(&mut buffer) + ); + } + + #[test] + fn newline_decode_strips_carriage_return() { + let mut decoder = Framer::new(128); + let mut buffer = BytesMut::from(&b"<13>first\r\n<13>second\n"[..]); + + assert_eq!(Ok(Some(Bytes::from("<13>first"))), decoder.decode(&mut buffer)); + assert_eq!(Ok(Some(Bytes::from("<13>second"))), decoder.decode(&mut buffer)); + assert_eq!(Ok(None), decoder.decode(&mut buffer)); + } + + #[test] + fn octet_decode_works_with_multiple_frames() { + let mut decoder = Framer::new(30); + let mut buffer = BytesMut::with_capacity(16); + + buffer.put(&b"28 abcdefghijklm"[..]); + assert_eq!(Ok(None), decoder.decode(&mut buffer)); + + // A frame starting with a number mid-message must not start a new message. + buffer.put(&b"3 nopqrstuvwxyz"[..]); + assert_eq!( + Ok(Some(Bytes::from("abcdefghijklm3 nopqrstuvwxyz"))), + decoder.decode(&mut buffer) + ); + } + + #[test] + fn octet_decode_moves_past_invalid_length() { + let mut decoder = Framer::new(16); + let mut buffer = BytesMut::from(&b"232>1 zork"[..]); + + assert_eq!( + Err(FrameError::fatal(RejectReason::InvalidLength)), + decoder.decode(&mut buffer) + ); + assert_eq!(b"zork"[..], buffer); + } + + #[test] + fn octet_decode_moves_past_invalid_utf8() { + let mut decoder = Framer::new(16); + let mut buffer = BytesMut::from(&[b'4', b' ', 0xf0, 0x28, 0x8c, 0xbc][..]); + + assert_eq!( + Err(FrameError::fatal(RejectReason::InvalidUtf8)), + decoder.decode(&mut buffer) + ); + assert_eq!(b""[..], buffer); + } + + #[test] + fn octet_decode_moves_past_exceeded_frame_length() { + let mut decoder = Framer::new(16); + let mut buffer = BytesMut::from( + &b"32thisshouldbelongerthanthmaxframeasizewhichmeansitwillnotbedecoded\n"[..], + ); + + assert_eq!( + Err(FrameError::fatal(RejectReason::Oversize)), + decoder.decode(&mut buffer) + ); + assert_eq!(b""[..], buffer); + } + + #[test] + fn octet_decode_rejects_exceeded_frame_length() { + let mut decoder = Framer::new(16); + let mut buffer = BytesMut::from(&b"26 abcdefghijklmnopqrstuvwxyzand here we are"[..]); + + assert_eq!(Ok(None), decoder.decode(&mut buffer)); + assert_eq!( + Err(FrameError::fatal(RejectReason::Oversize)), + decoder.decode(&mut buffer) + ); + assert_eq!(b"and here we are"[..], buffer); + } + + #[test] + fn octet_decode_rejects_exceeded_frame_length_multiple_frames() { + let mut decoder = Framer::new(16); + let mut buffer = BytesMut::from(&b"26 abc"[..]); + let _ = decoder.decode(&mut buffer); + + buffer.put(&b"defghijklmnopqrstuvwxyzand here we are"[..]); + assert_eq!( + Err(FrameError::fatal(RejectReason::Oversize)), + decoder.decode(&mut buffer) + ); + assert_eq!(b"and here we are"[..], buffer); + } + + #[test] + fn octet_decode_discards_oversize_body_arriving_in_pieces() { + let mut decoder = Framer::new(16); + let mut buffer = BytesMut::from(&b"26 abc"[..]); + assert_eq!(Ok(None), decoder.decode(&mut buffer)); + + for _ in 0..10 { + buffer.put(&b"ab"[..]); + assert_eq!(Ok(None), decoder.decode(&mut buffer)); + } + + buffer.put(&b"xyzrest"[..]); + assert_eq!( + Err(FrameError::fatal(RejectReason::Oversize)), + decoder.decode(&mut buffer) + ); + assert_eq!(b"rest"[..], buffer); + } + + #[test] + fn octet_decode_moves_past_exceeded_frame_length_multiple_frames() { + let mut decoder = Framer::new(16); + let mut buffer = BytesMut::from( + &b"32thisshouldbelongerthanthmaxframeasizewhichmeansitwillnotbedecoded"[..], + ); + let _ = decoder.decode(&mut buffer); + + buffer.put(&b"wemustcontinuetodiscard\n32 something valid"[..]); + assert_eq!( + Err(FrameError::fatal(RejectReason::Oversize)), + decoder.decode(&mut buffer) + ); + assert_eq!(b"32 something valid"[..], buffer); + } + + #[test] + fn oversize_body_starting_with_a_digit_is_not_decoded_as_a_frame() { + let mut decoder = Framer::new(16); + let mut buffer = BytesMut::from(&b"26 "[..]); + assert_eq!(Ok(None), decoder.decode(&mut buffer)); + + // A read boundary leaving a digit at the front of the discarded body. + buffer.put(&b"9 abcdefghZ"[..]); + assert_eq!(Ok(None), decoder.decode(&mut buffer)); + assert_eq!(b""[..], buffer); + + buffer.put(&b"aaaaaaaaaaaaaaa5 next"[..]); + assert_eq!( + Err(FrameError::fatal(RejectReason::Oversize)), + decoder.decode(&mut buffer) + ); + assert_eq!(b"5 next"[..], buffer); + } + + #[test] + fn newline_oversize_is_recoverable() { + let mut decoder = Framer::new(16); + let mut buffer = BytesMut::from(&b"<13>aaaaaaaaaaaaaaaaaaaaaaaaaaaa\n<13>ok\n"[..]); + + let results = decode_all(&mut decoder, &mut buffer); + assert_eq!( + vec![ + Err(FrameError::recoverable(RejectReason::Oversize)), + Ok(Bytes::from("<13>ok")), + ], + results + ); + } + + #[test] + fn newline_invalid_utf8_is_fatal() { + let mut decoder = Framer::new(64); + let mut buffer = BytesMut::from(&[b'<', b'1', b'3', b'>', 0xf0, 0x28, b'\n'][..]); + + assert_eq!( + Err(FrameError::fatal(RejectReason::InvalidUtf8)), + decoder.decode(&mut buffer) + ); + } + + #[test] + fn decode_eof_releases_trailing_line() { + let mut decoder = Framer::new(64); + let mut buffer = BytesMut::from(&b"<13>trailing"[..]); + + assert_eq!(Ok(None), decoder.decode(&mut buffer)); + assert_eq!( + Ok(Some(Bytes::from("<13>trailing"))), + decoder.decode_eof(&mut buffer) + ); + assert_eq!(Ok(None), decoder.decode_eof(&mut buffer)); + } + + #[test] + fn decode_eof_holds_back_incomplete_octet_frame() { + let mut decoder = Framer::new(64); + let mut buffer = BytesMut::from(&b"10 partial"[..]); + + assert_eq!(Ok(None), decoder.decode(&mut buffer)); + assert_eq!(Ok(None), decoder.decode_eof(&mut buffer)); + } + + #[test] + fn frames_decode_identically_at_every_split_point() { + let inputs: Vec<&[u8]> = vec![ + b"11 hello there7 goodbye<13>newline\n<13>another\r\n", + b"5 first<13>second\n6 third\n", + b"<13>only a line\n", + b"13 embedded\nline<13>after\n", + ]; + + for input in inputs { + let mut whole = BytesMut::from(input); + let expected = decode_all(&mut Framer::new(64), &mut whole); + + for split in 1..input.len() { + let mut framer = Framer::new(64); + let mut buffer = BytesMut::new(); + let mut actual = Vec::new(); + + buffer.put(&input[..split]); + actual.extend(decode_all(&mut framer, &mut buffer)); + buffer.put(&input[split..]); + actual.extend(decode_all(&mut framer, &mut buffer)); + + assert_eq!( + expected, + actual, + "split {} of {:?}", + split, + String::from_utf8_lossy(input) + ); + } + } + } + + #[test] + fn frames_decode_identically_byte_at_a_time() { + let input: &[u8] = b"11 hello there<13>newline\n7 goodbye"; + let mut whole = BytesMut::from(input); + let expected = decode_all(&mut Framer::new(64), &mut whole); + + let mut framer = Framer::new(64); + let mut buffer = BytesMut::new(); + let mut actual = Vec::new(); + for byte in input { + buffer.put_u8(*byte); + actual.extend(decode_all(&mut framer, &mut buffer)); + } + + assert_eq!(expected, actual); + } +} diff --git a/ext/hyper_ruby/src/syslog/proxy.rs b/ext/hyper_ruby/src/syslog/proxy.rs new file mode 100644 index 0000000..6f333cb --- /dev/null +++ b/ext/hyper_ruby/src/syslog/proxy.rs @@ -0,0 +1,286 @@ +// PROXY protocol v2 header parsing. Only the binary v2 header is accepted; the +// v1 text header and a missing header are errors. + +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr}; + +const SIGNATURE: [u8; 12] = [ + 0x0d, 0x0a, 0x0d, 0x0a, 0x00, 0x0d, 0x0a, 0x51, 0x55, 0x49, 0x54, 0x0a, +]; +const V1_SIGNATURE: &[u8] = b"PROXY "; +/// Signature, version/command, family and the length of the address block. +const PREFIX_LEN: usize = 16; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum ProxyError { + Signature, + Version, + Command, + Family, + TooLong, +} + +impl ProxyError { + pub(crate) fn as_str(self) -> &'static str { + match self { + ProxyError::Signature => "signature", + ProxyError::Version => "version", + ProxyError::Command => "command", + ProxyError::Family => "family", + ProxyError::TooLong => "too_long", + } + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum ProxyHeader { + /// A valid prefix so far, but more bytes are needed. + Incomplete, + /// A complete header; `source` is absent for LOCAL and unspecified families. + Complete { + source: Option, + length: usize, + }, +} + +/// Parse a PROXY v2 header from the front of `buf`, rejecting anything longer +/// than `max_length` bytes. +pub(crate) fn parse(buf: &[u8], max_length: usize) -> Result { + if buf.is_empty() { + return Ok(ProxyHeader::Incomplete); + } + + let seen = buf.len().min(V1_SIGNATURE.len()); + if buf[..seen] == V1_SIGNATURE[..seen] { + return Err(ProxyError::Version); + } + + let seen = buf.len().min(SIGNATURE.len()); + if buf[..seen] != SIGNATURE[..seen] { + return Err(ProxyError::Signature); + } + + if buf.len() < PREFIX_LEN { + return Ok(ProxyHeader::Incomplete); + } + + let version_command = buf[12]; + if version_command >> 4 != 0x2 { + return Err(ProxyError::Version); + } + + let command = version_command & 0x0f; + if command > 0x1 { + return Err(ProxyError::Command); + } + + let family = buf[13]; + let address_length = u16::from_be_bytes([buf[14], buf[15]]) as usize; + let length = PREFIX_LEN + address_length; + if length > max_length { + return Err(ProxyError::TooLong); + } + + if buf.len() < length { + return Ok(ProxyHeader::Incomplete); + } + + // LOCAL connections (health checks) carry no meaningful address. + let source = if command == 0x0 { + None + } else { + match family { + // AF_INET, stream or datagram. + 0x11 | 0x12 => { + if address_length < 12 { + return Err(ProxyError::Family); + } + let mut octets = [0u8; 4]; + octets.copy_from_slice(&buf[16..20]); + Some(IpAddr::V4(Ipv4Addr::from(octets))) + } + // AF_INET6, stream or datagram. + 0x21 | 0x22 => { + if address_length < 36 { + return Err(ProxyError::Family); + } + let mut octets = [0u8; 16]; + octets.copy_from_slice(&buf[16..32]); + Some(IpAddr::V6(Ipv6Addr::from(octets))) + } + // Unspecified and Unix families name no network peer; the + // transport's own peer stands. + 0x00 | 0x31 | 0x32 => None, + _ => return Err(ProxyError::Family), + } + }; + + Ok(ProxyHeader::Complete { source, length }) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn v2_header(family: u8, command: u8, address: &[u8]) -> Vec { + let mut header = SIGNATURE.to_vec(); + header.push(0x20 | command); + header.push(family); + header.extend_from_slice(&(address.len() as u16).to_be_bytes()); + header.extend_from_slice(address); + header + } + + fn ipv4_address(source: [u8; 4]) -> Vec { + let mut address = source.to_vec(); + address.extend_from_slice(&[10, 0, 0, 1]); // destination + address.extend_from_slice(&514u16.to_be_bytes()); + address.extend_from_slice(&6514u16.to_be_bytes()); + address + } + + #[test] + fn parses_ipv4_source() { + let header = v2_header(0x11, 0x1, &ipv4_address([192, 0, 2, 7])); + assert_eq!( + Ok(ProxyHeader::Complete { + source: Some(IpAddr::V4(Ipv4Addr::new(192, 0, 2, 7))), + length: header.len(), + }), + parse(&header, 1024) + ); + } + + #[test] + fn parses_ipv6_source() { + let mut address = vec![0u8; 36]; + address[15] = 1; + let header = v2_header(0x21, 0x1, &address); + assert_eq!( + Ok(ProxyHeader::Complete { + source: Some(IpAddr::V6(Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 1))), + length: header.len(), + }), + parse(&header, 1024) + ); + } + + #[test] + fn ignores_trailing_payload_and_tlvs() { + let mut address = ipv4_address([198, 51, 100, 3]); + address.extend_from_slice(&[0x03, 0x00, 0x04, 1, 2, 3, 4]); // a TLV + let header = v2_header(0x11, 0x1, &address); + let prefix_len = header.len(); + + let mut buf = header; + buf.extend_from_slice(b"<13>a message\n"); + + assert_eq!( + Ok(ProxyHeader::Complete { + source: Some(IpAddr::V4(Ipv4Addr::new(198, 51, 100, 3))), + length: prefix_len, + }), + parse(&buf, 1024) + ); + } + + #[test] + fn unix_and_unspecified_families_have_no_source() { + for family in [0x00, 0x31, 0x32] { + let header = v2_header(family, 0x1, &[0u8; 216]); + assert_eq!( + Ok(ProxyHeader::Complete { + source: None, + length: header.len(), + }), + parse(&header, 1024), + "family {:#x}", + family + ); + } + } + + #[test] + fn local_command_has_no_source() { + let header = v2_header(0x00, 0x0, &[]); + assert_eq!( + Ok(ProxyHeader::Complete { + source: None, + length: header.len(), + }), + parse(&header, 1024) + ); + } + + #[test] + fn rejects_v1_header() { + assert_eq!( + Err(ProxyError::Version), + parse(b"PROXY TCP4 192.0.2.1 10.0.0.1 514 6514\r\n", 1024) + ); + assert_eq!(Err(ProxyError::Version), parse(b"PRO", 1024)); + } + + #[test] + fn rejects_missing_header() { + assert_eq!(Err(ProxyError::Signature), parse(b"<13>a message\n", 1024)); + } + + #[test] + fn rejects_bad_version_command_and_family() { + let mut header = v2_header(0x11, 0x1, &ipv4_address([192, 0, 2, 7])); + header[12] = 0x31; + assert_eq!(Err(ProxyError::Version), parse(&header, 1024)); + + let mut header = v2_header(0x11, 0x1, &ipv4_address([192, 0, 2, 7])); + header[12] = 0x27; + assert_eq!(Err(ProxyError::Command), parse(&header, 1024)); + + let header = v2_header(0x41, 0x1, &ipv4_address([192, 0, 2, 7])); + assert_eq!(Err(ProxyError::Family), parse(&header, 1024)); + + let header = v2_header(0x11, 0x1, &[0u8; 4]); + assert_eq!(Err(ProxyError::Family), parse(&header, 1024)); + } + + #[test] + fn rejects_oversize_header() { + let header = v2_header(0x11, 0x1, &vec![0u8; 600]); + assert_eq!(Err(ProxyError::TooLong), parse(&header, 256)); + } + + #[test] + fn parses_identically_at_every_split_point() { + let mut buf = v2_header(0x11, 0x1, &ipv4_address([203, 0, 113, 9])); + let prefix_len = buf.len(); + buf.extend_from_slice(b"7 hello!"); + + let expected = ProxyHeader::Complete { + source: Some(IpAddr::V4(Ipv4Addr::new(203, 0, 113, 9))), + length: prefix_len, + }; + + for split in 0..buf.len() { + let result = parse(&buf[..split], 1024); + if split < prefix_len { + assert_eq!(Ok(ProxyHeader::Incomplete), result, "split {}", split); + } else { + assert_eq!(Ok(expected), result, "split {}", split); + } + } + } + + #[test] + fn detects_garbled_signature_as_soon_as_it_differs() { + let mut buf = v2_header(0x11, 0x1, &ipv4_address([203, 0, 113, 9])); + buf[5] = 0xff; + + for split in 6..buf.len() { + assert_eq!( + Err(ProxyError::Signature), + parse(&buf[..split], 1024), + "split {}", + split + ); + } + } +} diff --git a/hyper_ruby.gemspec b/hyper_ruby.gemspec index a0da765..239cae6 100644 --- a/hyper_ruby.gemspec +++ b/hyper_ruby.gemspec @@ -11,7 +11,7 @@ Gem::Specification.new do |spec| spec.summary = "Hyper-backed ruby web server" spec.description = "Hyper-backed ruby web server" spec.homepage = "https://github.com/betterstack/hyper-ruby" - spec.license = "MIT" + spec.licenses = ["MIT", "MPL-2.0"] spec.required_ruby_version = ">= 3.0.0" spec.required_rubygems_version = ">= 3.3.11" diff --git a/lib/hyper_ruby.rb b/lib/hyper_ruby.rb index dd14579..e473767 100644 --- a/lib/hyper_ruby.rb +++ b/lib/hyper_ruby.rb @@ -3,7 +3,46 @@ require_relative "hyper_ruby/version" require_relative "hyper_ruby/hyper_ruby" +# Server#configure takes a hash; alongside the HTTP keys (bind_address, +# tokio_threads, debug, recv_timeout, send_timeout, channel_capacity, +# max_connection_age) the syslog listeners accept: +# +# syslog_stream_path Unix socket path for the stream listener +# syslog_stream_bind "host:port" for the stream listener +# syslog_udp_bind "host:port" for the datagram listener (SO_REUSEPORT) +# syslog_proxy_protocol require a PROXY v2 header per connection (default false) +# syslog_proxy_header_timeout milliseconds allowed for that header (default 5000) +# syslog_max_frame_bytes largest stream frame (default 102400) +# syslog_max_pending undelivered messages allowed across all listeners +# (default 1000) +# syslog_max_pending_per_connection undelivered messages allowed for one connection, so +# a stalled sender cannot take the lot (default 64) +# syslog_max_connections stream connections accepted at once (default 10000) +# syslog_idle_timeout_ms close a connection after this long without data +# (default 0, no timeout) +# syslog_udp_max_datagram_bytes largest datagram accepted; one that fills the buffer +# counts as truncated and is dropped (default 65536) +# syslog_udp_so_rcvbuf SO_RCVBUF for the datagram socket +# syslog_work_ratio syslog messages a worker may take ahead of a waiting +# request (default 4) +# +# Syslog messages go to the same Server#run_worker block as requests, which tests +# what it was given: a HyperRuby::Request is answered with a Response, while a +# HyperRuby::SyslogMessage (readers message, peer_ip, transport, received_at_ns, +# message_id, attempt) is answered with an admission verdict. transport is :stream +# or :datagram; peer_ip is nil when the transport names no peer, such as a Unix +# socket connection with no PROXY header. message_id is stable across the retries +# of one message and unique for the life of the process, and attempt starts at 1, +# so a block can make its retries idempotent. +# +# A truthy return means the message was admitted; anything else, or a raised +# exception, means it was refused: a stream connection holds the message and +# retries it while its reads stall, a datagram is dropped and counted. Returning +# a Response for a syslog message is not a verdict: it is logged and the message +# is refused. +# +# Server#syslog_listening? reports listener readiness, and Server#syslog_stats +# returns the transport counters. module HyperRuby class Error < StandardError; end - # Your code goes here... end diff --git a/test/test_helper.rb b/test/test_helper.rb index 62ed8df..e133472 100644 --- a/test/test_helper.rb +++ b/test/test_helper.rb @@ -38,6 +38,34 @@ def with_configured_server(config, request_handler, &block) workers.map(&:join) if workers end + # Starts a server with syslog listeners configured; the worker block sends + # syslog messages to the given handler and requests to the request handler. + def with_syslog_server(config, syslog_handler, worker_count: 1, request_handler: nil, &block) + server = HyperRuby::Server.new + server.configure(config) + server.start + + workers = worker_count.times.map do + Thread.new do + server.run_worker do |work| + if work.is_a?(HyperRuby::SyslogMessage) + syslog_handler.call(work) + elsif request_handler + request_handler.call(work) + else + HyperRuby::Response.new(200, {}, "") + end + end + end + end + + block.call(server) + + ensure + server.stop if server + workers.map(&:join) if workers + end + def with_unix_socket_server(request_handler, &block) server = HyperRuby::Server.new server.configure({ diff --git a/test/test_syslog.rb b/test/test_syslog.rb new file mode 100644 index 0000000..b35a146 --- /dev/null +++ b/test/test_syslog.rb @@ -0,0 +1,709 @@ +# frozen_string_literal: true + +require "test_helper" +require "socket" +require "ipaddr" +require "net/http" +require "concurrent" + +class TestSyslog < HyperRubyTest + PROXY_V2_SIGNATURE = "\x0d\x0a\x0d\x0a\x00\x0d\x0a\x51\x55\x49\x54\x0a".b + + def setup + @collector = Collector.new + end + + # Collects delivered messages and decides what the handler returns. + class Collector + attr_reader :messages + attr_writer :delay, :accept_from_attempt, :refuse_matching + + def initialize + @mutex = Mutex.new + @messages = [] + @accept = true + @delay = 0 + @accept_from_attempt = 1 + @refuse_matching = nil + end + + def handler + lambda do |syslog| + sleep(@delay) if @delay > 0 + message = syslog.message + @mutex.synchronize do + @messages << { + message: message, peer: syslog.peer_ip, transport: syslog.transport, + received_at_ns: syslog.received_at_ns, message_id: syslog.message_id, + attempt: syslog.attempt + } + end + + next false if @refuse_matching && message.include?(@refuse_matching) + next false if syslog.attempt < @accept_from_attempt + + @accept + end + end + + def accept! + @accept = true + end + + def refuse! + @accept = false + end + + def bodies + @mutex.synchronize { @messages.map { |m| m[:message] } } + end + + def count + @mutex.synchronize { @messages.size } + end + + def for_message(body) + @mutex.synchronize { @messages.select { |m| m[:message] == body } } + end + end + + def test_worker_block_receives_a_syslog_message_object + seen = Queue.new + handler = lambda do |syslog| + seen << { class: syslog.class, inspect: syslog.inspect, message: syslog.message } + true + end + + with_syslog_server(syslog_config(stream: true), handler) do |server| + connect_stream(server) do |socket| + socket.write("<13>typed\n") + wait_until { !seen.empty? } + end + end + + delivered = seen.pop + assert_equal HyperRuby::SyslogMessage, delivered[:class] + assert_equal "<13>typed", delivered[:message] + assert_match(/HyperRuby::SyslogMessage transport=:stream/, delivered[:inspect]) + end + + def test_messages_survive_garbage_collection + kept = Queue.new + handler = lambda do |syslog| + kept << syslog + true + end + + with_syslog_server(syslog_config(stream: true), handler) do |server| + connect_stream(server) do |socket| + GC.stress = true + begin + 3.times { |i| socket.write("<13>collected #{i}\n") } + socket.flush + wait_until(timeout: 30) { kept.size == 3 } + ensure + GC.stress = false + end + end + + assert_equal 3, server.syslog_stats[:messages_delivered] + end + + # The objects outlive the worker call that yielded them. + GC.start(full_mark: true, immediate_sweep: true) + GC.compact + + messages = Array.new(kept.size) { kept.pop } + assert_equal ["<13>collected 0", "<13>collected 1", "<13>collected 2"], + messages.map(&:message) + assert_equal [:stream, :stream, :stream], messages.map(&:transport) + assert_equal [1, 1, 1], messages.map(&:attempt) + assert_equal ["127.0.0.1"] * 3, messages.map(&:peer_ip) + assert messages.map(&:message_id).all?(&:positive?) + end + + def test_returning_a_response_for_a_syslog_message_refuses_it + handler = lambda do |_syslog| + HyperRuby::Response.new(200, {}, "") + end + + with_syslog_server(syslog_config(udp: true), handler) do |server| + send_datagram(server, "<13>answered with a response") + wait_until { server.syslog_stats[:udp_dropped] == 1 } + + stats = server.syslog_stats + assert_equal 1, stats[:handler_errors] + assert_equal 0, stats[:messages_delivered] + assert_equal 0, stats[:deliveries_refused] + end + end + + def test_octet_counted_and_newline_frames + with_syslog_server(syslog_config(stream: true), @collector.handler) do |server| + before = Time.now.to_f * 1_000_000_000 + connect_stream(server) do |socket| + socket.write("11 hello there") + socket.write("<13>a newline frame\n") + socket.write("<13>a crlf frame\r\n") + wait_until { @collector.count == 3 } + end + + assert_equal ["hello there", "<13>a newline frame", "<13>a crlf frame"], @collector.bodies + + first = @collector.messages.first + assert_equal "127.0.0.1", first[:peer] + assert_equal :stream, first[:transport] + assert_equal Encoding::UTF_8, first[:message].encoding + assert_operator first[:received_at_ns], :>=, before + assert_equal 1, first[:attempt] + assert_equal 3, @collector.messages.map { |m| m[:message_id] }.uniq.size + assert_equal 3, server.syslog_stats[:messages_delivered] + end + end + + def test_frames_split_across_writes + with_syslog_server(syslog_config(stream: true), @collector.handler) do |server| + payload = "24 <13>an octet counted one<13>and a line\n" + connect_stream(server) do |socket| + payload.each_char do |char| + socket.write(char) + socket.flush + end + wait_until { @collector.count == 2 } + end + + assert_equal ["<13>an octet counted one", "<13>and a line"], @collector.bodies + end + end + + def test_newline_oversize_frame_is_skipped_but_connection_survives + config = syslog_config(stream: true).merge(syslog_max_frame_bytes: 32) + with_syslog_server(config, @collector.handler) do |server| + connect_stream(server) do |socket| + socket.write("<13>#{'a' * 64}\n") + socket.write("<13>short enough\n") + wait_until { @collector.count == 1 } + + assert_equal ["<13>short enough"], @collector.bodies + assert_equal 1, server.syslog_stats[:frames_rejected][:oversize] + end + end + end + + def test_octet_counted_oversize_frame_closes_connection + config = syslog_config(stream: true).merge(syslog_max_frame_bytes: 32) + with_syslog_server(config, @collector.handler) do |server| + connect_stream(server) do |socket| + socket.write("64 #{'a' * 64}") + # The declared bytes are discarded as they arrive, and the frame is + # rejected once the last of them is in. + sleep 0.05 + socket.write("<13>never read\n") + assert_closed(socket) + end + + assert_equal [], @collector.bodies + assert_equal 1, server.syslog_stats[:frames_rejected][:oversize] + end + end + + def test_oversize_body_starting_with_a_digit_is_not_delivered + config = syslog_config(stream: true).merge(syslog_max_frame_bytes: 16) + with_syslog_server(config, @collector.handler) do |server| + connect_stream(server) do |socket| + socket.write("26 ") + sleep 0.05 + socket.write("9 abcdefghZ") + sleep 0.05 + socket.write("aaaaaaaaaaaaaaa") + assert_closed(socket) + end + + assert_equal [], @collector.bodies + assert_equal 1, server.syslog_stats[:frames_rejected][:oversize] + end + end + + def test_invalid_utf8_closes_connection + with_syslog_server(syslog_config(stream: true), @collector.handler) do |server| + connect_stream(server) do |socket| + socket.write("11 valid frame") + socket.write("4 \xf0\x28\x8c\xbc".b) + socket.write("<13>never read\n") + assert_closed(socket) + end + + assert_equal ["valid frame"], @collector.bodies + assert_equal 1, server.syslog_stats[:frames_rejected][:invalid_utf8] + end + end + + def test_unix_socket_stream_listener + path = "/tmp/hyper_ruby_syslog_test_#{Process.pid}.sock" + config = syslog_config.merge(syslog_stream_path: path) + + with_syslog_server(config, @collector.handler) do |server| + assert server.syslog_listening? + + socket = UNIXSocket.new(path) + begin + socket.write("<13>over a unix socket\n") + wait_until { @collector.count == 1 } + ensure + socket.close + end + + assert_equal ["<13>over a unix socket"], @collector.bodies + assert_nil @collector.messages.first[:peer] + end + + refute File.exist?(path), "the socket file should be removed on stop" + end + + def test_udp_datagram_with_embedded_newlines_is_one_message + with_syslog_server(syslog_config(udp: true), @collector.handler) do |server| + send_datagram(server, "<13>first line\n<13>second line\n") + wait_until { @collector.count == 1 } + + message = @collector.messages.first + assert_equal "<13>first line\n<13>second line\n", message[:message] + assert_equal Encoding::ASCII_8BIT, message[:message].encoding + assert_equal :datagram, message[:transport] + assert_equal "127.0.0.1", message[:peer] + assert_equal 1, server.syslog_stats[:messages_delivered] + end + end + + def test_udp_datagram_keeps_invalid_utf8_bytes + with_syslog_server(syslog_config(udp: true), @collector.handler) do |server| + send_datagram(server, "<13>\xf0\x28\x8c\xbc".b) + wait_until { @collector.count == 1 } + + assert_equal "<13>\xf0\x28\x8c\xbc".b, @collector.messages.first[:message] + assert_equal 1, server.syslog_stats[:messages_delivered] + end + end + + def test_truncated_udp_datagram_is_dropped_and_counted + config = syslog_config(udp: true).merge(syslog_udp_max_datagram_bytes: 64) + with_syslog_server(config, @collector.handler) do |server| + send_datagram(server, "<13>#{'a' * 256}") + send_datagram(server, "<13>small one") + wait_until { @collector.count == 1 } + + assert_equal ["<13>small one"], @collector.bodies + assert_equal 1, server.syslog_stats[:udp_truncated] + end + end + + def test_proxy_protocol_supplies_the_peer_address + config = syslog_config(stream: true).merge(syslog_proxy_protocol: true) + with_syslog_server(config, @collector.handler) do |server| + connect_stream(server) do |socket| + socket.write(proxy_v2_header("192.0.2.7")) + socket.write("<13>behind a proxy\n") + wait_until { @collector.count == 1 } + end + + assert_equal ["<13>behind a proxy"], @collector.bodies + assert_equal "192.0.2.7", @collector.messages.first[:peer] + assert_equal 0, server.syslog_stats[:proxy_header_errors] + end + end + + def test_proxy_protocol_header_split_across_writes + config = syslog_config(stream: true).merge(syslog_proxy_protocol: true) + with_syslog_server(config, @collector.handler) do |server| + connect_stream(server) do |socket| + header = proxy_v2_header("198.51.100.4") + header.each_char do |char| + socket.write(char) + socket.flush + end + socket.write("<13>split header\n") + wait_until { @collector.count == 1 } + end + + assert_equal "198.51.100.4", @collector.messages.first[:peer] + end + end + + def test_missing_proxy_header_closes_connection + config = syslog_config(stream: true).merge(syslog_proxy_protocol: true) + with_syslog_server(config, @collector.handler) do |server| + connect_stream(server) do |socket| + socket.write("<13>no proxy header\n") + assert_closed(socket) + end + + assert_equal [], @collector.bodies + wait_until { server.syslog_stats[:proxy_header_errors] == 1 } + assert_equal 1, server.syslog_stats[:proxy_header_errors] + assert_equal 0, server.syslog_stats[:proxy_read_errors] + end + end + + def test_garbled_proxy_header_closes_connection + config = syslog_config(stream: true).merge(syslog_proxy_protocol: true) + with_syslog_server(config, @collector.handler) do |server| + connect_stream(server) do |socket| + header = proxy_v2_header("192.0.2.7").dup + header[5] = "\xff".b + socket.write(header) + socket.write("<13>never read\n") + assert_closed(socket) + end + + assert_equal [], @collector.bodies + wait_until { server.syslog_stats[:proxy_header_errors] == 1 } + assert_equal 1, server.syslog_stats[:proxy_header_errors] + end + end + + def test_connection_closed_before_the_proxy_header_counts_as_a_read_error + config = syslog_config(stream: true).merge(syslog_proxy_protocol: true) + with_syslog_server(config, @collector.handler) do |server| + socket = TCPSocket.new("127.0.0.1", @stream_port) + socket.write(PROXY_V2_SIGNATURE[0, 6]) + socket.close + + wait_until { server.syslog_stats[:proxy_read_errors] == 1 } + assert_equal 1, server.syslog_stats[:proxy_read_errors] + assert_equal 0, server.syslog_stats[:proxy_header_errors] + end + end + + def test_refused_messages_stop_reads_and_are_retried + config = syslog_config(stream: true).merge(syslog_max_pending_per_connection: 1) + with_syslog_server(config, @collector.handler) do |server| + @collector.refuse! + + connect_stream(server) do |socket| + 20.times { |i| socket.write("<13>message #{i}\n") } + socket.flush + + # The refused message is retried, and nothing behind it is delivered. + wait_until { @collector.count > 3 } + assert_equal ["<13>message 0"], @collector.bodies.uniq + assert_operator server.syslog_stats[:deliveries_refused], :>=, 3 + + @collector.accept! + wait_until { @collector.bodies.uniq.size == 20 } + assert_equal 20, @collector.bodies.uniq.size + end + end + end + + def test_retries_repeat_the_message_id_and_count_attempts + @collector.accept_from_attempt = 3 + + with_syslog_server(syslog_config(stream: true), @collector.handler) do |server| + connect_stream(server) do |socket| + socket.write("<13>retried\n") + socket.write("<13>plain\n") + wait_until { server.syslog_stats[:messages_delivered] == 2 } + end + + retried = @collector.for_message("<13>retried") + assert_equal [1, 2, 3], retried.map { |m| m[:attempt] } + assert_equal 1, retried.map { |m| m[:message_id] }.uniq.size + + plain = @collector.for_message("<13>plain") + assert_equal [1, 2, 3], plain.map { |m| m[:attempt] } + assert_equal 1, plain.map { |m| m[:message_id] }.uniq.size + refute_equal retried.first[:message_id], plain.first[:message_id] + assert_equal 2, server.syslog_stats[:messages_delivered] + assert_equal 4, server.syslog_stats[:deliveries_refused] + end + end + + def test_one_stalled_connection_does_not_block_another + config = syslog_config(stream: true).merge( + syslog_max_pending: 100, + syslog_max_pending_per_connection: 2 + ) + + with_syslog_server(config, @collector.handler) do |server| + @collector.refuse_matching = "stalled" + + connect_stream(server) do |stalled| + 20.times { |i| stalled.write("<13>stalled #{i}\n") } + stalled.flush + wait_until { @collector.count >= 2 } + + healthy = TCPSocket.new("127.0.0.1", @stream_port) + begin + 5.times { |i| healthy.write("<13>healthy #{i}\n") } + healthy.flush + wait_until { @collector.bodies.count { |body| body.include?("healthy") } == 5 } + ensure + healthy.close + end + + assert_equal 5, @collector.bodies.uniq.count { |body| body.include?("healthy") } + assert_operator @collector.bodies.uniq.count { |body| body.include?("stalled") }, :<=, 3 + assert_equal 5, server.syslog_stats[:messages_delivered] + end + end + end + + def test_refused_udp_datagrams_are_dropped_and_counted + with_syslog_server(syslog_config(udp: true), @collector.handler) do |server| + @collector.refuse! + send_datagram(server, "<13>refused datagram") + wait_until { server.syslog_stats[:udp_dropped] == 1 } + + assert_equal 1, server.syslog_stats[:udp_dropped] + assert_equal 0, server.syslog_stats[:messages_delivered] + assert_equal 1, @collector.count + end + end + + def test_handler_exception_refuses_the_message_and_counts_separately + raised = Queue.new + handler = lambda do |_syslog| + raised << true + raise "handler failure" + end + + with_syslog_server(syslog_config(udp: true), handler) do |server| + send_datagram(server, "<13>boom") + wait_until { server.syslog_stats[:udp_dropped] == 1 } + + stats = server.syslog_stats + assert_equal 1, stats[:handler_errors] + assert_equal 0, stats[:deliveries_refused] + assert_equal 0, stats[:messages_delivered] + refute raised.empty? + end + end + + def test_shutdown_delivers_already_framed_messages + with_syslog_server(syslog_config(stream: true), @collector.handler) do |server| + @collector.delay = 0.02 + + socket = TCPSocket.new("127.0.0.1", @stream_port) + begin + 10.times { |i| socket.write("<13>message #{i}\n") } + socket.flush + + # Let the listener read and frame everything before it is told to stop. + wait_until { @collector.count >= 1 } + sleep 0.1 + + server.stop + assert_equal 10, @collector.count + assert_equal 10, server.syslog_stats[:messages_delivered] + assert_equal 0, server.syslog_stats[:abandoned_at_shutdown] + refute server.syslog_listening? + ensure + socket.close + end + end + end + + def test_shutdown_counts_messages_it_could_not_admit + with_syslog_server(syslog_config(stream: true), @collector.handler) do |server| + @collector.refuse! + + socket = TCPSocket.new("127.0.0.1", @stream_port) + begin + socket.write("<13>never admitted\n") + socket.flush + wait_until { server.syslog_stats[:deliveries_refused] >= 1 } + + server.stop + stats = server.syslog_stats + assert_equal 1, stats[:abandoned_at_shutdown] + assert_equal 0, stats[:messages_delivered] + assert_equal 0, stats[:pending] + ensure + socket.close + end + end + end + + def test_connections_beyond_the_limit_are_refused + config = syslog_config(stream: true).merge(syslog_max_connections: 1) + with_syslog_server(config, @collector.handler) do |server| + connect_stream(server) do |socket| + socket.write("<13>the only connection\n") + wait_until { @collector.count == 1 } + + refused = TCPSocket.new("127.0.0.1", @stream_port) + begin + assert_closed(refused) + ensure + refused.close + end + + assert_equal 1, server.syslog_stats[:connections_refused] + socket.write("<13>still connected\n") + wait_until { @collector.count == 2 } + assert_equal 2, @collector.count + end + end + end + + def test_idle_connections_are_closed + config = syslog_config(stream: true).merge(syslog_idle_timeout_ms: 200) + with_syslog_server(config, @collector.handler) do |server| + connect_stream(server) do |socket| + socket.write("<13>then silence\n") + wait_until { @collector.count == 1 } + assert_closed(socket) + end + + wait_until { server.syslog_stats[:connections_closed] == 1 } + assert_equal 1, server.syslog_stats[:connections_closed] + end + end + + def test_listening_and_connection_counters + with_syslog_server(syslog_config(stream: true), @collector.handler) do |server| + assert server.syslog_listening? + + connect_stream(server) do |socket| + socket.write("<13>counted\n") + wait_until { @collector.count == 1 } + end + wait_until { server.syslog_stats[:connections_closed] == 1 } + + stats = server.syslog_stats + assert_equal 1, stats[:connections_opened] + assert_equal 1, stats[:connections_closed] + assert_equal 0, stats[:connections_refused] + assert_equal({ oversize: 0, invalid_utf8: 0, invalid_length: 0 }, stats[:frames_rejected]) + end + end + + def test_http_still_works_alongside_syslog + with_syslog_server(syslog_config(stream: true), @collector.handler) do |server| + response = Net::HTTP.get_response(URI("http://127.0.0.1:#{@http_port}/")) + assert_equal "200", response.code + + connect_stream(server) do |socket| + socket.write("<13>after an http request\n") + wait_until { @collector.count == 1 } + end + + assert_equal ["<13>after an http request"], @collector.bodies + end + end + + def test_syslog_delivery_outpaces_requests_under_http_load + # A handler slow enough that the request queue never empties, which is what + # would let requests starve syslog messages. + handled = Concurrent::AtomicFixnum.new(0) + slow_request = lambda do |_request| + sleep 0.005 + handled.increment + HyperRuby::Response.new(200, {}, "") + end + + with_syslog_server(syslog_config(stream: true), @collector.handler, + request_handler: slow_request) do |server| + flooding = true + floods = 8.times.map do + Thread.new do + Net::HTTP.start("127.0.0.1", @http_port) do |http| + http.request(Net::HTTP::Get.new("/")) while flooding + end + rescue IOError, EOFError, SystemCallError + nil + end + end + + begin + wait_until { handled.value > 5 } + before = handled.value + + sockets = 5.times.map { TCPSocket.new("127.0.0.1", @stream_port) } + begin + sockets.each_with_index do |socket, connection| + 10.times { |i| socket.write("<13>under load #{connection}-#{i}\n") } + socket.flush + end + wait_until(timeout: 10) { @collector.count == 50 } + ensure + sockets.each(&:close) + end + + during = handled.value - before + assert_equal 50, @collector.count + assert_equal 50, server.syslog_stats[:messages_delivered] + assert_operator during, :<, 25, + "syslog should outpace requests at the configured work ratio, saw #{during} requests" + ensure + flooding = false + floods.each(&:join) + end + end + end + + private + + @@next_port = 3400 + + def next_port + @@next_port += 1 + end + + def syslog_config(stream: false, udp: false) + @http_port = next_port + config = { bind_address: "127.0.0.1:#{@http_port}", tokio_threads: 1 } + + if stream + @stream_port = next_port + config[:syslog_stream_bind] = "127.0.0.1:#{@stream_port}" + end + + if udp + @udp_port = next_port + config[:syslog_udp_bind] = "127.0.0.1:#{@udp_port}" + end + + config + end + + def connect_stream(_server) + socket = TCPSocket.new("127.0.0.1", @stream_port) + yield socket + ensure + socket.close if socket && !socket.closed? + end + + def send_datagram(_server, payload, port: nil) + socket = UDPSocket.new + socket.send(payload, 0, "127.0.0.1", port || @udp_port) + ensure + socket.close if socket + end + + def proxy_v2_header(source_ip) + address = IPAddr.new(source_ip).hton + IPAddr.new("10.0.0.1").hton + [514, 6514].pack("nn") + PROXY_V2_SIGNATURE + [0x21, 0x11, address.bytesize].pack("CCn") + address + end + + def assert_closed(socket, timeout: 3) + deadline = Time.now + timeout + loop do + remaining = deadline - Time.now + flunk "connection was not closed by the server" if remaining <= 0 + + begin + return if socket.read_nonblock(1024, exception: false) == nil + rescue EOFError, Errno::ECONNRESET, IOError + return + end + + IO.select([socket], nil, nil, 0.05) + end + end + + def wait_until(timeout: 5) + deadline = Time.now + timeout + sleep(0.01) until yield || Time.now > deadline + yield + end +end