use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use bytes::{Bytes, BytesMut};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, ReadHalf, WriteHalf, split};
use tokio::net::TcpStream;
use tokio::sync::{mpsc, oneshot};
use crate::message::Message;
use crate::parser::{Frame, extract_frame};
use crate::session::{Command, Connection, SessionHandle};
use crate::tags;
pub(crate) type SessionKey = (String, String, String);
pub(crate) enum Tls {
None,
#[cfg(feature = "tls")]
Server(tokio_rustls::TlsAcceptor),
#[cfg(feature = "tls")]
Client(crate::tls::ClientTls),
}
pub(crate) fn spawn_io_tasks<S>(
stream: S,
leftover: BytesMut,
) -> (mpsc::Receiver<Bytes>, mpsc::Sender<Bytes>)
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
let (read_half, write_half) = split(stream);
let (in_tx, in_rx) = mpsc::channel(256);
let (out_tx, out_rx) = mpsc::channel::<Bytes>(256);
tokio::spawn(read_task(read_half, leftover, in_tx));
tokio::spawn(write_task(write_half, out_rx));
(in_rx, out_tx)
}
async fn read_task<S: AsyncRead + Send + 'static>(
mut read_half: ReadHalf<S>,
mut buf: BytesMut,
in_tx: mpsc::Sender<Bytes>,
) {
loop {
while let Frame::Message(raw) = extract_frame(&mut buf) {
if in_tx.send(raw).await.is_err() {
return; }
}
match read_half.read_buf(&mut buf).await {
Ok(0) | Err(_) => return, Ok(_) => {}
}
}
}
async fn write_task<S: AsyncWrite + Send + 'static>(
mut write_half: WriteHalf<S>,
mut out_rx: mpsc::Receiver<Bytes>,
) {
while let Some(raw) = out_rx.recv().await {
if write_half.write_all(&raw).await.is_err() {
return;
}
}
let _ = write_half.shutdown().await;
}
pub(crate) async fn run_acceptor(
listener: tokio::net::TcpListener,
registry: Arc<HashMap<SessionKey, SessionHandle>>,
tls: Tls,
) {
let tls = Arc::new(tls);
loop {
match listener.accept().await {
Ok((stream, peer)) => {
let _ = stream.set_nodelay(true);
let registry = registry.clone();
let tls = tls.clone();
tokio::spawn(async move {
if let Err(e) = accept_connection(stream, ®istry, &tls).await {
tracing::warn!("connection from {peer} rejected: {e}");
}
});
}
Err(e) => {
tracing::warn!("accept error: {e}");
tokio::time::sleep(Duration::from_millis(100)).await;
}
}
}
}
async fn accept_connection(
stream: TcpStream,
registry: &HashMap<SessionKey, SessionHandle>,
tls: &Tls,
) -> std::result::Result<(), String> {
match tls {
Tls::None => handle_connection(stream, registry).await,
#[cfg(feature = "tls")]
Tls::Server(acceptor) => {
let stream =
acceptor.accept(stream).await.map_err(|e| format!("TLS handshake failed: {e}"))?;
handle_connection(stream, registry).await
}
#[cfg(feature = "tls")]
Tls::Client(_) => Err("acceptor was given client TLS config".into()),
}
}
async fn handle_connection<S>(
mut stream: S,
registry: &HashMap<SessionKey, SessionHandle>,
) -> std::result::Result<(), String>
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
let mut buf = BytesMut::with_capacity(8192);
let first = tokio::time::timeout(Duration::from_secs(30), async {
loop {
if let Frame::Message(raw) = extract_frame(&mut buf) {
return Ok::<_, String>(raw);
}
match stream.read_buf(&mut buf).await {
Ok(0) => return Err("peer closed before sending a message".into()),
Ok(_) => {}
Err(e) => return Err(format!("read error: {e}")),
}
}
})
.await
.map_err(|_| "timed out waiting for first message".to_string())??;
let msg = Message::parse(&first, false).map_err(|e| format!("unparseable first message: {e}"))?;
let get = |tag| {
msg.header
.get_raw(tag)
.map(|v| String::from_utf8_lossy(v).into_owned())
.ok_or_else(|| format!("first message missing tag {tag}"))
};
let key: SessionKey = (
get(tags::BEGIN_STRING)?,
get(tags::TARGET_COMP_ID)?,
get(tags::SENDER_COMP_ID)?,
);
let handle = registry.get(&key).ok_or_else(|| {
format!("no session configured for {}:{}->{}", key.0, key.1, key.2)
})?;
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
loop {
match handle.status().await {
Ok(s) if s.connected => {
if tokio::time::Instant::now() >= deadline {
return Err("session already has an active connection".into());
}
tokio::time::sleep(Duration::from_millis(20)).await;
}
Ok(_) => break,
Err(_) => return Err("session task not running".into()),
}
}
let (in_rx, out_tx) = spawn_io_tasks(stream, buf);
let (relay_tx, relay_rx) = mpsc::channel(256);
relay_tx.send(first).await.map_err(|e| e.to_string())?;
let mut in_rx = in_rx;
tokio::spawn(async move {
while let Some(raw) = in_rx.recv().await {
if relay_tx.send(raw).await.is_err() {
return;
}
}
});
handle
.cmd_tx
.send(Command::Attach(Connection {
inbound: relay_rx,
outbound: out_tx,
disconnected: None,
}))
.await
.map_err(|_| "session task not running".to_string())
}
pub(crate) async fn run_initiator(
host: String,
port: u16,
reconnect_interval: Duration,
handle: SessionHandle,
tls: Tls,
) {
loop {
match TcpStream::connect((host.as_str(), port)).await {
Ok(stream) => {
let _ = stream.set_nodelay(true);
let disconnected = attach_initiator(stream, &handle, &tls).await;
match disconnected {
Ok(Some(disc_rx)) => {
let _ = disc_rx.await;
}
Ok(None) => return, Err(e) => {
tracing::info!(session = %handle.id, "TLS/connect setup failed: {e}");
}
}
}
Err(e) => {
tracing::info!(session = %handle.id, "connect to {host}:{port} failed: {e}");
}
}
tokio::time::sleep(reconnect_interval).await;
}
}
async fn attach_initiator(
stream: TcpStream,
handle: &SessionHandle,
tls: &Tls,
) -> std::result::Result<Option<oneshot::Receiver<()>>, String> {
let (in_rx, out_tx) = match tls {
Tls::None => spawn_io_tasks(stream, BytesMut::with_capacity(8192)),
#[cfg(feature = "tls")]
Tls::Client(client) => {
let stream = client
.connector
.connect(client.server_name.clone(), stream)
.await
.map_err(|e| format!("TLS handshake failed: {e}"))?;
spawn_io_tasks(stream, BytesMut::with_capacity(8192))
}
#[cfg(feature = "tls")]
Tls::Server(_) => return Err("initiator was given server TLS config".into()),
};
let (disc_tx, disc_rx) = oneshot::channel();
if handle
.cmd_tx
.send(Command::Attach(Connection {
inbound: in_rx,
outbound: out_tx,
disconnected: Some(disc_tx),
}))
.await
.is_err()
{
return Ok(None); }
Ok(Some(disc_rx))
}