use std::future::Future;
use std::time::Duration;
use crate::cbor::Value;
use crate::cert_chain::{self, CertChainError};
use crate::connection::{self, Session};
use crate::content;
use crate::dht::{self, DhtError, Record};
use crate::frame::{CallResponse, StreamMode};
use crate::identity::KeyPair;
use crate::manifest::Mcid;
use crate::stream::{self, StreamHandle};
use crate::transport::Trust;
fn now_ms() -> i128 {
use std::time::{SystemTime, UNIX_EPOCH};
SystemTime::now()
.duration_since(UNIX_EPOCH)
.expect("system clock before 1970")
.as_millis() as i128
}
const RESOLVE_RETRIES: u32 = 50;
const RESOLVE_RETRY_DELAY: Duration = Duration::from_millis(100);
#[derive(Debug)]
pub enum ResolveError {
ProcedureNotAdvertised,
NoTrustedAdvertisement,
StationEndpointNotFound,
StationEndpointSignerMismatch,
Dht(DhtError),
NoAuthorizedAdvertisement(CertChainError),
}
impl std::fmt::Display for ResolveError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ResolveError::ProcedureNotAdvertised => {
write!(
f,
"direct_dial: procedure has no direct-dial advertisement in the DHT"
)
}
ResolveError::NoTrustedAdvertisement => write!(
f,
"direct_dial: every candidate advertisement failed signature verification"
),
ResolveError::StationEndpointNotFound => write!(
f,
"direct_dial: resolved station published no reachable station_endpoint"
),
ResolveError::StationEndpointSignerMismatch => {
write!(f, "direct_dial: station_endpoint signer mismatch")
}
ResolveError::Dht(e) => write!(f, "direct_dial: {e}"),
ResolveError::NoAuthorizedAdvertisement(e) => write!(
f,
"direct_dial: no candidate advertisement is cert-chain-authorized for the expected org: {e}"
),
}
}
}
impl std::error::Error for ResolveError {}
#[derive(Debug, Clone)]
pub struct Resolved {
pub station: [u8; 32],
pub host: String,
pub port: u16,
}
pub async fn resolve(
session: &mut Session,
id: &KeyPair,
realm: [u8; 32],
procedure: &str,
) -> Result<Resolved, ResolveError> {
let uri = dht::discovery_uri(realm, procedure);
let key = dht::procedure_key(&uri);
let mut recs: Vec<Record> = Vec::new();
for _ in 0..RESOLVE_RETRIES {
match dht::find_records(session, id, key).await {
Ok(found) if !found.is_empty() => {
recs = found;
break;
}
Ok(_) => {}
Err(e) => return Err(ResolveError::Dht(e)),
}
tokio::time::sleep(RESOLVE_RETRY_DELAY).await;
}
if recs.is_empty() {
return Err(ResolveError::ProcedureNotAdvertised);
}
let adv = first_trusted_advertisement(&recs).ok_or(ResolveError::NoTrustedAdvertisement)?;
resolve_station_endpoint(session, id, adv.serving_station).await
}
fn first_trusted_advertisement(recs: &[Record]) -> Option<dht::ProcedureAdvertisement> {
recs.iter().find_map(|rec| {
dht::verify(rec).ok()?;
dht::read_procedure_advertisement(rec).ok()
})
}
pub async fn resolve_with_cert_chain(
session: &mut Session,
id: &KeyPair,
realm: [u8; 32],
procedure: &str,
realm_ca_pem: &[u8],
expected_org: &str,
) -> Result<Resolved, ResolveError> {
let uri = dht::discovery_uri(realm, procedure);
let key = dht::procedure_key(&uri);
let mut recs: Vec<Record> = Vec::new();
for _ in 0..RESOLVE_RETRIES {
match dht::find_records(session, id, key).await {
Ok(found) if !found.is_empty() => {
recs = found;
break;
}
Ok(_) => {}
Err(e) => return Err(ResolveError::Dht(e)),
}
tokio::time::sleep(RESOLVE_RETRY_DELAY).await;
}
if recs.is_empty() {
return Err(ResolveError::ProcedureNotAdvertised);
}
let adv = first_authorized_advertisement(&recs, realm_ca_pem, expected_org)?;
resolve_station_endpoint(session, id, adv.serving_station).await
}
fn first_authorized_advertisement(
recs: &[Record],
realm_ca_pem: &[u8],
expected_org: &str,
) -> Result<dht::ProcedureAdvertisement, ResolveError> {
let mut last_cert_err: Option<CertChainError> = None;
for rec in recs {
if dht::verify(rec).is_err() {
continue;
}
match cert_chain::verify_advertisement_cert_chain(realm_ca_pem, rec, expected_org) {
Ok(()) => {
if let Ok(adv) = dht::read_procedure_advertisement(rec) {
return Ok(adv);
}
}
Err(e) => last_cert_err = Some(e),
}
}
match last_cert_err {
Some(e) => Err(ResolveError::NoAuthorizedAdvertisement(e)),
None => Err(ResolveError::NoTrustedAdvertisement),
}
}
async fn resolve_station_endpoint(
session: &mut Session,
id: &KeyPair,
station: [u8; 32],
) -> Result<Resolved, ResolveError> {
let key = dht::station_endpoint_key(station);
for _ in 0..RESOLVE_RETRIES {
let rec = match dht::find_record(session, id, key).await {
Ok(rec) => rec,
Err(DhtError::NotFound) => {
tokio::time::sleep(RESOLVE_RETRY_DELAY).await;
continue;
}
Err(e) => return Err(ResolveError::Dht(e)),
};
if rec.key != station {
return Err(ResolveError::StationEndpointSignerMismatch);
}
match dht::verify(&rec) {
Ok(()) => {}
Err(dht::VerifyError::Expired) => {
tokio::time::sleep(RESOLVE_RETRY_DELAY).await;
continue;
}
Err(_) => return Err(ResolveError::NoTrustedAdvertisement),
}
let ep =
dht::read_station_endpoint(&rec).map_err(|_| ResolveError::StationEndpointNotFound)?;
let Some(host) = ep.host_advertised.into_iter().next() else {
return Err(ResolveError::StationEndpointNotFound);
};
return Ok(Resolved {
station,
host,
port: ep.quic_port,
});
}
Err(ResolveError::StationEndpointNotFound)
}
#[derive(Debug)]
pub enum CallError {
Resolve(ResolveError),
Dial(connection::HandshakeError),
TrustViolation {
resolved: [u8; 32],
dialed: [u8; 32],
},
Call(connection::CallError),
}
impl std::fmt::Display for CallError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
CallError::Resolve(e) => write!(f, "{e}"),
CallError::Dial(e) => write!(f, "direct_dial: dialing resolved station: {e}"),
CallError::TrustViolation { resolved, dialed } => write!(
f,
"direct_dial: trust violation -- resolved station {} but the dialed peer proved identity {}",
hex_of(resolved),
hex_of(dialed)
),
CallError::Call(e) => write!(f, "direct_dial: {e}"),
}
}
}
impl std::error::Error for CallError {}
fn hex_of(b: &[u8; 32]) -> String {
b.iter().map(|byte| format!("{byte:02x}")).collect()
}
pub async fn call(
resolve_via: &mut Session,
id: &KeyPair,
realm: [u8; 32],
procedure: &str,
payload: Value,
timeout: Duration,
) -> Result<CallResponse, CallError> {
let resolved = resolve(resolve_via, id, realm, procedure)
.await
.map_err(CallError::Resolve)?;
let mut target = tokio::time::timeout(
timeout,
connection::connect(&resolved.host, resolved.port, Trust::Insecure, id),
)
.await
.unwrap_or(Err(connection::HandshakeError::Timeout))
.map_err(CallError::Dial)?;
if target.station.node_id != resolved.station {
let dialed = target.station.node_id;
target.close("trust_violation", None, id).await;
return Err(CallError::TrustViolation {
resolved: resolved.station,
dialed,
});
}
let deadline_ms = now_ms() + timeout.as_millis() as i128;
let result = target
.call(procedure, realm, payload, deadline_ms, id, timeout)
.await
.map_err(CallError::Call);
target.close("normal", None, id).await;
result
}
pub async fn call_with_ucan(
resolve_via: &mut Session,
id: &KeyPair,
realm: [u8; 32],
procedure: &str,
payload: Value,
timeout: Duration,
ucan_token: Vec<u8>,
) -> Result<CallResponse, CallError> {
let resolved = resolve(resolve_via, id, realm, procedure)
.await
.map_err(CallError::Resolve)?;
let mut target = tokio::time::timeout(
timeout,
connection::connect(&resolved.host, resolved.port, Trust::Insecure, id),
)
.await
.unwrap_or(Err(connection::HandshakeError::Timeout))
.map_err(CallError::Dial)?;
if target.station.node_id != resolved.station {
let dialed = target.station.node_id;
target.close("trust_violation", None, id).await;
return Err(CallError::TrustViolation {
resolved: resolved.station,
dialed,
});
}
let deadline_ms = now_ms() + timeout.as_millis() as i128;
let result = target
.call_with_ucan(
procedure,
realm,
payload,
deadline_ms,
id,
timeout,
ucan_token,
)
.await
.map_err(CallError::Call);
target.close("normal", None, id).await;
result
}
#[allow(clippy::too_many_arguments)]
pub async fn call_with_cert_chain(
resolve_via: &mut Session,
id: &KeyPair,
realm: [u8; 32],
procedure: &str,
realm_ca_pem: &[u8],
expected_org: &str,
payload: Value,
timeout: Duration,
) -> Result<CallResponse, CallError> {
let resolved = resolve_with_cert_chain(
resolve_via,
id,
realm,
procedure,
realm_ca_pem,
expected_org,
)
.await
.map_err(CallError::Resolve)?;
let mut target = tokio::time::timeout(
timeout,
connection::connect(&resolved.host, resolved.port, Trust::Insecure, id),
)
.await
.unwrap_or(Err(connection::HandshakeError::Timeout))
.map_err(CallError::Dial)?;
if target.station.node_id != resolved.station {
let dialed = target.station.node_id;
target.close("trust_violation", None, id).await;
return Err(CallError::TrustViolation {
resolved: resolved.station,
dialed,
});
}
let deadline_ms = now_ms() + timeout.as_millis() as i128;
let result = target
.call(procedure, realm, payload, deadline_ms, id, timeout)
.await
.map_err(CallError::Call);
target.close("normal", None, id).await;
result
}
pub async fn advertise_direct(
session: &mut Session,
id: &KeyPair,
realm: [u8; 32],
procedure: &str,
ttl: Duration,
) -> Result<(), AdvertiseDirectError> {
let advertise_spec = crate::frame::AdvertiseSpec::new(realm, procedure, id.node_id());
session
.advertise(&advertise_spec, id)
.await
.map_err(AdvertiseDirectError::Advertise)?;
let uri = dht::discovery_uri(realm, procedure);
let rec = dht::new_procedure_advertisement(id.node_id(), uri, session.station.node_id, ttl);
let rec = dht::sign(rec, id);
dht::put_record(session, id, &rec)
.await
.map_err(AdvertiseDirectError::Dht)
}
pub async fn advertise_direct_with_cert_chain(
session: &mut Session,
id: &KeyPair,
realm: [u8; 32],
procedure: &str,
ttl: Duration,
cert_chain_pem: Vec<u8>,
) -> Result<(), AdvertiseDirectError> {
let advertise_spec = crate::frame::AdvertiseSpec::new(realm, procedure, id.node_id());
session
.advertise(&advertise_spec, id)
.await
.map_err(AdvertiseDirectError::Advertise)?;
let uri = dht::discovery_uri(realm, procedure);
let rec = dht::new_procedure_advertisement_with_cert_chain(
id.node_id(),
uri,
session.station.node_id,
ttl,
cert_chain_pem,
);
let rec = dht::sign(rec, id);
dht::put_record(session, id, &rec)
.await
.map_err(AdvertiseDirectError::Dht)
}
#[derive(Debug)]
pub enum AdvertiseDirectError {
Advertise(connection::SendFrameError),
Dht(DhtError),
}
impl std::fmt::Display for AdvertiseDirectError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
AdvertiseDirectError::Advertise(e) => write!(f, "direct_dial: sending ADVERTISE: {e}"),
AdvertiseDirectError::Dht(e) => write!(f, "direct_dial: {e}"),
}
}
}
impl std::error::Error for AdvertiseDirectError {}
#[allow(clippy::too_many_arguments)]
pub async fn keep_advertised_direct<F>(
session: &mut Session,
id: &KeyPair,
realm: [u8; 32],
procedure: &str,
ttl: Duration,
interval: Duration,
stop: F,
on_error: impl Fn(AdvertiseDirectError),
) where
F: Future<Output = ()>,
{
tokio::pin!(stop);
let mut ticker = tokio::time::interval(interval);
loop {
tokio::select! {
_ = &mut stop => return,
_ = ticker.tick() => {
if let Err(e) = advertise_direct(session, id, realm, procedure, ttl).await {
on_error(e);
}
}
}
}
}
#[derive(Debug)]
pub enum DialAndVerifyError {
Dial(connection::HandshakeError),
TrustViolation {
resolved: [u8; 32],
dialed: [u8; 32],
},
}
impl std::fmt::Display for DialAndVerifyError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
DialAndVerifyError::Dial(e) => write!(f, "direct_dial: dialing resolved station: {e}"),
DialAndVerifyError::TrustViolation { resolved, dialed } => write!(
f,
"direct_dial: trust violation -- resolved station {} but the dialed peer proved identity {}",
hex_of(resolved),
hex_of(dialed)
),
}
}
}
impl std::error::Error for DialAndVerifyError {}
async fn dial_and_verify(
host: &str,
port: u16,
station: [u8; 32],
id: &KeyPair,
timeout: Duration,
) -> Result<Session, DialAndVerifyError> {
let target = tokio::time::timeout(
timeout,
connection::connect(host, port, Trust::Insecure, id),
)
.await
.unwrap_or(Err(connection::HandshakeError::Timeout))
.map_err(DialAndVerifyError::Dial)?;
if target.station.node_id != station {
let dialed = target.station.node_id;
target.close("trust_violation", None, id).await;
return Err(DialAndVerifyError::TrustViolation {
resolved: station,
dialed,
});
}
Ok(target)
}
#[derive(Debug)]
pub enum OpenStreamDirectError {
Resolve(ResolveError),
Dial(DialAndVerifyError),
Open(stream::OpenError),
}
impl std::fmt::Display for OpenStreamDirectError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
OpenStreamDirectError::Resolve(e) => write!(f, "{e}"),
OpenStreamDirectError::Dial(e) => write!(f, "{e}"),
OpenStreamDirectError::Open(e) => write!(f, "direct_dial: open stream: {e}"),
}
}
}
impl std::error::Error for OpenStreamDirectError {}
#[allow(clippy::too_many_arguments)]
pub async fn open_stream_direct(
resolve_via: &mut Session,
id: &KeyPair,
realm: [u8; 32],
procedure: &str,
mode: StreamMode,
args: Value,
deadline_ms: i128,
timeout: Duration,
) -> Result<(Session, StreamHandle), OpenStreamDirectError> {
let resolved = resolve(resolve_via, id, realm, procedure)
.await
.map_err(OpenStreamDirectError::Resolve)?;
let mut target = dial_and_verify(&resolved.host, resolved.port, resolved.station, id, timeout)
.await
.map_err(OpenStreamDirectError::Dial)?;
match StreamHandle::open(&mut target, procedure, realm, mode, args, deadline_ms, id).await {
Ok(handle) => Ok((target, handle)),
Err(e) => {
target.close("normal", None, id).await;
Err(OpenStreamDirectError::Open(e))
}
}
}
#[allow(clippy::too_many_arguments)]
pub async fn open_stream_direct_with_cert_chain(
resolve_via: &mut Session,
id: &KeyPair,
realm: [u8; 32],
procedure: &str,
realm_ca_pem: &[u8],
expected_org: &str,
mode: StreamMode,
args: Value,
deadline_ms: i128,
timeout: Duration,
) -> Result<(Session, StreamHandle), OpenStreamDirectError> {
let resolved = resolve_with_cert_chain(
resolve_via,
id,
realm,
procedure,
realm_ca_pem,
expected_org,
)
.await
.map_err(OpenStreamDirectError::Resolve)?;
let mut target = dial_and_verify(&resolved.host, resolved.port, resolved.station, id, timeout)
.await
.map_err(OpenStreamDirectError::Dial)?;
match StreamHandle::open(&mut target, procedure, realm, mode, args, deadline_ms, id).await {
Ok(handle) => Ok((target, handle)),
Err(e) => {
target.close("normal", None, id).await;
Err(OpenStreamDirectError::Open(e))
}
}
}
#[derive(Debug)]
pub enum PutDirectError {
Resolve(ResolveError),
Dial(DialAndVerifyError),
Put(content::PutError),
}
impl std::fmt::Display for PutDirectError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
PutDirectError::Resolve(e) => write!(f, "{e}"),
PutDirectError::Dial(e) => write!(f, "{e}"),
PutDirectError::Put(e) => write!(f, "direct_dial: {e}"),
}
}
}
impl std::error::Error for PutDirectError {}
pub async fn put_direct(
resolve_via: &mut Session,
id: &KeyPair,
station: [u8; 32],
data: &[u8],
name: impl Into<String>,
timeout: Duration,
) -> Result<Mcid, PutDirectError> {
let resolved = resolve_station_endpoint(resolve_via, id, station)
.await
.map_err(PutDirectError::Resolve)?;
let mut target = dial_and_verify(&resolved.host, resolved.port, resolved.station, id, timeout)
.await
.map_err(PutDirectError::Dial)?;
let result = content::put(&mut target, data, name, id)
.await
.map_err(PutDirectError::Put);
target.close("normal", None, id).await;
result
}
#[derive(Debug)]
pub struct ContentNotAnnounced;
impl std::fmt::Display for ContentNotAnnounced {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"direct_dial: content has no verifiable announcement in the DHT"
)
}
}
impl std::error::Error for ContentNotAnnounced {}
#[derive(Debug)]
pub enum GetDirectError {
Dht(DhtError),
NotAnnounced(ContentNotAnnounced),
EndpointParse(String),
Dial(DialAndVerifyError),
Get(content::GetError),
}
impl std::fmt::Display for GetDirectError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
GetDirectError::Dht(e) => write!(f, "direct_dial: find content providers: {e}"),
GetDirectError::NotAnnounced(e) => write!(f, "{e}"),
GetDirectError::EndpointParse(endpoint) => {
write!(
f,
"direct_dial: content provider endpoint {endpoint:?}: not a URL or host:port"
)
}
GetDirectError::Dial(e) => write!(f, "{e}"),
GetDirectError::Get(e) => write!(f, "direct_dial: {e}"),
}
}
}
impl std::error::Error for GetDirectError {}
pub async fn get_direct(
resolve_via: &mut Session,
id: &KeyPair,
mcid: Mcid,
timeout: Duration,
) -> Result<Vec<u8>, GetDirectError> {
let recs = dht::find_records(resolve_via, id, dht::content_key(mcid))
.await
.map_err(GetDirectError::Dht)?;
let adv = first_trusted_content_provider(&recs)
.ok_or(GetDirectError::NotAnnounced(ContentNotAnnounced))?;
let (host, port) = parse_seed_url(&adv.endpoint)
.ok_or_else(|| GetDirectError::EndpointParse(adv.endpoint.clone()))?;
let mut target = dial_and_verify(&host, port, adv.announcer_node, id, timeout)
.await
.map_err(GetDirectError::Dial)?;
let result = content::get(&mut target, mcid, id)
.await
.map_err(GetDirectError::Get);
target.close("normal", None, id).await;
result
}
fn first_trusted_content_provider(recs: &[Record]) -> Option<dht::ContentAnnouncement> {
recs.iter().find_map(|rec| {
dht::verify(rec).ok()?;
let adv = dht::read_content_announcement(rec).ok()?;
(adv.announcer_node == rec.key).then_some(adv)
})
}
fn parse_seed_url(seed: &str) -> Option<(String, u16)> {
if let Some(rest) = seed
.strip_prefix("https://")
.or_else(|| seed.strip_prefix("http://"))
{
let hostport = rest.split('/').next().unwrap_or(rest);
let (host, port_str) = hostport.rsplit_once(':')?;
return Some((host.to_string(), port_str.parse().ok()?));
}
let (host, port_str) = seed.rsplit_once(':')?;
Some((host.to_string(), port_str.parse().ok()?))
}