const MAX_FRAME_LENGTH: usize = 16 * 1024 * 1024;
use std::{collections::HashMap, net::SocketAddr, path::PathBuf};
use bytes::Bytes;
use tokio::{
io::AsyncWriteExt,
net::{
tcp::{OwnedReadHalf, OwnedWriteHalf},
TcpListener,
},
select,
sync::mpsc::{channel, error::SendError, Receiver, Sender},
task::JoinHandle,
};
use tokio_stream::StreamExt;
use tokio_util::{
codec::{FramedRead, LengthDelimitedCodec},
sync::CancellationToken,
};
use crate::{
hex8,
msg::{self, Addr, EncodedMessage, YgwMessage},
protobuf::{self, ygw::MessageType},
recorder::Recorder,
replay_server::start_replay_server,
Link, Result, YgwError, YgwLinkNodeProperties, YgwNode,
};
pub enum CtrlMessage {
NewYamcsConnection(YamcsConnection),
YamcsConnectionClosed(SocketAddr),
}
pub struct Server {
nodes: Vec<Box<dyn YgwNode>>,
addr: SocketAddr,
record_replay_conf: Option<(PathBuf, SocketAddr)>,
}
pub struct ServerBuilder {
nodes: Vec<Box<dyn YgwNode>>,
addr: SocketAddr,
record_replay_conf: Option<(PathBuf, SocketAddr)>,
}
impl Default for ServerBuilder {
fn default() -> Self {
Self::new()
}
}
impl ServerBuilder {
pub fn new() -> Self {
Self {
addr: ([127, 0, 0, 1], 7897).into(),
nodes: Vec::new(),
record_replay_conf: None,
}
}
pub fn set_addr(mut self, addr: SocketAddr) -> Self {
self.addr = addr;
self
}
pub fn add_node(mut self, node: Box<dyn YgwNode>) -> Self {
self.nodes.push(node);
self
}
pub fn with_record_replay_conf(
mut self,
record_replay_conf: Option<(PathBuf, SocketAddr)>,
) -> Self {
self.record_replay_conf = record_replay_conf;
self
}
pub fn build(self) -> Server {
Server {
nodes: self.nodes,
addr: self.addr,
record_replay_conf: self.record_replay_conf,
}
}
}
pub struct ServerHandle {
pub addr: SocketAddr,
pub jh: JoinHandle<Result<()>>,
pub cancel_token: CancellationToken,
}
impl Server {
pub async fn start(mut self) -> Result<ServerHandle> {
let mut node_tx_map = HashMap::new();
let mut node_data = HashMap::new();
let mut node_id = 0;
let (encoder_tx, encoder_rx) = tokio::sync::mpsc::channel(100);
let socket = TcpListener::bind(self.addr)
.await
.map_err(|e| YgwError::IOError(format!("Cannot bind to {}", self.addr), e))?;
let addr = socket.local_addr()?;
let cancel_token = CancellationToken::new();
for node in self.nodes.drain(..) {
let props = node.properties();
let node_name = props.name.clone();
let (tx, rx) = tokio::sync::mpsc::channel(100);
node_data.insert(node_id, NodeData::new(node_id, props, node.sub_links()));
let encoder_tx = encoder_tx.clone();
let node_name_clone = node_name.clone();
log::info!("Starting node {} with id {}", node_name, node_id);
tokio::spawn(async move {
match node.run(node_id, encoder_tx, rx).await {
Ok(()) => {
log::info!("Node {} exited normally", node_name_clone);
}
Err(e) => {
println!("Node {} exited with error: {:?}", node_name_clone, e);
log::error!("Node {} exited with error: {:?}", node_name_clone, e);
}
}
});
node_tx_map.insert(node_id, tx);
node_id += 1;
}
let (ctrl_tx, ctrl_rx) = channel(10);
let (decoder_tx, decoder_rx) = tokio::sync::mpsc::channel(100);
let cancel_token2 = cancel_token.clone();
let accepter_jh =
tokio::spawn(
async move { accepter_task(ctrl_tx, socket, decoder_tx, cancel_token2).await },
);
let cancel_token2 = cancel_token.clone();
let cancel_token3 = cancel_token.clone();
let encoder_jh = tokio::spawn(async move {
if let Err(e) = encoder_task(
ctrl_rx,
encoder_rx,
node_data,
self.record_replay_conf,
cancel_token2,
)
.await
{
log::error!("Encoder task failed: {:?}", e);
cancel_token3.cancel();
Err(e)
} else {
Ok(())
}
});
let decoder_jh = tokio::spawn(async move { decoder_task(decoder_rx, node_tx_map).await });
let jh: JoinHandle<Result<()>> = tokio::spawn(async move {
let (res1, res2, res3) =
futures::future::join3(accepter_jh, encoder_jh, decoder_jh).await;
res1.map_err(|e| YgwError::from(e))
.and(res2.map_err(|e| YgwError::from(e)))
.and(res3.map_err(|e| YgwError::from(e)))
.map(|_| ())
});
Ok(ServerHandle {
jh,
addr,
cancel_token,
})
}
}
impl ServerHandle {
pub async fn run(self) -> Result<()> {
let ServerHandle {
jh, cancel_token, ..
} = self;
tokio::spawn(async move {
#[cfg(unix)]
{
use tokio::signal::unix::{signal, SignalKind};
let mut sigint = signal(SignalKind::interrupt())?;
let mut sigterm = signal(SignalKind::terminate())?;
tokio::select! {
_ = sigint.recv() => log::info!("Received SIGINT"),
_ = sigterm.recv() => log::info!("Received SIGTERM"),
}
}
#[cfg(windows)]
{
tokio::signal::ctrl_c().await?;
log::info!("Received Ctrl+C");
}
cancel_token.cancel();
Ok::<(), std::io::Error>(())
});
let server_result = match jh.await {
Ok(result) => result,
Err(e) => Err(e.into()),
};
server_result
}
}
#[derive(Debug)]
pub struct YamcsConnection {
addr: SocketAddr,
writer_jh: JoinHandle<Result<()>>,
reader_jh: JoinHandle<Result<()>>,
chan_tx: Sender<EncodedMessage>,
drop_if_full: bool,
}
impl PartialEq for YamcsConnection {
fn eq(&self, other: &Self) -> bool {
self.addr == other.addr
}
}
struct NodeData {
node_id: u32,
props: YgwLinkNodeProperties,
links: Vec<Link>,
para_defs: protobuf::ygw::ParameterDefinitionList,
cmd_defs: protobuf::ygw::CommandDefinitionList,
cmd_opts: protobuf::ygw::CommandOptionList,
para_values: HashMap<String, protobuf::ygw::ParameterData>,
link_status: HashMap<u32, protobuf::ygw::LinkStatus>,
}
impl NodeData {
fn new(node_id: u32, props: &YgwLinkNodeProperties, links: &[Link]) -> Self {
Self {
node_id,
props: props.clone(),
links: links.to_vec(),
para_defs: protobuf::ygw::ParameterDefinitionList {
definitions: Vec::new(),
},
cmd_defs: protobuf::ygw::CommandDefinitionList {
definitions: Vec::new(),
},
cmd_opts: protobuf::ygw::CommandOptionList {
options: Vec::new(),
},
para_values: HashMap::new(),
link_status: HashMap::new(),
}
}
fn node_to_proto(&self) -> protobuf::ygw::Node {
protobuf::ygw::Node {
id: self.node_id,
name: self.props.name.clone(),
description: Some(self.props.description.clone()),
tm_packet: if self.props.tm_packet {
Some(true)
} else {
None
},
tm_frame: if self.props.tm_frame {
Some(true)
} else {
None
},
tc: if self.props.tc { Some(true) } else { None },
tc_frame: if self.props.tc_frame {
Some(true)
} else {
None
},
links: self.links.iter().map(|l| l.to_proto()).collect(),
}
}
}
async fn encoder_task(
mut ctrl_rx: Receiver<CtrlMessage>,
mut encoder_rx: Receiver<YgwMessage>,
mut nodes: HashMap<u32, NodeData>,
recorder_replay_conf: Option<(PathBuf, SocketAddr)>,
cancel_token: CancellationToken,
) -> Result<()> {
let mut connections: Vec<YamcsConnection> = Vec::new();
let mut rn = 0;
let recorder_tx: Option<Sender<EncodedMessage>> = match recorder_replay_conf {
None => None,
Some((dir, replay_addr)) => {
log::info!(
"Encoder: starting recorder with recording directory {}",
dir.display()
);
let (mut recorder, last_rn) = Recorder::new(&dir)?;
if let Some(last_rn) = last_rn {
rn = last_rn + 1;
}
let (recorder_tx, recorder_rx) = tokio::sync::mpsc::channel(100);
let (query_tx, query_rx) = tokio::sync::mpsc::channel(16);
tokio::spawn(async move {
if let Err(e) = recorder.record(recorder_rx, query_rx).await {
log::error!("Recorder exited with error: {:?}", e);
}
});
log::info!(
"Encoder: starting the replay server listening on {}",
replay_addr
);
tokio::spawn(async move {
if let Err(e) = start_replay_server(replay_addr, query_tx, cancel_token).await {
log::error!("Replay server exited with error: {:?}", e);
}
});
Some(recorder_tx)
}
};
let mut ctrl_select = true;
loop {
select! {
msg = encoder_rx.recv() => {
match msg {
Some(msg) => {
rn+=1;
let enc_msg = msg.encode(rn);
if let Some(ref recorder_tx) = recorder_tx {
if let Err(e) = recorder_tx.send(enc_msg.clone()).await {
log::warn!("Error sending data to recorder: {:?}", e);
}
}
send_data_to_all(&mut connections, enc_msg).await;
match msg {
YgwMessage::ParameterDefinitions(addr, pdefs) => {
if let Some(node) = nodes.get_mut(&addr.node_id()) {
for def in pdefs.definitions {
if let Some(other_with_same_id) = node.para_defs.definitions.iter().find(|x| x.relative_name != def.relative_name && x.id == def.id) {
log::warn!("Parameter {} with ID {} would override existing parameter {} with same ID, not updating it", def.relative_name, def.id, other_with_same_id.relative_name);
continue;
}
if let Some(pos) = node.para_defs.definitions.iter().position(|x| x.relative_name == def.relative_name) {
node.para_defs.definitions[pos] = def;
} else {
node.para_defs.definitions.push(def);
}
}
}
},
YgwMessage::ParameterData(addr, pvals) => {
if let Some(node) = nodes.get_mut(&addr.node_id()) {
node.para_values.insert(pvals.group.clone(), pvals);
}
},
YgwMessage::LinkStatus(addr, link_status) => {
if let Some(node) = nodes.get_mut(&addr.node_id()) {
node.link_status.insert(addr.link_id(), link_status);
}
},
YgwMessage::CommandDefinitions(addr, cmd_defs) => {
if let Some(node) = nodes.get_mut(&addr.node_id()) {
for def in cmd_defs.definitions {
if let Some(other_with_same_id) = node.cmd_defs.definitions.iter().find(|x| x.relative_name != def.relative_name && x.ygw_cmd_id == def.ygw_cmd_id) {
log::warn!("Command {} with ID {} would override existing command {} with same ID, not updating it", def.relative_name, def.ygw_cmd_id, other_with_same_id.relative_name);
continue;
}
if let Some(pos) = node.cmd_defs.definitions.iter().position(|x| x.relative_name == def.relative_name) {
node.cmd_defs.definitions[pos] = def;
} else {
node.cmd_defs.definitions.push(def);
}
}
}
},
YgwMessage::CommandOptions(addr, cmd_opts) => {
if let Some(node) = nodes.get_mut(&addr.node_id()) {
node.cmd_opts.options.extend(cmd_opts.options);
}
},
_ => {}
}
},
None => {
log::debug!("Encoder: channel from nodes closed");
break
}
}
}
msg = ctrl_rx.recv(), if ctrl_select => {
match msg {
Some(CtrlMessage::NewYamcsConnection(yc)) => {
if let Err(_)= send_initial_data(&yc, &nodes).await {
log::warn!("Encoder: error sending initial data message to {}", yc.addr);
continue;
}
connections.push(yc);
},
Some(CtrlMessage::YamcsConnectionClosed(addr)) => connections.retain(|yc| yc.addr != addr),
None => {
log::debug!("Encoder: channel from accepter closed, waiting for all nodes to quit");
ctrl_select = false;
},
}
}
}
}
log::debug!("Encoder task exiting");
Ok(())
}
async fn send_data_to_all(connections: &mut Vec<YamcsConnection>, msg: EncodedMessage) {
let mut idx = 0;
while idx < connections.len() {
let msg1 = msg.clone();
let yc = &connections[idx];
if yc.drop_if_full {
if let Err(_) = yc.chan_tx.try_send(msg1) {
log::warn!("Channel to {} is full, dropping connection", yc.addr);
yc.reader_jh.abort();
yc.writer_jh.abort();
connections.remove(idx);
continue;
}
} else if let Err(_) = yc.chan_tx.send(msg1).await {
connections.remove(idx);
continue;
}
idx += 1;
}
}
async fn send_initial_data(
yc: &YamcsConnection,
nodes: &HashMap<u32, NodeData>,
) -> std::result::Result<(), SendError<EncodedMessage>> {
let nl = protobuf::ygw::NodeList {
nodes: nodes.iter().map(|(_, nd)| nd.node_to_proto()).collect(),
};
let buf = msg::encode_node_info(&nl);
yc.chan_tx.send(buf).await?;
for nd in nodes.values() {
if !nd.para_defs.definitions.is_empty() {
let buf = msg::encode_message(
0,
&Addr::new(nd.node_id, 0),
MessageType::ParameterDefinitions,
&nd.para_defs,
);
yc.chan_tx.send(buf).await?;
}
}
for nd in nodes.values() {
if !nd.cmd_defs.definitions.is_empty() {
let buf = msg::encode_message(
0,
&Addr::new(nd.node_id, 0),
MessageType::CommandDefinitions,
&nd.cmd_defs,
);
yc.chan_tx.send(buf).await?;
}
}
for nd in nodes.values() {
if !nd.cmd_opts.options.is_empty() {
let buf = msg::encode_message(
0,
&Addr::new(nd.node_id, 0),
MessageType::CommandOptions,
&nd.cmd_opts,
);
yc.chan_tx.send(buf).await?;
}
}
for nd in nodes.values() {
for pdata in nd.para_values.values() {
let buf = msg::encode_message(
0,
&Addr::new(nd.node_id, 0),
MessageType::ParameterData,
pdata,
);
yc.chan_tx.send(buf).await?;
}
}
for nd in nodes.values() {
for (&link_id, lstatus) in nd.link_status.iter() {
let buf = msg::encode_message(
0,
&Addr::new(nd.node_id, link_id),
MessageType::LinkStatus,
lstatus,
);
yc.chan_tx.send(buf).await?;
}
}
Ok(())
}
async fn decoder_task(
mut decoder_rx: Receiver<Bytes>,
mut nodes: HashMap<u32, Sender<YgwMessage>>,
) -> Result<()> {
loop {
match decoder_rx.recv().await {
Some(mut buf) => match YgwMessage::decode(&mut buf) {
Ok(msg) => {
let node_id = msg.node_id();
match nodes.get(&node_id) {
Some(tx) => {
if let Err(_) = tx.send(msg).await {
log::warn!("Channel to node {} closed", node_id);
nodes.remove(&node_id);
}
}
None => {
log::warn!("Received message for unknown node {} ", node_id);
}
}
}
Err(err) => log::warn!("Cannot decode data {}: {:?}", hex8(&buf), err),
},
None => break,
};
}
log::debug!("Decoder task exiting");
Ok(())
}
async fn accepter_task(
ctrl_tx: Sender<CtrlMessage>,
srv_sock: TcpListener,
decoder_tx: Sender<Bytes>,
cancel_token: CancellationToken,
) -> Result<()> {
loop {
tokio::select! {
res = srv_sock.accept() => {
match res {
Ok((sock, addr)) => {
log::info!("New Yamcs connection from {}", addr);
let (read_sock, write_sock) = sock.into_split();
let (chan_tx, chan_rx) = channel(100);
let decoder_tx2 = decoder_tx.clone();
let ctrl_tx2 = ctrl_tx.clone();
let cancel_token2 = cancel_token.clone();
let reader_jh = tokio::spawn(async move { reader_task(ctrl_tx2, addr, read_sock, decoder_tx2, cancel_token2).await });
let writer_jh = tokio::spawn(async move { writer_task(write_sock, chan_rx).await });
let yc = YamcsConnection {
addr,
reader_jh,
writer_jh,
chan_tx,
drop_if_full: false,
};
if let Err(_) = ctrl_tx.send(CtrlMessage::NewYamcsConnection(yc)).await {
break;
}
}
Err(err) => {
log::error!("Failed to accept connection: {}", err);
break;
}
}
},
_ = cancel_token.cancelled() => {
log::debug!("Accepter task received cancel signal.");
break;
}
}
}
log::debug!("Accepter task exiting");
Ok(())
}
async fn reader_task(
ctrl_tx: Sender<CtrlMessage>,
addr: SocketAddr,
read_sock: OwnedReadHalf,
decoder_tx: Sender<Bytes>,
cancel_token: CancellationToken,
) -> Result<()> {
let mut codec = LengthDelimitedCodec::new();
codec.set_max_frame_length(MAX_FRAME_LENGTH);
let mut stream = FramedRead::new(read_sock, codec);
loop {
select! {
result = stream.next() => {
match result {
Some(Ok(buf)) => {
let buf = buf.freeze();
log::trace!("Received message {:}", hex8(&buf));
if let Err(_) = decoder_tx.send(buf).await {
break;
}
}
Some(Err(e)) => {
log::warn!("Error reading from {}: {:?}", addr, e);
let _ = ctrl_tx.send(CtrlMessage::YamcsConnectionClosed(addr)).await;
return Err(YgwError::IOError(format!("Error reading from {addr}"), e));
}
None => {
log::info!("Yamcs connection {} closed", addr);
let _ = ctrl_tx.send(CtrlMessage::YamcsConnectionClosed(addr)).await;
break;
}
}
}
_ = cancel_token.cancelled() => {
log::debug!("Reader task for {} received cancel signal", addr);
break;
}
}
}
Ok(())
}
async fn writer_task(mut sock: OwnedWriteHalf, mut chan: Receiver<EncodedMessage>) -> Result<()> {
loop {
match chan.recv().await {
Some(msg) => {
sock.write_all(&msg).await?;
}
None => break,
}
}
log::debug!("Writer task exiting");
Ok(())
}
#[cfg(test)]
mod tests {
use std::{
io::ErrorKind,
time::{Duration, Instant},
};
use async_trait::async_trait;
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::TcpStream,
sync::mpsc,
};
use tokio_util::codec::Framed;
use crate::{
msg::{Addr, TmPacket},
protobuf::{ygw::CommandId, ygw::PreparedCommand},
Link,
};
use super::*;
#[tokio::test]
async fn test_frame_too_long() {
let (addr, _node_id, _node_tx, _node_rx) = setup_test().await;
let mut conn = TcpStream::connect(addr).await.unwrap();
conn.write_u32((MAX_FRAME_LENGTH + 1) as u32).await.unwrap();
let mut buf = vec![0; 1024];
let _ = conn.read_buf(&mut buf).await.unwrap();
let r = conn.read_u32().await.unwrap_err();
assert_eq!(ErrorKind::UnexpectedEof, r.kind());
}
#[tokio::test]
async fn test_tm() {
let (addr, node_id, node_tx, _node_rx) = setup_test().await;
let conn = TcpStream::connect(addr).await.unwrap();
let mut stream = Framed::new(conn, LengthDelimitedCodec::new());
tokio::task::yield_now().await;
node_tx
.send(YgwMessage::TmPacket(
Addr::new(node_id, 0),
TmPacket {
data: vec![1, 2, 3, 7],
acq_time: protobuf::now(),
},
))
.await
.unwrap();
let _ = stream.next().await.unwrap().unwrap();
let buf2 = stream.next().await.unwrap().unwrap();
assert_eq!(34, buf2.len());
assert_eq!([1, 2, 3, 7], buf2[30..34]);
}
#[tokio::test]
async fn test_tc() {
env_logger::init();
let (addr, node_id, _node_tx, mut node_rx) = setup_test().await;
let mut conn = TcpStream::connect(addr).await.unwrap();
let pc = prepared_cmd();
let msg = YgwMessage::Tc(Addr::new(node_id, 0), pc.clone());
let enc_msg = msg.encode(0);
conn.write_all(&enc_msg).await.unwrap();
let msg1 = node_rx.recv().await.unwrap();
assert_eq!(msg, msg1);
}
async fn _test_performance() {
let (addr, node_id, node_tx, _node_rx) = setup_test().await;
let conn = TcpStream::connect(addr).await.unwrap();
let mut stream = Framed::new(conn, LengthDelimitedCodec::new());
let n = 1_000_000;
let client_handle = tokio::spawn(async move {
let mut count = 0;
let mut t0 = Instant::now();
while let Some(Ok(_)) = stream.next().await {
if count == 0 {
t0 = Instant::now();
}
count += 1;
if count == n {
break;
}
}
let d = t0.elapsed();
println!(
"Received {} messages in {:?}: speed {:.2} msg/millisec",
count,
d,
(count as f64) / (d.as_millis() as f64)
);
});
tokio::time::sleep(Duration::from_secs(1)).await;
tokio::spawn(async move {
let t0 = Instant::now();
for _ in 0..n {
node_tx
.send(YgwMessage::TmPacket(
Addr::new(node_id, 0),
TmPacket {
data: vec![0; 1024],
acq_time: protobuf::now(),
},
))
.await
.unwrap();
}
let d = t0.elapsed();
println!(
"Sent {} messages; speed {:.2} msg/millisec {} nanosec/message",
n,
(n as f64) / (d.as_millis() as f64),
d.as_nanos() / n
);
})
.await
.unwrap();
client_handle.await.unwrap();
}
async fn setup_test() -> (SocketAddr, u32, Sender<YgwMessage>, Receiver<YgwMessage>) {
let (tx, mut rx) = mpsc::channel(1);
let props = YgwLinkNodeProperties::new("test_node", "test node")
.tm_packet(true)
.tc(true);
let dn = DummyNode { props, tx };
let addr = ([127, 0, 0, 1], 0).into();
let server = ServerBuilder::new()
.set_addr(addr)
.add_node(Box::new(dn))
.build();
let server_handle = server.start().await.unwrap();
let x = rx.recv().await.unwrap();
(server_handle.addr, x.0, x.1, x.2)
}
fn prepared_cmd() -> PreparedCommand {
PreparedCommand {
command_id: CommandId {
generation_time: 100,
origin: String::from("test"),
sequence_number: 10,
command_name: None,
},
assignments: Vec::new(),
extra: HashMap::new(),
binary: Some(vec![1, 2, 3]),
ygw_cmd_id: Some(10),
}
}
struct DummyNode {
props: YgwLinkNodeProperties,
tx: mpsc::Sender<(u32, Sender<YgwMessage>, Receiver<YgwMessage>)>,
}
#[async_trait]
impl YgwNode for DummyNode {
fn properties(&self) -> &YgwLinkNodeProperties {
&self.props
}
fn sub_links(&self) -> &[Link] {
&[]
}
async fn run(
self: Box<Self>,
node_id: u32,
tx: Sender<YgwMessage>,
rx: Receiver<YgwMessage>,
) -> Result<()> {
self.tx.send((node_id, tx, rx)).await.unwrap();
tokio::time::sleep(std::time::Duration::from_secs(200)).await;
Ok(())
}
}
}