use std::{collections::BTreeSet, net::SocketAddr, sync::Arc};
use tokio::sync::Mutex;
use async_trait::async_trait;
use iroh::{
Endpoint, NodeAddr, RelayMode, Watcher,
endpoint::{Connection, RecvStream, SendStream},
};
use serde::{Deserialize, Serialize};
#[allow(unused_imports)] use tokio::io::AsyncWriteExt;
use tokio::sync::oneshot;
use super::{SyncTransport, shared::*};
use crate::{
Result,
sync::{
error::SyncError,
handler::SyncHandler,
peer_types::Address,
protocol::{RequestContext, SyncRequest, SyncResponse},
},
};
const SYNC_ALPN: &[u8] = b"eidetica/v0";
#[derive(Debug, Clone, Serialize, Deserialize)]
struct NodeAddrInfo {
node_id: String,
direct_addresses: BTreeSet<SocketAddr>,
}
impl From<&NodeAddr> for NodeAddrInfo {
fn from(node_addr: &NodeAddr) -> Self {
Self {
node_id: node_addr.node_id.to_string(),
direct_addresses: node_addr.direct_addresses.iter().cloned().collect(),
}
}
}
impl TryFrom<NodeAddrInfo> for NodeAddr {
type Error = crate::Error;
fn try_from(info: NodeAddrInfo) -> Result<Self> {
let node_id = info.node_id.parse().map_err(|e| {
SyncError::SerializationError(format!("Invalid NodeId '{}': {}", info.node_id, e))
})?;
Ok(NodeAddr::from_parts(
node_id,
None, info.direct_addresses,
))
}
}
#[derive(Debug, Clone)]
pub struct IrohTransportBuilder {
relay_mode: RelayMode,
}
impl IrohTransportBuilder {
pub fn new() -> Self {
Self {
relay_mode: RelayMode::Default, }
}
pub fn relay_mode(mut self, mode: RelayMode) -> Self {
self.relay_mode = mode;
self
}
pub fn build(self) -> Result<IrohTransport> {
Ok(IrohTransport {
endpoint: Arc::new(Mutex::new(None)),
server_state: ServerState::new(),
handler: None,
config: IrohTransportConfig {
relay_mode: self.relay_mode,
},
})
}
}
impl Default for IrohTransportBuilder {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone)]
struct IrohTransportConfig {
relay_mode: RelayMode,
}
pub struct IrohTransport {
endpoint: Arc<Mutex<Option<Endpoint>>>,
server_state: ServerState,
handler: Option<Arc<dyn SyncHandler>>,
config: IrohTransportConfig,
}
impl IrohTransport {
pub const TRANSPORT_TYPE: &'static str = "iroh";
pub fn new() -> Result<Self> {
IrohTransportBuilder::new().build()
}
pub fn builder() -> IrohTransportBuilder {
IrohTransportBuilder::new()
}
async fn ensure_endpoint(&self) -> Result<Endpoint> {
let mut endpoint_lock = self.endpoint.lock().await;
if endpoint_lock.is_none() {
let builder = Endpoint::builder()
.alpns(vec![SYNC_ALPN.to_vec()])
.relay_mode(self.config.relay_mode.clone());
let endpoint = builder.bind().await.map_err(|e| {
SyncError::TransportInit(format!("Failed to create Iroh endpoint: {e}"))
})?;
*endpoint_lock = Some(endpoint);
}
Ok(endpoint_lock.as_ref().unwrap().clone())
}
async fn start_server_loop(
&self,
endpoint: Endpoint,
ready_tx: oneshot::Sender<()>,
shutdown_rx: oneshot::Receiver<()>,
handler: Arc<dyn SyncHandler>,
) -> Result<()> {
let mut shutdown_rx = shutdown_rx;
let _ = ready_tx.send(());
tokio::spawn(async move {
loop {
tokio::select! {
_ = &mut shutdown_rx => {
break;
}
connection_result = endpoint.accept() => {
match connection_result {
Some(connecting) => {
let handler_clone = handler.clone();
tokio::spawn(async move {
if let Ok(conn) = connecting.await {
Self::handle_connection(conn, handler_clone).await;
}
});
}
None => break, }
}
}
}
});
Ok(())
}
async fn handle_connection(conn: Connection, handler: Arc<dyn SyncHandler>) {
let remote_node_id = match conn.remote_node_id() {
Ok(node_id) => node_id,
Err(e) => {
tracing::error!("Failed to get remote node ID: {e}");
return;
}
};
let remote_address = Address {
transport_type: Self::TRANSPORT_TYPE.to_string(),
address: remote_node_id.to_string(),
};
while let Ok((send_stream, recv_stream)) = conn.accept_bi().await {
let handler_clone = handler.clone();
let remote_addr_clone = remote_address.clone();
tokio::spawn(Self::handle_stream(
send_stream,
recv_stream,
handler_clone,
remote_addr_clone,
));
}
}
async fn handle_stream(
mut send_stream: SendStream,
mut recv_stream: RecvStream,
handler: Arc<dyn SyncHandler>,
remote_address: Address,
) {
let buffer: Vec<u8> = match recv_stream.read_to_end(1024 * 1024).await {
Ok(buffer) => buffer,
Err(e) => {
tracing::error!("Failed to read stream: {e}");
return;
}
};
let request: SyncRequest = match JsonHandler::deserialize_request(&buffer) {
Ok(req) => req,
Err(e) => {
tracing::error!("Failed to deserialize request: {e}");
return;
}
};
let peer_pubkey = match &request {
SyncRequest::SyncTree(sync_tree_request) => sync_tree_request.peer_pubkey.clone(),
_ => None,
};
let context = RequestContext {
remote_address: Some(remote_address),
peer_pubkey,
};
let response = handler.handle_request(&request, &context).await;
match JsonHandler::serialize_response(&response) {
Ok(response_bytes) => {
if let Err(e) = send_stream.write_all(&response_bytes).await {
tracing::error!("Failed to write response: {e}");
return;
}
if let Err(e) = send_stream.finish() {
tracing::error!("Failed to finish stream: {e}");
}
}
Err(e) => {
tracing::error!("Failed to serialize response: {e}");
}
}
}
}
#[async_trait]
impl SyncTransport for IrohTransport {
fn can_handle_address(&self, address: &Address) -> bool {
address.transport_type == Self::TRANSPORT_TYPE
}
async fn start_server(&mut self, _addr: &str, handler: Arc<dyn SyncHandler>) -> Result<()> {
if self.server_state.is_running() {
return Err(SyncError::ServerAlreadyRunning {
address: "iroh-endpoint".to_string(),
}
.into());
}
self.handler = Some(handler);
let endpoint = self.ensure_endpoint().await?;
let endpoint_clone = endpoint.clone();
let node_addr = endpoint.node_addr().initialized().await;
let node_addr_info = NodeAddrInfo::from(&node_addr);
let node_addr_str = serde_json::to_string(&node_addr_info)
.map_err(|e| SyncError::TransportInit(format!("Failed to serialize NodeAddr: {e}")))?;
let (ready_tx, ready_rx) = oneshot::channel();
let (shutdown_tx, shutdown_rx) = oneshot::channel();
self.start_server_loop(
endpoint_clone,
ready_tx,
shutdown_rx,
self.handler.clone().unwrap(),
)
.await?;
wait_for_ready(ready_rx, "iroh-endpoint").await?;
self.server_state.server_started(node_addr_str, shutdown_tx);
Ok(())
}
async fn stop_server(&mut self) -> Result<()> {
if !self.server_state.is_running() {
return Err(SyncError::ServerNotRunning.into());
}
self.server_state.stop_server();
Ok(())
}
async fn send_request(&self, address: &Address, request: &SyncRequest) -> Result<SyncResponse> {
if !self.can_handle_address(address) {
return Err(SyncError::UnsupportedTransport {
transport_type: address.transport_type.clone(),
}
.into());
}
let endpoint = self.ensure_endpoint().await?;
let node_addr_info: NodeAddrInfo = serde_json::from_str(&address.address).map_err(|e| {
SyncError::SerializationError(format!(
"Failed to parse NodeAddrInfo from '{}': {}",
address.address, e
))
})?;
let node_addr = NodeAddr::try_from(node_addr_info)?;
let conn = endpoint.connect(node_addr, SYNC_ALPN).await.map_err(|e| {
SyncError::ConnectionFailed {
address: address.address.clone(),
reason: e.to_string(),
}
})?;
let (mut send_stream, mut recv_stream) = conn
.open_bi()
.await
.map_err(|e| SyncError::Network(format!("Failed to open stream: {e}")))?;
let request_bytes = JsonHandler::serialize_request(request)?;
send_stream
.write_all(&request_bytes)
.await
.map_err(|e| SyncError::Network(format!("Failed to write request: {e}")))?;
send_stream
.finish()
.map_err(|e| SyncError::Network(format!("Failed to finish send stream: {e}")))?;
let response_bytes: Vec<u8> = recv_stream
.read_to_end(1024 * 1024)
.await
.map_err(|e| SyncError::Network(format!("Failed to read response: {e}")))?;
let response: SyncResponse = JsonHandler::deserialize_response(&response_bytes)?;
Ok(response)
}
fn is_server_running(&self) -> bool {
self.server_state.is_running()
}
fn get_server_address(&self) -> Result<String> {
self.server_state.get_address().map_err(|e| e.into())
}
}