use crate::runtime::BlockId;
use crate::runtime::BlockMessage;
use crate::runtime::Error;
use crate::runtime::Pmt;
use crate::runtime::PortId;
use crate::runtime::PortIndex;
use crate::runtime::block_inbox::BlockEndpoint;
use crate::runtime::port_id_matches;
#[derive(Debug)]
struct MessageHandler {
port: PortIndex,
endpoint: BlockEndpoint,
}
#[derive(Debug)]
struct MessageOutput {
name: &'static str,
handlers: Vec<MessageHandler>,
}
impl MessageOutput {
fn new(name: &'static str) -> MessageOutput {
MessageOutput {
name,
handlers: Vec::new(),
}
}
fn name(&self) -> &str {
self.name
}
fn connect(&mut self, port: PortIndex, dst: BlockEndpoint) {
self.handlers.push(MessageHandler {
port,
endpoint: dst,
});
}
async fn post(&mut self, p: Pmt) {
for handler in &self.handlers {
let _ = handler
.endpoint
.send(BlockMessage::Post {
port_id: handler.port,
data: p.clone(),
})
.await;
}
}
}
#[derive(Debug)]
pub struct MessageOutputs {
block_id: BlockId,
outputs: Vec<MessageOutput>,
}
impl MessageOutputs {
pub fn new(block_id: BlockId, outputs: &'static [&'static str]) -> Self {
let outputs = outputs.iter().copied().map(MessageOutput::new).collect();
MessageOutputs { block_id, outputs }
}
pub async fn post(&mut self, id: impl Into<PortId>, p: Pmt) -> Result<(), Error> {
let id = id.into();
let block_id = self.block_id;
self.output_mut(&id)
.ok_or(Error::InvalidMessagePort(block_id, id))?
.post(p)
.await;
Ok(())
}
pub(crate) fn connect(
&mut self,
src_port: &PortId,
dst_block_endpoint: BlockEndpoint,
dst_port: &PortId,
) -> Result<(), Error> {
let block_id = self.block_id;
let PortId::Index(dst_port) = dst_port else {
return Err(Error::InvalidMessagePort(block_id, dst_port.clone()));
};
self.output_mut(src_port)
.ok_or_else(|| Error::InvalidMessagePort(block_id, src_port.clone()))?
.connect(*dst_port, dst_block_endpoint);
Ok(())
}
pub async fn notify_finished(&mut self) {
for o in self.outputs.iter_mut() {
o.post(Pmt::Finished).await;
}
}
fn output_mut(&mut self, port: &PortId) -> Option<&mut MessageOutput> {
self.outputs
.iter_mut()
.enumerate()
.find(|(index, item)| port_id_matches(port, *index, item.name()))
.map(|(_, item)| item)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::runtime::block_inbox::BlockInbox;
use crate::runtime::block_inbox::BlockNotifier;
use crate::runtime::block_on;
use crate::runtime::channel::mpsc::channel;
#[test]
fn handler_sends_through_endpoint() {
let (tx, rx) = channel(1);
let endpoint = BlockInbox::new(tx, BlockNotifier::new()).into();
let mut outputs = MessageOutputs::new(BlockId(0), &["out"]);
outputs
.connect(&PortId::from("out"), endpoint, &PortId::index(0))
.unwrap();
block_on(outputs.post("out", Pmt::U32(7))).unwrap();
assert!(matches!(
rx.try_recv().ok(),
Some(BlockMessage::Post { port_id, data })
if port_id == PortIndex::new(0) && data == Pmt::U32(7)
));
}
#[test]
fn handler_accepts_indexed_ports() {
let (tx, rx) = channel(1);
let endpoint = BlockInbox::new(tx, BlockNotifier::new()).into();
let mut outputs = MessageOutputs::new(BlockId(0), &["out"]);
outputs
.connect(&PortId::index(0), endpoint, &PortId::index(0))
.unwrap();
block_on(outputs.post(PortId::index(0), Pmt::U32(7))).unwrap();
assert!(matches!(
rx.try_recv().ok(),
Some(BlockMessage::Post { port_id, data })
if port_id == PortIndex::new(0) && data == Pmt::U32(7)
));
}
#[test]
fn post_invalid_port_reports_block_id() {
let mut outputs = MessageOutputs::new(BlockId(7), &["out"]);
let result = block_on(outputs.post("missing", Pmt::U32(7)));
assert!(matches!(
result,
Err(Error::InvalidMessagePort(block_id, port))
if block_id == BlockId(7) && port == PortId::from("missing")
));
}
}