use bytes::{Buf, BytesMut};
use std::{
borrow::Cow,
io::Read,
net::{Shutdown, SocketAddr, TcpStream},
time::Duration,
};
use crate::{
consts::WireOp,
wire::{parse_event_notification, EventNotification},
};
use rsfbclient_core::{Charset, FbError};
const EPB_VERSION1: u8 = 1;
pub const MAX_EVENT_NAME_LEN: usize = u8::MAX as usize;
pub const MAX_EVENT_BLOCK_LEN: usize = u16::MAX as usize;
const AUX_CONNECT_TIMEOUT: Duration = Duration::from_secs(180);
pub fn normalize_event_name(name: &str) -> Result<&str, FbError> {
let name = name.trim_end_matches(' ');
if name.is_empty() {
return Err(FbError::from("A firebird event name cannot be empty"));
}
Ok(name)
}
fn encode_event_name<'a>(charset: &Charset, name: &'a str) -> Result<Cow<'a, [u8]>, FbError> {
let name = normalize_event_name(name)?;
let encoded = charset.encode(name)?;
if encoded.len() > MAX_EVENT_NAME_LEN {
return Err(FbError::from(format!(
"A firebird event name is limited to {} bytes, but '{}' uses {} once encoded in {}",
MAX_EVENT_NAME_LEN,
name,
encoded.len(),
charset.on_firebird
)));
}
Ok(encoded)
}
pub fn event_block<'a, I>(charset: &Charset, events: I) -> Result<Vec<u8>, FbError>
where
I: IntoIterator<Item = (&'a str, u32)>,
{
let mut epb = vec![EPB_VERSION1];
for (name, count) in events {
let name = encode_event_name(charset, name)?;
epb.push(name.len() as u8);
epb.extend_from_slice(&name);
epb.extend_from_slice(&count.to_le_bytes());
}
if epb.len() > MAX_EVENT_BLOCK_LEN {
return Err(FbError::from(format!(
"The event parameter block is limited to {} bytes, but the requested events need {}",
MAX_EVENT_BLOCK_LEN,
epb.len()
)));
}
Ok(epb)
}
pub fn parse_event_block(charset: &Charset, epb: &[u8]) -> Result<Vec<(String, u32)>, FbError> {
let (version, mut rest) = epb
.split_first()
.ok_or_else(|| FbError::from("Empty event parameter block"))?;
if *version != EPB_VERSION1 {
return Err(FbError::from(format!(
"Unsupported event parameter block version: {}",
version
)));
}
let mut events = Vec::new();
while let Some((len, tail)) = rest.split_first() {
let len = *len as usize;
if tail.len() < len + 4 {
return Err(FbError::from("Truncated event parameter block"));
}
let (name, tail) = tail.split_at(len);
let name = charset.decode(name)?;
let (count, tail) = tail.split_at(4);
let count = u32::from_le_bytes([count[0], count[1], count[2], count[3]]);
events.push((name, count));
rest = tail;
}
Ok(events)
}
pub fn event_count(events: &[(String, u32)], name: &str) -> Result<u32, FbError> {
events
.iter()
.find(|(event, _)| event == name)
.map(|(_, count)| *count)
.ok_or_else(|| {
FbError::from(format!(
"The server notification does not hold the '{}' event",
name
))
})
}
fn connect_aux(addr: SocketAddr, timeout: Duration) -> Result<TcpStream, FbError> {
let socket = TcpStream::connect_timeout(&addr, timeout)?;
set_keep_alive(&socket);
Ok(socket)
}
#[cfg(unix)]
fn set_keep_alive(socket: &TcpStream) {
use std::os::{
fd::AsRawFd,
raw::{c_int, c_void},
};
extern "C" {
fn setsockopt(
socket: c_int,
level: c_int,
name: c_int,
value: *const c_void,
option_len: u32,
) -> c_int;
}
let (level, name) = keep_alive_option();
let enable: c_int = 1;
unsafe {
setsockopt(
socket.as_raw_fd(),
level,
name,
&enable as *const c_int as *const c_void,
std::mem::size_of::<c_int>() as u32,
);
}
}
#[cfg(windows)]
fn set_keep_alive(socket: &TcpStream) {
use std::os::{
raw::{c_char, c_int},
windows::io::AsRawSocket,
};
#[link(name = "ws2_32")]
extern "system" {
fn setsockopt(
socket: usize,
level: c_int,
name: c_int,
value: *const c_char,
option_len: c_int,
) -> c_int;
}
let (level, name) = keep_alive_option();
let enable: c_int = 1;
unsafe {
setsockopt(
socket.as_raw_socket() as usize,
level,
name,
&enable as *const c_int as *const c_char,
std::mem::size_of::<c_int>() as c_int,
);
}
}
#[cfg(not(any(unix, windows)))]
fn set_keep_alive(_socket: &TcpStream) {}
#[cfg(any(unix, windows))]
fn keep_alive_option() -> (std::os::raw::c_int, std::os::raw::c_int) {
if cfg!(all(
any(target_os = "linux", target_os = "android"),
not(any(
target_arch = "mips",
target_arch = "mips32r6",
target_arch = "mips64",
target_arch = "mips64r6",
target_arch = "sparc",
target_arch = "sparc64"
))
)) {
(1, 9)
} else {
(0xffff, 0x0008)
}
}
pub struct EventChannel {
db_handle: u32,
charset: Charset,
socket: TcpStream,
buff: BytesMut,
}
impl EventChannel {
pub fn open(
db_handle: u32,
charset: Charset,
main_peer: SocketAddr,
port: u16,
) -> Result<Self, FbError> {
let mut addr = main_peer;
addr.set_port(port);
Ok(Self {
db_handle,
charset,
socket: connect_aux(addr, AUX_CONNECT_TIMEOUT)?,
buff: BytesMut::with_capacity(256),
})
}
pub fn db_handle(&self) -> u32 {
self.db_handle
}
pub fn recv_event(&mut self, event_id: u32) -> Result<Vec<(String, u32)>, FbError> {
loop {
let notification = self.recv_notification()?;
if notification.event_id == event_id {
return parse_event_block(&self.charset, ¬ification.epb);
}
}
}
fn recv_notification(&mut self) -> Result<EventNotification, FbError> {
loop {
self.fill(4)?;
let op_code = self.peek_u32(0);
if op_code == WireOp::Dummy as u32 {
self.buff.advance(4);
continue;
}
if op_code == WireOp::Exit as u32 || op_code == WireOp::Disconnect as u32 {
return Err(FbError::from(
"The server closed the event channel while waiting for an event",
));
}
if op_code != WireOp::Event as u32 {
return Err(FbError::from(format!(
"Unexpected operation {} on the event channel",
op_code
)));
}
self.fill(12)?;
let epb_len = self.peek_u32(8) as usize;
if epb_len > MAX_EVENT_BLOCK_LEN {
return Err(FbError::from(format!(
"The server announced an event parameter block of {} bytes, over the {} bytes limit",
epb_len, MAX_EVENT_BLOCK_LEN
)));
}
let len = 12 + epb_len.next_multiple_of(4) + 12;
self.fill(len)?;
let mut packet = self.buff.split_to(len).freeze();
return parse_event_notification(&mut packet);
}
}
fn peek_u32(&self, offset: usize) -> u32 {
u32::from_be_bytes([
self.buff[offset],
self.buff[offset + 1],
self.buff[offset + 2],
self.buff[offset + 3],
])
}
fn fill(&mut self, len: usize) -> Result<(), FbError> {
let mut chunk = [0; 512];
while self.buff.len() < len {
let read = self.socket.read(&mut chunk)?;
if read == 0 {
return Err(FbError::from(
"The event channel was closed while waiting for an event",
));
}
self.buff.extend_from_slice(&chunk[..read]);
}
Ok(())
}
}
impl Drop for EventChannel {
fn drop(&mut self) {
let _ = self.socket.shutdown(Shutdown::Both);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::wire::{cancel_events, connect_request, parse_aux_port, que_events};
use bytes::{BufMut, Bytes};
use rsfbclient_core::charset::{ASCII, ISO_8859_1, UTF_8, WIN_1252};
use std::{io::Write, net::TcpListener, thread, time::Instant};
#[cfg(unix)]
fn keep_alive_enabled(socket: &TcpStream) -> bool {
use std::os::{
fd::AsRawFd,
raw::{c_int, c_void},
};
extern "C" {
fn getsockopt(
socket: c_int,
level: c_int,
name: c_int,
value: *mut c_void,
option_len: *mut u32,
) -> c_int;
}
let (level, name) = keep_alive_option();
let mut enabled: c_int = 0;
let mut len = std::mem::size_of::<c_int>() as u32;
let res = unsafe {
getsockopt(
socket.as_raw_fd(),
level,
name,
&mut enabled as *mut c_int as *mut c_void,
&mut len,
)
};
assert_eq!(res, 0, "getsockopt(SO_KEEPALIVE) failed");
enabled != 0
}
#[cfg(windows)]
fn keep_alive_enabled(socket: &TcpStream) -> bool {
use std::os::{
raw::{c_char, c_int},
windows::io::AsRawSocket,
};
#[link(name = "ws2_32")]
extern "system" {
fn getsockopt(
socket: usize,
level: c_int,
name: c_int,
value: *mut c_char,
option_len: *mut c_int,
) -> c_int;
}
let (level, name) = keep_alive_option();
let mut enabled: c_int = 0;
let mut len = std::mem::size_of::<c_int>() as c_int;
let res = unsafe {
getsockopt(
socket.as_raw_socket() as usize,
level,
name,
&mut enabled as *mut c_int as *mut c_char,
&mut len,
)
};
assert_eq!(res, 0, "getsockopt(SO_KEEPALIVE) failed");
enabled != 0
}
fn event_packet(event_id: u32, events: &[(&str, u32)]) -> Bytes {
let epb = event_block(&UTF_8, events.iter().copied()).unwrap();
let mut packet = BytesMut::new();
packet.put_u32(WireOp::Event as u32);
packet.put_u32(1); packet.put_u32(epb.len() as u32);
packet.put_slice(&epb);
packet.put_slice(&vec![0; epb.len().next_multiple_of(4) - epb.len()]); packet.put_u32(0); packet.put_u32(0); packet.put_u32(event_id);
packet.freeze()
}
#[test]
fn event_block_single_name() {
let epb = event_block(&UTF_8, [("evento", 0)]).unwrap();
assert_eq!(
epb,
vec![
1, 6, b'e', b'v', b'e', b'n', b't', b'o', 0, 0, 0, 0, ]
);
}
#[test]
fn event_block_counter_is_little_endian() {
let epb = event_block(&UTF_8, [("a", 0x0102_0304)]).unwrap();
assert_eq!(epb, vec![1, 1, b'a', 0x04, 0x03, 0x02, 0x01]);
}
#[test]
fn event_block_multiple_names() {
let epb = event_block(&UTF_8, [("ab", 1), ("c", 2)]).unwrap();
assert_eq!(epb, vec![1, 2, b'a', b'b', 1, 0, 0, 0, 1, b'c', 2, 0, 0, 0]);
}
#[test]
fn event_block_strips_trailing_blanks() {
assert_eq!(
event_block(&UTF_8, [("ab ", 0)]).unwrap(),
event_block(&UTF_8, [("ab", 0)]).unwrap()
);
}
#[test]
fn event_name_limits_are_counted_in_encoded_bytes() {
let ascii = "a".repeat(MAX_EVENT_NAME_LEN);
assert!(encode_event_name(&UTF_8, &ascii).is_ok());
let too_long = "a".repeat(MAX_EVENT_NAME_LEN + 1);
assert!(encode_event_name(&UTF_8, &too_long).is_err());
let accented = "é".repeat(128);
assert_eq!(accented.chars().count(), 128);
assert_eq!(accented.len(), 256);
assert!(encode_event_name(&UTF_8, &accented).is_err());
assert_eq!(
encode_event_name(&ISO_8859_1, &accented).unwrap().len(),
128
);
}
#[test]
fn event_name_cannot_be_empty() {
assert!(normalize_event_name("").is_err());
assert!(normalize_event_name(" ").is_err());
assert!(event_block(&UTF_8, [("", 0)]).is_err());
assert!(event_block(&ISO_8859_1, [(" ", 0)]).is_err());
}
#[test]
fn event_name_is_encoded_with_the_connection_charset() {
assert_eq!(
event_block(&ISO_8859_1, [("é", 0)]).unwrap(),
vec![1, 1, 0xe9, 0, 0, 0, 0]
);
assert_eq!(
event_block(&UTF_8, [("é", 0)]).unwrap(),
vec![1, 2, 0xc3, 0xa9, 0, 0, 0, 0]
);
}
#[test]
fn event_name_in_windows_1252() {
assert_eq!(
event_block(&WIN_1252, [("caf\u{e9}\u{20ac}", 0)]).unwrap(),
vec![1, 5, b'c', b'a', b'f', 0xe9, 0x80, 0, 0, 0, 0]
);
assert!(event_block(&ISO_8859_1, [("\u{20ac}", 0)]).is_err());
}
#[test]
fn event_name_that_the_charset_cannot_represent_is_rejected() {
let err = event_block(&ASCII, [("é", 0)]).unwrap_err();
assert!(
format!("{}", err).contains("ascii"),
"unexpected error: {}",
err
);
assert!(event_block(&ISO_8859_1, [("日本", 0)]).is_err());
}
#[test]
fn event_names_roundtrip_through_a_non_utf8_charset() {
let name = "caf\u{e9}";
let epb = event_block(&WIN_1252, [(name, 7)]).unwrap();
let events = parse_event_block(&WIN_1252, &epb).unwrap();
assert_eq!(events, vec![(name.to_string(), 7)]);
assert_eq!(event_count(&events, name).unwrap(), 7);
}
#[test]
fn parse_event_block_rejects_bytes_the_charset_cannot_decode() {
assert!(parse_event_block(&UTF_8, &[1, 1, 0xff, 0, 0, 0, 0]).is_err());
assert!(parse_event_block(&ASCII, &[1, 1, 0xe9, 0, 0, 0, 0]).is_err());
assert_eq!(
parse_event_block(&ISO_8859_1, &[1, 1, 0xe9, 0, 0, 0, 0]).unwrap(),
vec![("é".to_string(), 0)]
);
}
#[test]
fn event_block_size_is_limited() {
let name = "a".repeat(MAX_EVENT_NAME_LEN);
let events = (0..300).map(|_| (name.as_str(), 0)).collect::<Vec<_>>();
assert!(event_block(&UTF_8, events).is_err());
}
#[test]
fn parse_event_block_roundtrip() {
let epb = event_block(&UTF_8, [("evento", 3), ("outro", 0)]).unwrap();
assert_eq!(
parse_event_block(&UTF_8, &epb).unwrap(),
vec![("evento".to_string(), 3), ("outro".to_string(), 0)]
);
}
#[test]
fn parse_event_block_rejects_invalid_input() {
assert!(parse_event_block(&UTF_8, &[]).is_err());
assert!(parse_event_block(&UTF_8, &[2, 1, b'a', 0, 0, 0, 0]).is_err());
assert!(parse_event_block(&UTF_8, &[1, 6, b'a', 0, 0, 0, 0]).is_err());
assert!(parse_event_block(&UTF_8, &[1, 1, b'a', 0, 0]).is_err());
}
#[test]
fn event_count_of_a_missing_name_is_an_error() {
let events = vec![("a".to_string(), 7)];
assert_eq!(event_count(&events, "a").unwrap(), 7);
assert!(event_count(&events, "b").is_err());
}
#[test]
fn connect_request_layout() {
assert_eq!(
connect_request(0x0a0b_0c0d).as_ref(),
[
0, 0, 0, 53, 0, 0, 0, 1, 0x0a, 0x0b, 0x0c, 0x0d, 0, 0, 0, 0, ]
);
}
#[test]
fn que_events_layout() {
let epb = event_block(&UTF_8, [("a", 0)]).unwrap();
assert_eq!(epb.len(), 7);
assert_eq!(
que_events(1, &epb, 42).as_ref(),
[
0, 0, 0, 48, 0, 0, 0, 1, 0, 0, 0, 7, 1, 1, b'a', 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 42, ]
);
}
#[test]
fn cancel_events_layout() {
assert_eq!(
cancel_events(1, 42).as_ref(),
[
0, 0, 0, 49, 0, 0, 0, 1, 0, 0, 0, 42, ]
);
}
#[test]
fn parse_event_notification_reads_the_counters() {
let mut packet = event_packet(42, &[("evento", 5)]);
let notification = parse_event_notification(&mut packet).unwrap();
assert_eq!(notification.event_id, 42);
assert_eq!(
parse_event_block(&UTF_8, ¬ification.epb).unwrap(),
vec![("evento".to_string(), 5)]
);
assert!(packet.is_empty());
}
#[test]
fn parse_event_notification_of_an_unpadded_block() {
let epb = event_block(&UTF_8, [("ab", 0)]).unwrap();
assert_eq!(epb.len(), 8);
let mut packet = event_packet(1, &[("ab", 0)]);
let notification = parse_event_notification(&mut packet).unwrap();
assert_eq!(
parse_event_block(&UTF_8, ¬ification.epb).unwrap(),
vec![("ab".to_string(), 0)]
);
assert!(packet.is_empty());
}
#[test]
fn parse_event_notification_rejects_another_operation() {
let mut packet = Bytes::from_static(&[0, 0, 0, 9]);
assert!(parse_event_notification(&mut packet).is_err());
}
#[test]
fn aux_port_is_read_in_network_byte_order() {
let sockaddr = [
2, 0, 0x80, 0x03, 192, 168, 1, 10, 0, 0, 0, 0, 0, 0, 0, 0, ];
assert_eq!(parse_aux_port(&sockaddr).unwrap(), 0x8003);
}
#[test]
fn aux_port_ignores_the_address_family_layout() {
let macos = [16, 2, 0x0b, 0xea, 127, 0, 0, 1];
let posix = [2, 0, 0x0b, 0xea, 127, 0, 0, 1];
assert_eq!(parse_aux_port(&macos).unwrap(), 3050);
assert_eq!(parse_aux_port(&posix).unwrap(), 3050);
}
#[test]
fn aux_port_rejects_invalid_data() {
assert!(parse_aux_port(&[]).is_err());
assert!(parse_aux_port(&[2, 0, 1]).is_err());
assert!(parse_aux_port(&[2, 0, 0, 0, 127, 0, 0, 1]).is_err());
}
#[test]
fn aux_connect_timeout_is_the_firebird_connection_timeout() {
assert_eq!(AUX_CONNECT_TIMEOUT, Duration::from_secs(180));
}
#[test]
#[cfg(any(unix, windows))]
fn aux_socket_enables_keep_alive() {
let listener = TcpListener::bind(("127.0.0.1", 0)).unwrap();
let addr = listener.local_addr().unwrap();
let plain = TcpStream::connect(addr).unwrap();
assert!(!keep_alive_enabled(&plain));
let stream = connect_aux(addr, AUX_CONNECT_TIMEOUT).unwrap();
assert!(keep_alive_enabled(&stream));
}
#[test]
fn aux_socket_is_blocking() {
let listener = TcpListener::bind(("127.0.0.1", 0)).unwrap();
let addr = listener.local_addr().unwrap();
let mut stream = connect_aux(addr, AUX_CONNECT_TIMEOUT).unwrap();
let (mut server, _) = listener.accept().unwrap();
let reader = thread::spawn(move || {
let mut byte = [0; 1];
stream.read_exact(&mut byte).map(|_| byte[0])
});
thread::sleep(Duration::from_millis(200));
assert!(!reader.is_finished());
server.write_all(&[42]).unwrap();
assert_eq!(reader.join().unwrap().unwrap(), 42);
}
#[test]
fn aux_connect_succeeds_within_the_timeout() {
let listener = TcpListener::bind(("127.0.0.1", 0)).unwrap();
let addr = listener.local_addr().unwrap();
let start = Instant::now();
let stream = connect_aux(addr, Duration::from_secs(30)).unwrap();
assert_eq!(stream.peer_addr().unwrap(), addr);
assert!(start.elapsed() < Duration::from_secs(5));
}
#[test]
fn aux_connect_reports_a_refused_connection() {
let addr = {
let listener = TcpListener::bind(("127.0.0.1", 0)).unwrap();
listener.local_addr().unwrap()
};
let start = Instant::now();
assert!(connect_aux(addr, AUX_CONNECT_TIMEOUT).is_err());
assert!(start.elapsed() < Duration::from_secs(5));
}
}