use std::{
collections::{hash_map::Entry, HashMap},
fmt::Display,
io::Read,
sync::{
atomic::{AtomicBool, Ordering},
Arc,
},
time::Duration,
};
use camino::Utf8PathBuf;
use chrono::{DateTime, Utc};
use error::DaemonResult;
use log::{error, info, warn};
use num_traits::FromPrimitive;
use tokio::{
select,
sync::{
mpsc::{channel, Receiver, Sender},
oneshot,
},
task::JoinHandle,
time::MissedTickBehavior,
};
use crate::{
filestore::{ChecksumType, FileStore},
pdu::{
header::{
CRCFlag, Condition, DeliveryCode, Direction, FileSizeFlag, FileStatusCode, PDUHeader,
SegmentedData, TransactionStatus, TransmissionMode,
},
ops::{EntityID, TransactionSeqNum},
PDUEncode, PDUError, PDUResult, PDU,
},
transaction::FaultHandlerAction,
transaction::{Metadata, TransactionConfig, TransactionID, TransactionState},
};
pub mod error;
pub mod segments;
pub mod timer;
pub mod transport;
pub mod user;
pub(crate) use timer::*;
pub use user::*;
use self::error::DaemonError;
use self::transport::PDUTransport;
use crate::transaction::{error::TransactionError, recv::RecvTransaction, send::SendTransaction};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PutRequest {
pub source_filename: Utf8PathBuf,
pub destination_filename: Utf8PathBuf,
pub destination_entity_id: EntityID,
pub transmission_mode: TransmissionMode,
}
#[derive(Debug)]
pub enum UserPrimitive {
Put(PutRequest, oneshot::Sender<TransactionID>),
Cancel(TransactionID),
Report(TransactionID, oneshot::Sender<Report>),
}
#[derive(Debug, Clone)]
#[cfg_attr(test, derive(PartialEq))]
pub struct Report {
pub id: TransactionID,
pub state: TransactionState,
pub status: TransactionStatus,
pub condition: Condition,
pub file_size: u64,
pub file_bytes_received: Option<u64>,
pub file_bytes_sent: Option<u64>,
pub empty_nak_received: bool,
pub direction: Option<Direction>,
pub file_name: String,
pub submit_date: DateTime<Utc>,
}
impl Report {
pub fn encode(self) -> Vec<u8> {
let mut buff = self.id.0.encode();
buff.extend(self.id.1.encode());
buff.push(self.state as u8);
buff.push(self.status as u8);
buff.push(self.condition as u8);
buff
}
pub fn decode<T: Read>(buffer: &mut T) -> PDUResult<Self> {
let id = {
let entity_id = EntityID::decode(buffer)?;
let sequence_num = TransactionSeqNum::decode(buffer)?;
TransactionID(entity_id, sequence_num)
};
let mut u8_buff = [0_u8; 1];
let state = {
buffer.read_exact(&mut u8_buff)?;
let possible = u8_buff[0];
TransactionState::from_u8(possible).ok_or(PDUError::InvalidState(possible))?
};
let status = {
buffer.read_exact(&mut u8_buff)?;
let possible = u8_buff[0];
TransactionStatus::from_u8(possible)
.ok_or(PDUError::InvalidTransactionStatus(possible))?
};
let condition = {
buffer.read_exact(&mut u8_buff)?;
let possible = u8_buff[0];
Condition::from_u8(possible).ok_or(PDUError::InvalidCondition(possible))?
};
Ok(Self {
id,
state,
status,
condition,
empty_nak_received: false,
file_size: 0,
file_bytes_received: None,
file_bytes_sent: None,
direction: None,
file_name: String::new(),
submit_date: Utc::now(),
})
}
}
#[derive(Debug, Clone)]
pub struct MetadataRecvIndication {
pub id: TransactionID,
pub source_filename: Utf8PathBuf,
pub destination_filename: Utf8PathBuf,
pub file_size: u64,
pub transmission_mode: TransmissionMode,
}
#[derive(Debug, Clone)]
pub struct FileSegmentIndication {
pub id: TransactionID,
pub offset: u64,
pub length: u64,
}
#[derive(Debug, Clone)]
pub struct FinishedIndication {
pub id: TransactionID,
pub report: Report,
pub file_status: FileStatusCode,
pub delivery_code: DeliveryCode,
}
#[derive(Debug, Clone)]
pub enum Indication {
Transaction(TransactionID),
EoFSent(TransactionID),
EoFRecv(TransactionID),
Finished(FinishedIndication),
MetadataRecv(MetadataRecvIndication),
FileSegmentRecv(FileSegmentIndication),
Report(Report),
}
#[derive(Debug, Copy, Clone, PartialEq, Eq)]
pub enum NakProcedure {
Immediate(Duration ),
Deferred(Duration ),
}
#[derive(Clone)]
pub struct EntityConfig {
pub fault_handler_override: HashMap<Condition, FaultHandlerAction>,
pub file_size_segment: u16,
pub default_transaction_max_count: u32,
pub inactivity_timeout: i64,
pub eof_timeout: i64,
pub nak_timeout: i64,
pub crc_flag: CRCFlag,
pub checksum_type: ChecksumType,
pub nak_procedure: NakProcedure,
pub local_entity_id: u16,
pub remote_entity_id: u16,
pub local_server_addr: &'static str,
pub remote_server_addr: &'static str,
pub progress_report_interval_secs: i64,
}
#[derive(Debug)]
pub enum Command {
Pdu(PDU),
Abandon,
Report(oneshot::Sender<Report>),
}
impl Display for Command {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{:?}", &self)
}
}
fn construct_metadata<T: FileStore + Send + 'static>(
filestore: &Arc<T>,
req: PutRequest,
config: EntityConfig,
) -> DaemonResult<Metadata> {
let file_size = match req.source_filename.file_name().is_none() {
true => 0_u64,
false => filestore
.get_size(&req.source_filename)
.map_err(DaemonError::SpawnSend)?,
};
Ok(Metadata {
source_filename: req.source_filename,
destination_filename: req.destination_filename,
file_size,
checksum_type: config.checksum_type,
})
}
type RecvSpawnerTuple = (
TransactionID,
Sender<Command>,
JoinHandle<Result<TransactionID, TransactionError>>,
);
type SendSpawnerTuple = (
Sender<Command>,
JoinHandle<Result<TransactionID, TransactionError>>,
);
pub struct Daemon<T: FileStore + Send + 'static> {
transaction_handles: Vec<JoinHandle<Result<TransactionID, TransactionError>>>,
transaction_channels: HashMap<TransactionID, Sender<Command>>,
transport_tx_map: HashMap<EntityID, Sender<(EntityID, PDU)>>,
transport_rx: Receiver<PDU>,
filestore: Arc<T>,
indication_tx: Sender<Indication>,
entity_configs: HashMap<EntityID, EntityConfig>,
default_config: EntityConfig,
entity_id: EntityID,
sequence_num: TransactionSeqNum,
terminate: Arc<AtomicBool>,
primitive_rx: Receiver<UserPrimitive>,
}
impl<T: FileStore + Send + Sync + 'static> Daemon<T> {
#[allow(clippy::too_many_arguments)]
pub fn new(
entity_id: EntityID,
sequence_num: TransactionSeqNum,
transport_map: HashMap<Vec<EntityID>, Box<dyn PDUTransport + Send>>,
filestore: Arc<T>,
entity_configs: HashMap<EntityID, EntityConfig>,
default_config: EntityConfig,
primitive_rx: Receiver<UserPrimitive>,
indication_tx: Sender<Indication>,
) -> Self {
let mut transport_tx_map: HashMap<EntityID, Sender<(EntityID, PDU)>> = HashMap::new();
let (pdu_send, pdu_receive) = channel(100);
let terminate = Arc::new(AtomicBool::new(false));
for (vec, mut transport) in transport_map.into_iter() {
let (remote_send, remote_receive) = channel(1);
vec.iter().for_each(|id| {
transport_tx_map.insert(*id, remote_send.clone());
});
let signal = terminate.clone();
let sender = pdu_send.clone();
tokio::task::spawn(async move {
transport.pdu_handler(signal, sender, remote_receive).await
});
}
Self {
transaction_handles: vec![],
transaction_channels: HashMap::new(),
transport_tx_map,
transport_rx: pdu_receive,
filestore,
indication_tx,
entity_configs,
default_config,
entity_id,
sequence_num,
terminate,
primitive_rx,
}
}
fn spawn_receive_transaction(
header: &PDUHeader,
transport_tx: Sender<(EntityID, PDU)>,
entity_config: EntityConfig,
filestore: Arc<T>,
indication_tx: Sender<Indication>,
) -> RecvSpawnerTuple {
let (transaction_tx, mut transaction_rx) = channel(100);
let config = TransactionConfig {
source_entity_id: header.source_entity_id,
destination_entity_id: header.destination_entity_id,
transmission_mode: header.transmission_mode,
sequence_number: header.transaction_sequence_number,
file_size_flag: header.large_file_flag,
fault_handler_override: entity_config.fault_handler_override.clone(),
file_size_segment: entity_config.file_size_segment,
crc_flag: header.crc_flag,
segment_metadata_flag: header.segment_metadata_flag,
max_count: entity_config.default_transaction_max_count,
inactivity_timeout: entity_config.inactivity_timeout,
eof_timeout: entity_config.eof_timeout,
nak_timeout: entity_config.nak_timeout,
progress_report_interval_secs: entity_config.progress_report_interval_secs,
};
let mut transaction = RecvTransaction::new(
config,
entity_config.nak_procedure,
filestore,
indication_tx,
);
let id = transaction.id();
let handle = tokio::task::spawn(async move {
transaction.send_report(None)?;
while transaction.get_state() != TransactionState::Terminated {
let timeout = transaction.until_timeout();
select! {
Ok(permit) = transport_tx.reserve(), if transaction.has_pdu_to_send() => {
transaction.send_pdu(permit)?
},
Some(command) = transaction_rx.recv() => {
match command {
Command::Pdu(pdu) => {
match transaction.process_pdu(pdu) {
Ok(()) => {}
Err(err @ TransactionError::UnexpectedPDU(..)) => {
info!("Transaction {} Received Unexpected PDU: {err}", transaction.id());
}
Err(err) => return Err(err)
}
}
Command::Abandon => transaction.shutdown(),
Command::Report(sender) => {
transaction.send_report(Some(sender))?
}
}
}
_ = tokio::time::sleep(timeout) => {
transaction.handle_timeout()?;
}
else => {
if transport_tx.is_closed(){
log::error!("Channel to transport unexpectedly severed for transaction {}.", transaction.id());
}
break;
}
};
}
transaction.send_report(None)?;
Ok(transaction.id())
});
(id, transaction_tx, handle)
}
fn spawn_send_transaction(
request: PutRequest,
transaction_id: TransactionID,
transport_tx: Sender<(EntityID, PDU)>,
entity_config: EntityConfig,
filestore: Arc<T>,
indication_tx: Sender<Indication>,
) -> DaemonResult<SendSpawnerTuple> {
let (transaction_tx, mut transaction_rx) = channel(10);
let destination_entity_id = request.destination_entity_id;
let transmission_mode = request.transmission_mode;
let mut config = TransactionConfig {
source_entity_id: transaction_id.0,
destination_entity_id,
transmission_mode,
sequence_number: transaction_id.1,
file_size_flag: FileSizeFlag::Small,
fault_handler_override: entity_config.fault_handler_override.clone(),
file_size_segment: entity_config.file_size_segment,
crc_flag: entity_config.crc_flag,
segment_metadata_flag: SegmentedData::NotPresent,
max_count: entity_config.default_transaction_max_count,
inactivity_timeout: entity_config.inactivity_timeout,
eof_timeout: entity_config.eof_timeout,
nak_timeout: entity_config.nak_timeout,
progress_report_interval_secs: entity_config.progress_report_interval_secs,
};
let metadata = construct_metadata(&filestore, request, entity_config)?;
let handle = tokio::task::spawn(async move {
config.file_size_flag = match metadata.file_size <= u32::MAX.into() {
true => FileSizeFlag::Small,
false => FileSizeFlag::Large,
};
let mut transaction = SendTransaction::new(config, metadata, filestore, indication_tx)?;
transaction.send_report(None)?;
while transaction.get_state() != TransactionState::Terminated {
let timeout = transaction.until_timeout();
select! {
Ok(permit) = transport_tx.reserve(), if transaction.has_pdu_to_send() => {
transaction.send_pdu(permit)?;
},
Some(command) = transaction_rx.recv() => {
match command {
Command::Pdu(pdu) => {
match transaction.process_pdu(pdu) {
Ok(()) => {}
Err(
err @ TransactionError::UnexpectedPDU(..),
) => {
info!("Received Unexpected PDU: {err}");
}
Err(err) => {
return Err(err);
}
}
}
Command::Abandon => transaction.shutdown(),
Command::Report(sender) => {
transaction.send_report(Some(sender))?
},
}
},
_ = tokio::time::sleep(timeout) => {
transaction.handle_timeout()?;
},
else => {
if transport_tx.is_closed(){
log::error!("Connection to transport unexpectedly severed for transaction {}.", transaction.id());
}
break;
}
};
}
transaction.send_report(None)?;
Ok(transaction_id)
});
Ok((transaction_tx, handle))
}
async fn process_primitive(&mut self, primitive: UserPrimitive) -> DaemonResult<()> {
match primitive {
UserPrimitive::Put(request, put_sender) => {
let sequence_number = self.sequence_num.get_and_increment();
let entity_config = self
.entity_configs
.get(&request.destination_entity_id)
.unwrap_or(&self.default_config)
.clone();
if let Some(transport_tx) = self
.transport_tx_map
.get(&request.destination_entity_id)
.cloned()
{
let id = TransactionID(self.entity_id, sequence_number);
let (sender, handle) = Self::spawn_send_transaction(
request,
id,
transport_tx,
entity_config,
self.filestore.clone(),
self.indication_tx.clone(),
)?;
self.transaction_handles.push(handle);
self.transaction_channels.insert(id, sender);
let _ = put_sender.send(id);
} else {
warn!(
"No Transport available for EntityID: {}. Skipping transaction creation.",
request.destination_entity_id
)
}
}
UserPrimitive::Cancel(id) => {
if let Some(channel) = self.transaction_channels.get(&id) {
channel
.send(Command::Abandon)
.await
.map_err(|err| DaemonError::from((id, err)))?;
}
}
UserPrimitive::Report(id, report_sender) => {
if let Some(channel) = self.transaction_channels.get(&id) {
channel
.send(Command::Report(report_sender))
.await
.map_err(|err| DaemonError::from((id, err)))?;
}
}
};
Ok(())
}
async fn forward_pdu(&mut self, pdu: PDU) -> DaemonResult<()> {
let transport_entity = match &pdu.header.direction {
Direction::ToSender => pdu.header.destination_entity_id,
Direction::ToReceiver => pdu.header.source_entity_id,
};
let key = TransactionID(
pdu.header.source_entity_id,
pdu.header.transaction_sequence_number,
);
let channel = match self.transaction_channels.entry(key) {
Entry::Occupied(entry) => entry.into_mut(),
Entry::Vacant(entry) => {
if let Some(transport) = self.transport_tx_map.get(&transport_entity).cloned() {
let entity_config = self
.entity_configs
.get(&key.0)
.unwrap_or(&self.default_config)
.clone();
match &pdu.header.direction {
Direction::ToReceiver => {
let (_id, channel, handle) = Self::spawn_receive_transaction(
&pdu.header,
transport,
entity_config,
self.filestore.clone(),
self.indication_tx.clone(),
);
self.transaction_handles.push(handle);
entry.insert(channel)
}
Direction::ToSender => {
error!("Received PDU sent back to sender but no transaction running. Unable to resume transaction.");
return Err(DaemonError::UnableToResume(TransactionID(
pdu.header.source_entity_id,
pdu.header.transaction_sequence_number,
)));
}
}
} else {
warn!(
"No Transport available for EntityID: {}. Skipping Transaction creation.",
transport_entity
);
return Ok(());
}
}
};
if channel.send(Command::Pdu(pdu.clone())).await.is_err() {
match pdu.header.direction {
Direction::ToReceiver => {
let entity_config = self
.entity_configs
.get(&key.0)
.unwrap_or(&self.default_config)
.clone();
if let Some(transport) = self.transport_tx_map.get(&transport_entity).cloned() {
let (id, new_channel, handle) = Self::spawn_receive_transaction(
&pdu.header,
transport,
entity_config,
self.filestore.clone(),
self.indication_tx.clone(),
);
self.transaction_handles.push(handle);
new_channel
.send(Command::Pdu(pdu.clone()))
.await
.map_err(|err| DaemonError::from((id, err)))?;
self.transaction_channels.insert(key, new_channel);
}
}
Direction::ToSender => {
error!("Received PDU sent back to sender but no transaction running. Unable to resume transaction.");
return Err(DaemonError::UnableToResume(TransactionID(
pdu.header.source_entity_id,
pdu.header.transaction_sequence_number,
)));
}
};
}
Ok(())
}
async fn cleanup_transactions(&mut self) {
let mut ind = 0;
while ind < self.transaction_handles.len() {
if self.transaction_handles[ind].is_finished() {
let handle = self.transaction_handles.remove(ind);
match handle.await {
Ok(Ok(id)) => {
let _ = self.transaction_channels.remove(&id);
}
Ok(Err(err)) => {
info!("Error occurred during transaction: {}", err)
}
Err(_) => error!("Unable to join handle!"),
};
} else {
ind += 1;
}
}
}
pub async fn manage_transactions(&mut self) -> DaemonResult<()> {
let cleanup = {
let mut interval = tokio::time::interval(Duration::from_secs(1));
interval.set_missed_tick_behavior(MissedTickBehavior::Delay);
interval
};
tokio::pin!(cleanup);
loop {
select! {
pdu = self.transport_rx.recv() => match pdu {
Some(pdu) => match self.forward_pdu(pdu).await{
Ok(_) => {},
Err(error @ DaemonError::TransactionCommunication(_, _)) => {
warn!("{error}");
},
Err(err) => {
if !self.terminate.load(Ordering::Relaxed) {
self.terminate.store(true, Ordering::Relaxed);
}
return Err(err);
}
},
None => {
if !self.terminate.load(Ordering::Relaxed) {
error!("Transport unexpectedly disconnected from daemon.");
self.terminate.store(true, Ordering::Relaxed);
}
break;
}
},
primitive = self.primitive_rx.recv() => match primitive {
Some(primitive) => match self.process_primitive(primitive).await{
Ok(_) => {},
Err(error @ DaemonError::SpawnSend(_)) => {
warn!("{error}");
},
Err(error @ DaemonError::TransactionCommunication(_, _)) => {
warn!("{error}");
}
Err(err) => {
if !self.terminate.load(Ordering::Relaxed) {
self.terminate.store(true, Ordering::Relaxed);
}
return Err(err);
}
},
None => {
info!("User triggered daemon shutdown.");
if !self.terminate.load(Ordering::Relaxed) {
self.terminate.store(true, Ordering::Relaxed);
}
break;
}
},
_ = cleanup.tick() => self.cleanup_transactions().await,
};
}
while let Some(handle) = self.transaction_handles.pop() {
match handle.await {
Ok(Ok(id)) => {
let _ = self.transaction_channels.remove(&id);
}
Ok(Err(err)) => {
info!("Error occurred during transaction: {}", err)
}
Err(_) => error!("Unable to join handle!"),
};
}
Ok(())
}
}
#[cfg(test)]
mod test {
use crate::{
daemon::NakProcedure,
filestore::{ChecksumType, NativeFileStore},
pdu::{self, CRCFlag, Condition, NegativeAcknowledgmentPDU, PDUPayload, U3},
transaction::FaultHandlerAction,
};
use super::*;
#[macro_export]
macro_rules! assert_err{
($expression:expr, $($pattern:tt)+) => {
match $expression {
$($pattern)+ => {},
ref e => panic!("expected {} but got {:?}", stringify!($($pattern)+), e)
}
}
}
#[tokio::test]
async fn pdu_to_sender_no_transaction() {
let (_send, recv) = channel(1);
let (indication_tx, _indication_rx) = channel(1);
let (_primitive_tx, primitive_rx) = channel(1);
let filestore = Arc::new(NativeFileStore::new("."));
let mut transport_tx_map = HashMap::<_, _>::new();
let (transport_tx, _) = channel(10);
transport_tx_map.insert(EntityID::from(1_u16), transport_tx);
let mut daemon = Daemon {
transaction_handles: vec![],
transaction_channels: HashMap::<_, _>::new(),
transport_tx_map,
transport_rx: recv,
filestore,
indication_tx,
entity_configs: HashMap::new(),
default_config: EntityConfig {
fault_handler_override: HashMap::from([(
Condition::PositiveLimitReached,
FaultHandlerAction::Abandon,
)]),
file_size_segment: 1024,
default_transaction_max_count: 2,
inactivity_timeout: 0,
eof_timeout: 1,
nak_timeout: 2,
crc_flag: CRCFlag::NotPresent,
checksum_type: ChecksumType::Modular,
nak_procedure: NakProcedure::Deferred(Duration::from_secs(0)),
local_entity_id: 0_u16,
remote_entity_id: 1_u16,
local_server_addr: "127.0.0.1:0",
remote_server_addr: "127.0.0.1:0",
progress_report_interval_secs: 1,
},
entity_id: EntityID::from(0_u16),
sequence_num: TransactionSeqNum::from(0_u32),
terminate: Arc::new(AtomicBool::new(false)),
primitive_rx,
};
let payload = PDUPayload::Directive(pdu::Operations::Nak(NegativeAcknowledgmentPDU {
start_of_scope: 0,
end_of_scope: 1_000_000,
segment_requests: vec![],
}));
let pdu = PDU {
header: PDUHeader {
version: U3::Zero,
pdu_type: pdu::PDUType::FileDirective,
direction: Direction::ToSender,
transmission_mode: pdu::TransmissionMode::Acknowledged,
crc_flag: CRCFlag::NotPresent,
large_file_flag: FileSizeFlag::Small,
pdu_data_field_length: payload.encoded_len(FileSizeFlag::Small),
segmentation_control: pdu::SegmentationControl::NotPreserved,
segment_metadata_flag: SegmentedData::NotPresent,
source_entity_id: EntityID::from(0_u16),
transaction_sequence_number: TransactionSeqNum::from(3_u32),
destination_entity_id: EntityID::from(1_u16),
},
payload,
};
let res = daemon.forward_pdu(pdu).await;
assert_err!(res, Err(DaemonError::UnableToResume(_)))
}
}