use crate::symmetric::CollectiveError;
use crate::transport::Transport;
use std::collections::{HashMap, VecDeque};
use std::str::FromStr;
use std::sync::{Arc, Condvar, Mutex};
use std::time::{Duration, Instant};
use iroh::endpoint::presets;
use iroh::{Endpoint, EndpointAddr, EndpointId, RelayUrl, SecretKey};
use tokio::runtime::Runtime;
pub use iroh::{RelayMap, RelayMode};
pub const RLX_PIPELINE_ALPN: &[u8] = b"rlx-pipeline/1";
const MAX_FRAME_BYTES: usize = 512 * 1024 * 1024;
const DIAL_TIMEOUT: Duration = Duration::from_secs(60);
fn terr(reason: impl Into<String>) -> CollectiveError {
CollectiveError::TransportError {
reason: reason.into(),
}
}
#[derive(Clone, Debug)]
pub struct IrohPeer {
id: EndpointId,
relay_url: Option<RelayUrl>,
direct_addrs: Vec<std::net::SocketAddr>,
}
impl IrohPeer {
pub fn new(id: EndpointId) -> Self {
Self {
id,
relay_url: None,
direct_addrs: Vec::new(),
}
}
pub fn from_id_str(s: &str) -> Result<Self, CollectiveError> {
let id = EndpointId::from_str(s.trim())
.map_err(|e| terr(format!("invalid endpoint id {s:?}: {e}")))?;
Ok(Self::new(id))
}
pub fn with_relay(mut self, url: &str) -> Result<Self, CollectiveError> {
self.relay_url = Some(
RelayUrl::from_str(url).map_err(|e| terr(format!("invalid relay url {url:?}: {e}")))?,
);
Ok(self)
}
pub fn with_direct(mut self, addr: std::net::SocketAddr) -> Self {
self.direct_addrs.push(addr);
self
}
fn to_addr(&self) -> EndpointAddr {
let mut addr = EndpointAddr::new(self.id);
if let Some(relay) = &self.relay_url {
addr = addr.with_relay_url(relay.clone());
}
for &d in &self.direct_addrs {
addr = addr.with_ip_addr(d);
}
addr
}
}
type Mailbox = Arc<(Mutex<HashMap<(u32, u32), VecDeque<Vec<u8>>>>, Condvar)>;
pub struct IrohTransport {
rank: u32,
world: u32,
endpoint_id: EndpointId,
peers: Vec<IrohPeer>,
alpn: Vec<u8>,
rt: Arc<Runtime>,
endpoint: Endpoint,
out: Mutex<HashMap<u32, Arc<OutEdge>>>,
mailbox: Mailbox,
}
struct OutEdge {
#[allow(dead_code)] conn: iroh::endpoint::Connection,
stream: Mutex<iroh::endpoint::SendStream>,
}
impl IrohTransport {
pub fn connect(
rank: u32,
world: u32,
secret_key: SecretKey,
peers: Vec<IrohPeer>,
alpn: &[u8],
) -> Result<Self, CollectiveError> {
Self::connect_with(rank, world, secret_key, peers, alpn, RelayMode::Disabled)
}
pub fn connect_relayed(
rank: u32,
world: u32,
secret_key: SecretKey,
peers: Vec<IrohPeer>,
alpn: &[u8],
) -> Result<Self, CollectiveError> {
Self::connect_with(rank, world, secret_key, peers, alpn, RelayMode::Default)
}
pub fn connect_with(
rank: u32,
world: u32,
secret_key: SecretKey,
peers: Vec<IrohPeer>,
alpn: &[u8],
relay_mode: RelayMode,
) -> Result<Self, CollectiveError> {
Self::connect_bind(rank, world, secret_key, peers, alpn, false, relay_mode)
}
pub fn connect_discovered(
rank: u32,
world: u32,
secret_key: SecretKey,
peers: Vec<IrohPeer>,
alpn: &[u8],
) -> Result<Self, CollectiveError> {
Self::connect_bind(
rank,
world,
secret_key,
peers,
alpn,
true,
RelayMode::Default,
)
}
fn connect_bind(
rank: u32,
world: u32,
secret_key: SecretKey,
peers: Vec<IrohPeer>,
alpn: &[u8],
use_n0: bool,
relay_mode: RelayMode,
) -> Result<Self, CollectiveError> {
if peers.len() != world as usize {
return Err(terr(format!(
"peers must have world_size ({world}) entries, got {}",
peers.len()
)));
}
let rt = Arc::new(
tokio::runtime::Builder::new_multi_thread()
.enable_all()
.worker_threads(2)
.thread_name("rlx-iroh")
.build()
.map_err(|e| terr(format!("tokio runtime: {e}")))?,
);
let alpn_v = alpn.to_vec();
let endpoint = rt
.block_on(async {
if use_n0 {
Endpoint::builder(presets::N0)
.secret_key(secret_key)
.alpns(vec![alpn_v.clone()])
.bind()
.await
} else {
Endpoint::builder(presets::Minimal)
.secret_key(secret_key)
.alpns(vec![alpn_v.clone()])
.relay_mode(relay_mode)
.bind()
.await
}
})
.map_err(|e| terr(format!("iroh bind: {e}")))?;
let endpoint_id = endpoint.id();
let mailbox: Mailbox = Arc::new((Mutex::new(HashMap::new()), Condvar::new()));
{
let ep = endpoint.clone();
let mb = mailbox.clone();
rt.spawn(async move {
while let Some(incoming) = ep.accept().await {
let mb = mb.clone();
tokio::spawn(async move {
let conn = match incoming.await {
Ok(c) => c,
Err(_) => return,
};
reader_loop(conn, mb).await;
});
}
});
}
Ok(Self {
rank,
world,
endpoint_id,
peers,
alpn: alpn_v,
rt,
endpoint,
out: Mutex::new(HashMap::new()),
mailbox,
})
}
pub fn connect_ephemeral(
rank: u32,
world: u32,
peers: Vec<IrohPeer>,
alpn: &[u8],
) -> Result<Self, CollectiveError> {
Self::connect(rank, world, SecretKey::generate(), peers, alpn)
}
pub fn endpoint_id(&self) -> EndpointId {
self.endpoint_id
}
pub fn endpoint_id_string(&self) -> String {
self.endpoint_id.to_string()
}
fn edge_to(&self, to: u32) -> Result<Arc<OutEdge>, CollectiveError> {
if let Some(e) = self.out.lock().unwrap().get(&to) {
return Ok(e.clone());
}
let addr = self.peers[to as usize].to_addr();
let deadline = Instant::now() + DIAL_TIMEOUT;
let (conn, stream) = loop {
let attempt = self.rt.block_on(async {
let conn = self
.endpoint
.connect(addr.clone(), &self.alpn)
.await
.map_err(|e| terr(e.to_string()))?;
let stream = conn.open_uni().await.map_err(|e| terr(e.to_string()))?;
Ok::<_, CollectiveError>((conn, stream))
});
match attempt {
Ok(cs) => break cs,
Err(e) => {
if Instant::now() >= deadline {
return Err(terr(format!("dial/open rank {to} failed: {e}")));
}
std::thread::sleep(Duration::from_millis(200));
}
}
};
let edge = Arc::new(OutEdge {
conn,
stream: Mutex::new(stream),
});
self.out.lock().unwrap().insert(to, edge.clone());
Ok(edge)
}
fn evict(&self, to: u32) {
self.out.lock().unwrap().remove(&to);
}
}
async fn reader_loop(conn: iroh::endpoint::Connection, mailbox: Mailbox) {
let mut recv = match conn.accept_uni().await {
Ok(r) => r,
Err(_) => return, };
let mut header = [0u8; 12];
loop {
if recv.read_exact(&mut header).await.is_err() {
return; }
let from = u32::from_le_bytes([header[0], header[1], header[2], header[3]]);
let tag = u32::from_le_bytes([header[4], header[5], header[6], header[7]]);
let len = u32::from_le_bytes([header[8], header[9], header[10], header[11]]) as usize;
if len > MAX_FRAME_BYTES {
return; }
let mut payload = vec![0u8; len];
if len > 0 && recv.read_exact(&mut payload).await.is_err() {
return;
}
let (lock, cv) = &*mailbox;
lock.lock()
.unwrap()
.entry((from, tag))
.or_default()
.push_back(payload);
cv.notify_all();
}
}
impl Transport for IrohTransport {
fn rank(&self) -> u32 {
self.rank
}
fn world_size(&self) -> u32 {
self.world
}
fn send_bytes(&self, to: u32, tag: u32, bytes: &[u8]) -> Result<(), CollectiveError> {
if to == self.rank {
let (lock, cv) = &*self.mailbox;
lock.lock()
.unwrap()
.entry((self.rank, tag))
.or_default()
.push_back(bytes.to_vec());
cv.notify_all();
return Ok(());
}
if bytes.len() > MAX_FRAME_BYTES {
return Err(terr(format!(
"frame of {} bytes exceeds MAX_FRAME_BYTES ({MAX_FRAME_BYTES})",
bytes.len()
)));
}
let mut header = [0u8; 12];
header[0..4].copy_from_slice(&self.rank.to_le_bytes());
header[4..8].copy_from_slice(&tag.to_le_bytes());
header[8..12].copy_from_slice(&(bytes.len() as u32).to_le_bytes());
for attempt in 0..2 {
let edge = self.edge_to(to)?;
let mut stream = edge.stream.lock().unwrap();
let res = self.rt.block_on(async {
stream
.write_all(&header)
.await
.map_err(|e| terr(e.to_string()))?;
stream
.write_all(bytes)
.await
.map_err(|e| terr(e.to_string()))?;
Ok::<(), CollectiveError>(())
});
drop(stream);
match res {
Ok(()) => return Ok(()),
Err(e) => {
self.evict(to);
if attempt == 1 {
return Err(e);
}
}
}
}
Ok(())
}
fn recv_bytes(&self, from: u32, tag: u32) -> Result<Vec<u8>, CollectiveError> {
let (lock, cv) = &*self.mailbox;
let mut guard = lock.lock().unwrap();
loop {
if let Some(q) = guard.get_mut(&(from, tag))
&& let Some(v) = q.pop_front()
{
return Ok(v);
}
guard = cv.wait(guard).unwrap();
}
}
fn recv_bytes_timeout(
&self,
from: u32,
tag: u32,
timeout: Duration,
) -> Result<Option<Vec<u8>>, CollectiveError> {
let (lock, cv) = &*self.mailbox;
let mut guard = lock.lock().unwrap();
let deadline = Instant::now() + timeout;
loop {
if let Some(q) = guard.get_mut(&(from, tag))
&& let Some(v) = q.pop_front()
{
return Ok(Some(v));
}
let now = Instant::now();
if now >= deadline {
return Ok(None);
}
let (g, res) = cv.wait_timeout(guard, deadline - now).unwrap();
guard = g;
if res.timed_out()
&& guard
.get_mut(&(from, tag))
.map(|q| q.is_empty())
.unwrap_or(true)
{
return Ok(None);
}
}
}
}
impl Drop for IrohTransport {
fn drop(&mut self) {
self.out.lock().unwrap().clear();
let ep = self.endpoint.clone();
let _ = self.rt.block_on(async move {
tokio::time::timeout(Duration::from_secs(2), ep.close()).await
});
}
}
fn hex_to_bytes(s: &str) -> Result<Vec<u8>, CollectiveError> {
let s = s.trim();
if !s.len().is_multiple_of(2) {
return Err(terr("hex string must have even length"));
}
(0..s.len())
.step_by(2)
.map(|i| {
u8::from_str_radix(&s[i..i + 2], 16)
.map_err(|_| terr(format!("bad hex byte {:?}", &s[i..i + 2])))
})
.collect()
}
fn derive_key(seed: &[u8], rank: u32) -> [u8; 32] {
let mut k = [0u8; 32];
for (i, b) in k.iter_mut().enumerate() {
*b = seed[i % seed.len()];
}
let rb = rank.to_le_bytes();
for i in 0..4 {
k[i] ^= rb[i];
k[28 + i] ^= rb[i].rotate_left(3).wrapping_add(0x5a);
}
k
}
pub fn process_group_from_env() -> Result<Arc<crate::transport::ProcessGroup>, CollectiveError> {
let var = |k: &str| std::env::var(k).ok();
let rank: u32 = var("RANK")
.ok_or_else(|| terr("RANK not set"))?
.trim()
.parse()
.map_err(|_| terr("RANK must be an integer"))?;
let world: u32 = var("WORLD")
.ok_or_else(|| terr("WORLD not set"))?
.trim()
.parse()
.map_err(|_| terr("WORLD must be an integer"))?;
if rank >= world {
return Err(terr(format!("RANK {rank} must be < WORLD {world}")));
}
let alpn = var("RLX_IROH_ALPN")
.map(String::into_bytes)
.unwrap_or_else(|| RLX_PIPELINE_ALPN.to_vec());
let (secret, peers): (SecretKey, Vec<IrohPeer>) = if let Some(csv) = var("RLX_IROH_PEERS") {
let peers: Vec<IrohPeer> = csv
.split(',')
.map(str::trim)
.filter(|s| !s.is_empty())
.map(IrohPeer::from_id_str)
.collect::<Result<_, _>>()?;
if peers.len() != world as usize {
return Err(terr(format!(
"RLX_IROH_PEERS has {} ids but WORLD={world}",
peers.len()
)));
}
let sk_hex = var("RLX_IROH_SECRET").ok_or_else(|| {
terr("RLX_IROH_PEERS set: also set RLX_IROH_SECRET (this rank's 64-hex key)")
})?;
let arr: [u8; 32] = hex_to_bytes(&sk_hex)?
.as_slice()
.try_into()
.map_err(|_| terr("RLX_IROH_SECRET must be 32 bytes (64 hex chars)"))?;
(SecretKey::from_bytes(&arr), peers)
} else if let Some(seed_hex) = var("RLX_IROH_SEED") {
let seed = hex_to_bytes(&seed_hex)?;
if seed.is_empty() {
return Err(terr("RLX_IROH_SEED is empty"));
}
let peers: Vec<IrohPeer> = (0..world)
.map(|r| IrohPeer::new(SecretKey::from_bytes(&derive_key(&seed, r)).public()))
.collect();
(SecretKey::from_bytes(&derive_key(&seed, rank)), peers)
} else {
return Err(terr(
"set RLX_IROH_SEED (shared) or RLX_IROH_PEERS+RLX_IROH_SECRET to configure the iroh group",
));
};
let t = IrohTransport::connect_discovered(rank, world, secret, peers, &alpn)?;
Ok(Arc::new(crate::transport::ProcessGroup::new(Arc::new(t))))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::transport::ProcessGroup;
#[test]
fn iroh_two_rank_roundtrip() {
let t0 = IrohTransport::connect_ephemeral(0, 2, placeholder_peers(2), RLX_PIPELINE_ALPN)
.expect("rank0 up");
let t1 = IrohTransport::connect_ephemeral(1, 2, placeholder_peers(2), RLX_PIPELINE_ALPN)
.expect("rank1 up");
let id0 = t0.endpoint_id_string();
let id1 = t1.endpoint_id_string();
assert_ne!(id0, id1);
assert_eq!(t0.rank(), 0);
assert_eq!(t1.world_size(), 2);
let g0 = ProcessGroup::new(Arc::new(t0));
g0.transport().send_bytes(0, 7, b"ping").unwrap();
assert_eq!(g0.transport().recv_bytes(0, 7).unwrap(), b"ping");
}
fn placeholder_peers(n: u32) -> Vec<IrohPeer> {
(0..n)
.map(|_| IrohPeer::new(SecretKey::generate().public()))
.collect()
}
}