use std::net::SocketAddr;
use std::path::Path;
use std::time::Duration;
use crate::wire::DEFAULT_MAX_DATAGRAM;
use bytes::Bytes;
use iroh::endpoint::{
presets, Connection, ConnectionError, IdleTimeout, PathId, QuicTransportConfig, VarInt,
};
use iroh::{Endpoint, EndpointAddr, EndpointId, RelayMode, RelayUrl, SecretKey};
pub mod auth;
pub mod ratelimit;
#[expect(
clippy::expect_used,
reason = "300s is far below IdleTimeout's varint ceiling; the conversion is statically infallible"
)]
#[allow(
clippy::duration_suboptimal_units,
reason = "`from_secs(300)` is the intended, readable idle timeout"
)]
fn koh_transport_config() -> QuicTransportConfig {
QuicTransportConfig::builder()
.keep_alive_interval(Duration::from_secs(5))
.max_idle_timeout(Some(
IdleTimeout::try_from(Duration::from_secs(300)).expect("300s fits in IdleTimeout"),
))
.build()
}
pub const ALPN: &[u8] = b"koh/iroh/1";
#[derive(Debug, thiserror::Error)]
pub enum SetupError {
#[error("io error: {0}")]
Io(#[from] std::io::Error),
#[error("secret key file is not 32 bytes of hex")]
BadKeyFile,
#[error("could not parse endpoint id: {0}")]
BadEndpointId(String),
#[error(transparent)]
Other(#[from] anyhow::Error),
}
pub fn load_or_create_secret_key(path: &Path) -> Result<SecretKey, SetupError> {
if path.exists() {
let text = std::fs::read_to_string(path)?;
let bytes = data_encoding::HEXLOWER_PERMISSIVE
.decode(text.trim().as_bytes())
.map_err(|_| SetupError::BadKeyFile)?;
let arr: [u8; 32] = bytes.try_into().map_err(|_| SetupError::BadKeyFile)?;
Ok(SecretKey::from_bytes(&arr))
} else {
let sk = generate_secret_key();
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
std::fs::write(path, data_encoding::HEXLOWER.encode(&sk.to_bytes()))?;
Ok(sk)
}
}
pub fn generate_secret_key() -> SecretKey {
use rand::RngCore;
let mut bytes = [0u8; 32];
rand::rngs::OsRng.fill_bytes(&mut bytes);
SecretKey::from_bytes(&bytes)
}
pub fn parse_endpoint_id(s: &str) -> Result<EndpointId, SetupError> {
s.trim()
.parse::<EndpointId>()
.map_err(|e| SetupError::BadEndpointId(e.to_string()))
}
pub fn format_endpoint_id(id: &EndpointId) -> String {
id.to_string()
}
fn parse_dns_spec(spec: &str) -> Option<SocketAddr> {
let spec = spec.trim();
spec.parse::<SocketAddr>().ok().or_else(|| {
spec.parse::<std::net::IpAddr>()
.ok()
.map(|ip| SocketAddr::new(ip, 53))
})
}
#[cfg_attr(
target_os = "android",
expect(
clippy::unnecessary_wraps,
reason = "Android always pins a nameserver (Some); the None arm is desktop-only"
)
)]
fn discovery_dns_resolver() -> Option<iroh::dns::DnsResolver> {
use iroh::dns::DnsResolver;
if let Some(addr) = std::env::var("KOH_DNS")
.ok()
.as_deref()
.and_then(parse_dns_spec)
{
return Some(DnsResolver::with_nameserver(addr));
}
#[cfg(target_os = "android")]
{
Some(DnsResolver::with_nameserver(SocketAddr::from((
[8, 8, 8, 8],
53,
))))
}
#[cfg(not(target_os = "android"))]
{
None
}
}
pub async fn bind_endpoint(secret: SecretKey, accept: bool) -> Result<Endpoint, SetupError> {
let mut builder = Endpoint::builder(presets::N0)
.secret_key(secret)
.transport_config(koh_transport_config());
if let Some(resolver) = discovery_dns_resolver() {
builder = builder.dns_resolver(resolver);
}
if accept {
builder = builder.alpns(vec![ALPN.to_vec()]);
}
let ep = builder
.bind()
.await
.map_err(|e| SetupError::Other(e.into()))?;
Ok(ep)
}
pub async fn bind_endpoint_local(secret: SecretKey, accept: bool) -> Result<Endpoint, SetupError> {
let mut builder = Endpoint::builder(presets::Minimal)
.secret_key(secret)
.transport_config(koh_transport_config());
if let Some(resolver) = discovery_dns_resolver() {
builder = builder.dns_resolver(resolver);
}
if accept {
builder = builder.alpns(vec![ALPN.to_vec()]);
}
let ep = builder
.bind()
.await
.map_err(|e| SetupError::Other(e.into()))?;
Ok(ep)
}
pub fn loopback_addr(ep: &Endpoint) -> EndpointAddr {
let mut addr = EndpointAddr::new(ep.id());
if let Some(port) = ep
.bound_sockets()
.iter()
.find(|s| s.is_ipv4())
.map(std::net::SocketAddr::port)
{
addr = addr.with_ip_addr(SocketAddr::from(([127, 0, 0, 1], port)));
}
addr
}
pub fn direct_addr(id: EndpointId, addr: SocketAddr) -> EndpointAddr {
EndpointAddr::new(id).with_ip_addr(addr)
}
pub fn relay_addr(id: EndpointId, relay: RelayUrl) -> EndpointAddr {
EndpointAddr::new(id).with_relay_url(relay)
}
pub async fn bind_endpoint_with_relay(
secret: SecretKey,
accept: bool,
relay: RelayUrl,
) -> Result<Endpoint, SetupError> {
let mut builder = Endpoint::builder(presets::Minimal)
.secret_key(secret)
.relay_mode(RelayMode::custom([relay]))
.transport_config(koh_transport_config());
if let Some(resolver) = discovery_dns_resolver() {
builder = builder.dns_resolver(resolver);
}
if accept {
builder = builder.alpns(vec![ALPN.to_vec()]);
}
let ep = builder
.bind()
.await
.map_err(|e| SetupError::Other(e.into()))?;
Ok(ep)
}
pub fn parse_relay_url(s: &str) -> Result<RelayUrl, SetupError> {
s.trim()
.parse::<RelayUrl>()
.map_err(|e| SetupError::Other(anyhow::anyhow!("bad relay url: {e}")))
}
#[derive(Clone)]
pub struct IrohChannel {
conn: Connection,
}
impl IrohChannel {
pub fn new(conn: Connection) -> Self {
Self { conn }
}
pub fn remote_id(&self) -> EndpointId {
self.conn.remote_id()
}
pub fn connection(&self) -> &Connection {
&self.conn
}
pub fn send(&self, datagram: &[u8]) -> bool {
match self.conn.send_datagram(Bytes::copy_from_slice(datagram)) {
Ok(()) => true,
Err(e) => {
tracing::trace!(error = %e, len = datagram.len(), "datagram send dropped");
false
}
}
}
pub async fn recv(&self) -> Result<Bytes, ConnectionError> {
self.conn.read_datagram().await
}
pub fn max_datagram_size(&self) -> usize {
self.conn
.max_datagram_size()
.unwrap_or(DEFAULT_MAX_DATAGRAM)
.max(64)
}
pub fn rtt_ms(&self) -> Option<f64> {
let to_ms = |d: Duration| d.as_secs_f64() * 1000.0;
if let Some(p) = self
.conn
.paths()
.iter()
.find(iroh::endpoint::Path::is_selected)
{
return Some(to_ms(p.rtt()));
}
if let Some(p) = self.conn.paths().iter().next() {
return Some(to_ms(p.rtt()));
}
self.conn.rtt(PathId::ZERO).map(to_ms)
}
pub fn close(&self, code: u32, reason: &[u8]) {
self.conn.close(VarInt::from_u32(code), reason);
}
pub async fn closed(&self) -> ConnectionError {
self.conn.closed().await
}
}
#[derive(Debug, Clone, Copy)]
pub struct MonoClock {
base: tokio::time::Instant,
}
impl Default for MonoClock {
fn default() -> Self {
Self::new()
}
}
impl MonoClock {
pub fn new() -> Self {
Self {
base: tokio::time::Instant::now(),
}
}
pub fn now_ms(&self) -> u64 {
self.base.elapsed().as_millis() as u64
}
pub fn base(&self) -> tokio::time::Instant {
self.base
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn secret_key_roundtrips_through_disk() {
let dir = std::env::temp_dir().join(format!("koh-key-test-{}", std::process::id()));
let path = dir.join("id.key");
let _ = std::fs::remove_dir_all(&dir);
let sk1 = load_or_create_secret_key(&path).unwrap();
let sk2 = load_or_create_secret_key(&path).unwrap();
assert_eq!(
sk1.to_bytes(),
sk2.to_bytes(),
"second load must reuse the key"
);
let id = sk1.public();
let s = format_endpoint_id(&id);
assert_eq!(parse_endpoint_id(&s).unwrap(), id);
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn parse_rejects_garbage() {
assert!(parse_endpoint_id("not-a-real-endpoint-id").is_err());
}
#[test]
fn dns_spec_accepts_ip_and_ip_port_rejects_junk() {
assert_eq!(
parse_dns_spec("1.1.1.1"),
Some(SocketAddr::from(([1, 1, 1, 1], 53)))
);
assert_eq!(
parse_dns_spec("8.8.8.8:5353"),
Some(SocketAddr::from(([8, 8, 8, 8], 5353)))
);
assert_eq!(
parse_dns_spec("2001:4860:4860::8888").map(|a| a.port()),
Some(53)
);
assert_eq!(
parse_dns_spec("[2001:4860:4860::8888]:53").map(|a| a.port()),
Some(53)
);
assert_eq!(
parse_dns_spec(" 9.9.9.9 "),
Some(SocketAddr::from(([9, 9, 9, 9], 53)))
);
assert_eq!(parse_dns_spec(""), None);
assert_eq!(parse_dns_spec("not-an-ip"), None);
assert_eq!(parse_dns_spec("8.8.8.8:"), None);
assert_eq!(parse_dns_spec("8.8.8.8:99999"), None);
}
#[test]
fn explicit_nameserver_resolver_builds() {
let _resolver =
iroh::dns::DnsResolver::with_nameserver(SocketAddr::from(([8, 8, 8, 8], 53)));
}
#[tokio::test]
async fn two_endpoints_exchange_datagram_over_loopback() {
let server = bind_endpoint_local(generate_secret_key(), true)
.await
.expect("bind server");
let client = bind_endpoint_local(generate_secret_key(), false)
.await
.expect("bind client");
let server_addr = loopback_addr(&server);
let srv = tokio::spawn(async move {
let incoming = server.accept().await.expect("accept");
let conn = incoming.await.expect("handshake");
let dg = conn.read_datagram().await.expect("read datagram");
conn.send_datagram(dg).expect("echo datagram"); conn.closed().await;
});
let conn = client
.connect(server_addr, ALPN)
.await
.expect("connect over loopback");
let chan = IrohChannel::new(conn);
assert!(
chan.send(b"ping-over-real-iroh"),
"datagram send should succeed"
);
let echoed = chan.recv().await.expect("recv echo");
assert_eq!(&echoed[..], b"ping-over-real-iroh");
chan.close(0, b"done");
let _ = srv.await;
}
}