use std::sync::Arc;
use std::time::Duration;
use crate::bolt4;
use crate::cbor::Value;
use crate::frame::{self, Decoded, HelloInfo};
use crate::identity::KeyPair;
use crate::transport::{self, ConnectError, Trust};
pub type BoxFuture<'a, T> = std::pin::Pin<Box<dyn std::future::Future<Output = T> + Send + 'a>>;
pub type CallHandler =
Arc<dyn Fn(Value) -> BoxFuture<'static, Result<Value, String>> + Send + Sync>;
pub const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(30);
pub const DEFAULT_CALL_TIMEOUT: Duration = Duration::from_secs(30);
const READ_CHUNK: usize = 64 * 1024;
pub struct FrameStream {
send: quinn::SendStream,
recv: quinn::RecvStream,
buf: Vec<u8>,
}
#[derive(Debug)]
pub enum SendFrameError {
Encode(frame::EncodeFrameError),
Write(quinn::WriteError),
}
impl std::fmt::Display for SendFrameError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
SendFrameError::Encode(e) => write!(f, "encoding frame: {e}"),
SendFrameError::Write(e) => write!(f, "writing to stream: {e}"),
}
}
}
impl std::error::Error for SendFrameError {}
#[derive(Debug)]
pub enum RecvFrameError {
Read(quinn::ReadError),
StreamClosed,
Decode(frame::DecodeFrameError),
Timeout,
}
impl std::fmt::Display for RecvFrameError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
RecvFrameError::Read(e) => write!(f, "reading from stream: {e}"),
RecvFrameError::StreamClosed => write!(f, "peer closed the stream"),
RecvFrameError::Decode(e) => write!(f, "decoding a frame: {e}"),
RecvFrameError::Timeout => write!(f, "timed out waiting for a frame"),
}
}
}
impl std::error::Error for RecvFrameError {}
#[derive(Debug)]
pub enum CallError {
Send(SendFrameError),
Recv(RecvFrameError),
}
impl std::fmt::Display for CallError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
CallError::Send(e) => write!(f, "sending CALL: {e}"),
CallError::Recv(e) => write!(f, "awaiting RESULT/ERROR: {e}"),
}
}
}
impl std::error::Error for CallError {}
impl FrameStream {
fn new(send: quinn::SendStream, recv: quinn::RecvStream) -> Self {
Self::with_buf(send, recv, Vec::new())
}
fn with_buf(send: quinn::SendStream, recv: quinn::RecvStream, buf: Vec<u8>) -> Self {
Self { send, recv, buf }
}
pub fn leftover_bytes(&self) -> &[u8] {
&self.buf
}
pub async fn send_frame(&mut self, frame: Value) -> Result<(), SendFrameError> {
let encoded = frame::encode(&frame).map_err(SendFrameError::Encode)?;
self.send
.write_all(&encoded)
.await
.map_err(SendFrameError::Write)
}
pub async fn recv_frame(&mut self) -> Result<Value, RecvFrameError> {
let mut chunk = vec![0u8; READ_CHUNK];
loop {
match frame::decode(&self.buf) {
Ok(Decoded::Frame(value, consumed)) => {
self.buf.drain(..consumed);
return Ok(value);
}
Ok(Decoded::More(_)) => {}
Err(e) => return Err(RecvFrameError::Decode(e)),
}
let n = self
.recv
.read(&mut chunk)
.await
.map_err(RecvFrameError::Read)?
.ok_or(RecvFrameError::StreamClosed)?;
self.buf.extend_from_slice(&chunk[..n]);
}
}
pub async fn recv_frame_timeout(&mut self, timeout: Duration) -> Result<Value, RecvFrameError> {
tokio::time::timeout(timeout, self.recv_frame())
.await
.unwrap_or(Err(RecvFrameError::Timeout))
}
pub async fn call(
&mut self,
procedure: &str,
realm: [u8; 32],
payload: Value,
deadline_ms: i128,
identity: &KeyPair,
timeout: Duration,
) -> Result<frame::CallResponse, CallError> {
let call_id: [u8; 16] = rand::random();
let spec = frame::CallSpec::new(
call_id,
procedure,
realm,
payload,
deadline_ms,
identity.node_id(),
);
let signed = frame::sign(frame::call(&spec), identity);
self.send_frame(signed).await.map_err(CallError::Send)?;
tokio::time::timeout(timeout, self.await_call_response(call_id))
.await
.unwrap_or(Err(CallError::Recv(RecvFrameError::Timeout)))
}
#[allow(clippy::too_many_arguments)]
pub async fn call_with_ucan(
&mut self,
procedure: &str,
realm: [u8; 32],
payload: Value,
deadline_ms: i128,
identity: &KeyPair,
timeout: Duration,
ucan_token: Vec<u8>,
) -> Result<frame::CallResponse, CallError> {
let call_id: [u8; 16] = rand::random();
let mut spec = frame::CallSpec::new(
call_id,
procedure,
realm,
payload,
deadline_ms,
identity.node_id(),
);
spec.ucan_token = ucan_token;
let signed = frame::sign(frame::call(&spec), identity);
self.send_frame(signed).await.map_err(CallError::Send)?;
tokio::time::timeout(timeout, self.await_call_response(call_id))
.await
.unwrap_or(Err(CallError::Recv(RecvFrameError::Timeout)))
}
async fn await_call_response(
&mut self,
call_id: [u8; 16],
) -> Result<frame::CallResponse, CallError> {
loop {
let value = self.recv_frame().await.map_err(CallError::Recv)?;
if frame::frame_call_id(&value) != Some(call_id) {
continue; }
if let Ok(response) = frame::parse_call_response(&value) {
return Ok(response);
}
}
}
}
pub struct Session {
connection: quinn::Connection,
control: FrameStream,
pub station: HelloInfo,
}
#[derive(Debug)]
pub enum HandshakeError {
Transport(ConnectError),
OpenStream(quinn::ConnectionError),
Write(quinn::WriteError),
Read(quinn::ReadError),
StreamClosed,
Timeout,
Encode(frame::EncodeFrameError),
Decode(frame::DecodeFrameError),
UnexpectedFrameType(frame::ParseHelloError),
SignatureInvalid(frame::VerifyError),
Refused {
refusal_code: Option<i128>,
},
}
impl std::fmt::Display for HandshakeError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
HandshakeError::Transport(e) => write!(f, "transport: {e}"),
HandshakeError::OpenStream(e) => write!(f, "opening control stream: {e}"),
HandshakeError::Write(e) => write!(f, "sending CONNECT: {e}"),
HandshakeError::Read(e) => write!(f, "reading from control stream: {e}"),
HandshakeError::StreamClosed => {
write!(f, "station closed the stream before HELLO arrived")
}
HandshakeError::Timeout => write!(
f,
"no HELLO within {HANDSHAKE_TIMEOUT:?} (likely a protocol mismatch)"
),
HandshakeError::Encode(e) => write!(f, "encoding CONNECT: {e}"),
HandshakeError::Decode(e) => write!(f, "decoding the station's response: {e}"),
HandshakeError::UnexpectedFrameType(e) => write!(f, "expected a HELLO frame: {e}"),
HandshakeError::SignatureInvalid(e) => write!(f, "HELLO signature check failed: {e}"),
HandshakeError::Refused { refusal_code } => {
write!(
f,
"station refused the connection (refusal_code = {refusal_code:?})"
)
}
}
}
}
impl std::error::Error for HandshakeError {}
pub async fn connect(
host: &str,
port: u16,
trust: Trust,
identity: &KeyPair,
) -> Result<Session, HandshakeError> {
tokio::time::timeout(
HANDSHAKE_TIMEOUT,
connect_inner(host, port, trust, identity),
)
.await
.unwrap_or(Err(HandshakeError::Timeout))
}
async fn connect_inner(
host: &str,
port: u16,
trust: Trust,
identity: &KeyPair,
) -> Result<Session, HandshakeError> {
let connection = transport::connect(host, port, trust)
.await
.map_err(HandshakeError::Transport)?;
let (mut send, mut recv) = connection
.open_bi()
.await
.map_err(HandshakeError::OpenStream)?;
let connect_spec =
crate::frame::ConnectSpec::new(identity.node_id(), identity.puzzle_evidence());
let connect_frame = frame::sign(frame::connect(&connect_spec), identity);
let encoded = frame::encode(&connect_frame).map_err(HandshakeError::Encode)?;
send.write_all(&encoded)
.await
.map_err(HandshakeError::Write)?;
let (hello_value, buf) = read_one_frame(&mut recv).await?;
let station = frame::parse_hello(&hello_value).map_err(HandshakeError::UnexpectedFrameType)?;
frame::verify(&hello_value, &station.node_id).map_err(HandshakeError::SignatureInvalid)?;
if !station.accepted {
return Err(HandshakeError::Refused {
refusal_code: station.refusal_code,
});
}
Ok(Session {
connection,
control: FrameStream::with_buf(send, recv, buf),
station,
})
}
async fn read_one_frame(recv: &mut quinn::RecvStream) -> Result<(Value, Vec<u8>), HandshakeError> {
let mut buf = Vec::new();
let mut chunk = vec![0u8; READ_CHUNK];
loop {
match frame::decode(&buf) {
Ok(Decoded::Frame(value, consumed)) => {
buf.drain(..consumed);
return Ok((value, buf));
}
Ok(Decoded::More(_)) => {}
Err(e) => return Err(HandshakeError::Decode(e)),
}
let n = recv
.read(&mut chunk)
.await
.map_err(HandshakeError::Read)?
.ok_or(HandshakeError::StreamClosed)?;
buf.extend_from_slice(&chunk[..n]);
}
}
impl Session {
pub fn remote_address(&self) -> std::net::SocketAddr {
self.connection.remote_address()
}
pub async fn open_dedicated_stream(&mut self) -> Result<FrameStream, quinn::ConnectionError> {
let (send, recv) = self.connection.open_bi().await?;
Ok(FrameStream::new(send, recv))
}
pub async fn accept_dedicated_stream(&mut self) -> Result<FrameStream, quinn::ConnectionError> {
let (send, recv) = self.connection.accept_bi().await?;
Ok(FrameStream::new(send, recv))
}
pub fn leftover_bytes(&self) -> &[u8] {
self.control.leftover_bytes()
}
pub async fn recv_frame(&mut self) -> Result<Value, RecvFrameError> {
self.control.recv_frame().await
}
pub async fn recv_frame_timeout(&mut self, timeout: Duration) -> Result<Value, RecvFrameError> {
self.control.recv_frame_timeout(timeout).await
}
pub async fn call(
&mut self,
procedure: &str,
realm: [u8; 32],
payload: Value,
deadline_ms: i128,
identity: &KeyPair,
timeout: Duration,
) -> Result<frame::CallResponse, CallError> {
let request_id: [u8; 16] = rand::random();
announce_rpc_sent(&mut *self, realm, identity, request_id).await;
let result = self
.control
.call(procedure, realm, payload, deadline_ms, identity, timeout)
.await;
announce_rpc_completed(&mut *self, realm, identity, request_id, &result).await;
result
}
#[allow(clippy::too_many_arguments)]
pub async fn call_with_ucan(
&mut self,
procedure: &str,
realm: [u8; 32],
payload: Value,
deadline_ms: i128,
identity: &KeyPair,
timeout: Duration,
ucan_token: Vec<u8>,
) -> Result<frame::CallResponse, CallError> {
let request_id: [u8; 16] = rand::random();
announce_rpc_sent(&mut *self, realm, identity, request_id).await;
let result = self
.control
.call_with_ucan(
procedure,
realm,
payload,
deadline_ms,
identity,
timeout,
ucan_token,
)
.await;
announce_rpc_completed(&mut *self, realm, identity, request_id, &result).await;
result
}
pub async fn publish(
&mut self,
spec: &frame::PublishSpec,
identity: &KeyPair,
) -> Result<(), SendFrameError> {
let unsigned = frame::publish(spec);
let with_publisher_sig = frame::sign_publisher(unsigned, identity);
let signed = frame::sign(with_publisher_sig, identity);
self.control.send_frame(signed).await
}
pub async fn subscribe(
&mut self,
spec: &frame::SubscribeSpec,
identity: &KeyPair,
) -> Result<(), SendFrameError> {
let signed = frame::sign(frame::subscribe(spec), identity);
self.control.send_frame(signed).await
}
pub async fn unsubscribe(
&mut self,
spec: &frame::UnsubscribeSpec,
identity: &KeyPair,
) -> Result<(), SendFrameError> {
let signed = frame::sign(frame::unsubscribe(spec), identity);
self.control.send_frame(signed).await
}
pub async fn advertise(
&mut self,
spec: &frame::AdvertiseSpec,
identity: &KeyPair,
) -> Result<(), SendFrameError> {
let signed = frame::sign(frame::advertise(spec), identity);
self.control.send_frame(signed).await
}
pub async fn unadvertise(
&mut self,
spec: &frame::UnadvertiseSpec,
identity: &KeyPair,
) -> Result<(), SendFrameError> {
let signed = frame::sign(frame::unadvertise(spec), identity);
self.control.send_frame(signed).await
}
pub async fn keep_advertised<F>(
&mut self,
spec: &frame::AdvertiseSpec,
identity: &KeyPair,
interval: Duration,
stop: F,
on_error: impl Fn(SendFrameError),
) where
F: std::future::Future<Output = ()>,
{
tokio::pin!(stop);
let mut ticker = tokio::time::interval(interval);
loop {
tokio::select! {
_ = &mut stop => return,
_ = ticker.tick() => {
if let Err(e) = self.advertise(spec, identity).await {
on_error(e);
}
}
}
}
}
pub async fn recv_event(
&mut self,
timeout: Duration,
) -> Result<frame::EventInfo, RecvEventError> {
let value = self
.control
.recv_frame_timeout(timeout)
.await
.map_err(RecvEventError::Recv)?;
frame::parse_event(&value).map_err(RecvEventError::Parse)
}
pub async fn serve_one_call<L>(
&mut self,
lookup: L,
identity: &KeyPair,
timeout: Duration,
) -> Result<(), ServeCallError>
where
L: Fn(&[u8; 32], &str) -> Option<CallHandler>,
{
self.serve_one_call_gated(
lookup,
|_, _| crate::ucan::Policy::open(),
identity,
timeout,
)
.await
}
pub async fn serve_one_call_gated<L, P>(
&mut self,
lookup: L,
policy: P,
identity: &KeyPair,
timeout: Duration,
) -> Result<(), ServeCallError>
where
L: Fn(&[u8; 32], &str) -> Option<CallHandler>,
P: Fn(&[u8; 32], &str) -> crate::ucan::Policy,
{
tokio::time::timeout(
timeout,
self.serve_one_call_gated_inner(lookup, policy, identity),
)
.await
.unwrap_or(Err(ServeCallError::Timeout))
}
async fn serve_one_call_gated_inner<L, P>(
&mut self,
lookup: L,
policy: P,
identity: &KeyPair,
) -> Result<(), ServeCallError>
where
L: Fn(&[u8; 32], &str) -> Option<CallHandler>,
P: Fn(&[u8; 32], &str) -> crate::ucan::Policy,
{
loop {
let value = self
.control
.recv_frame()
.await
.map_err(ServeCallError::Recv)?;
let Ok(call_info) = frame::parse_call(&value) else {
continue; };
let reply =
build_call_reply(call_info, &lookup, &policy, identity, Some(&mut *self)).await;
let signed = frame::sign(reply, identity);
self.control
.send_frame(signed)
.await
.map_err(ServeCallError::Send)?;
return Ok(());
}
}
const CLOSE_DRAIN: Duration = Duration::from_millis(250);
pub async fn close(mut self, reason: &str, detail: Option<&str>, identity: &KeyPair) {
let goodbye = frame::sign(frame::goodbye(reason, detail), identity);
if let Ok(encoded) = frame::encode(&goodbye) {
let _ = self.control.send.write_all(&encoded).await;
}
let _ = self.control.send.finish();
tokio::time::sleep(Self::CLOSE_DRAIN).await;
self.connection.close(0u32.into(), reason.as_bytes());
}
pub async fn run_publisher(
&mut self,
spec: &frame::PublishSpec,
identity: &KeyPair,
announce: bool,
) -> Result<(), SendFrameError> {
let publish_id: [u8; 16] = rand::random();
if announce {
let payload = Value::Map(vec![])
.with_field("publish_id", Value::Bytes(publish_id.to_vec()))
.with_field("topic", Value::Bytes(spec.topic.as_bytes().to_vec()));
let fact = frame::PublishSpec::new(
"pubsub.publish_started_v1",
spec.realm,
identity.node_id(),
rand::random(),
payload,
now_ms(),
);
let _ = self.publish(&fact, identity).await;
}
let result = self.publish(spec, identity).await;
if announce {
let payload =
Value::Map(vec![]).with_field("publish_id", Value::Bytes(publish_id.to_vec()));
let payload = match &result {
Ok(()) => payload.with_field("outcome", Value::text("completed")),
Err(e) => payload
.with_field("outcome", Value::text("failed"))
.with_field("reason", Value::text(e.to_string())),
};
let fact = frame::PublishSpec::new(
"pubsub.publish_completed_v1",
spec.realm,
identity.node_id(),
rand::random(),
payload,
now_ms(),
);
let _ = self.publish(&fact, identity).await;
}
result
}
pub async fn run_subscriber<F>(
&mut self,
spec: &frame::SubscribeSpec,
identity: &KeyPair,
stop: F,
mut handler: impl FnMut(frame::EventInfo),
) -> Result<(), RunSubscriberError>
where
F: std::future::Future<Output = ()>,
{
self.subscribe(spec, identity)
.await
.map_err(RunSubscriberError::Subscribe)?;
tokio::pin!(stop);
let result = loop {
tokio::select! {
_ = &mut stop => break Ok(()),
frame_result = self.control.recv_frame() => {
let value = match frame_result {
Ok(v) => v,
Err(e) => break Err(RunSubscriberError::Recv(e)),
};
let Ok(evt) = frame::parse_event(&value) else {
continue; };
handler(evt);
}
}
};
let _ = self
.unsubscribe(
&frame::UnsubscribeSpec::new(spec.topic.clone(), spec.realm, spec.subscriber),
identity,
)
.await;
result
}
}
fn now_ms() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.expect("system clock before 1970")
.as_millis() as u64
}
#[derive(Debug)]
pub enum RecvEventError {
Recv(RecvFrameError),
Parse(frame::ParseEventError),
}
impl std::fmt::Display for RecvEventError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
RecvEventError::Recv(e) => write!(f, "{e}"),
RecvEventError::Parse(e) => write!(f, "expected an EVENT frame: {e}"),
}
}
}
impl std::error::Error for RecvEventError {}
#[derive(Debug)]
pub enum ServeCallError {
Recv(RecvFrameError),
Send(SendFrameError),
Timeout,
}
impl std::fmt::Display for ServeCallError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ServeCallError::Recv(e) => write!(f, "{e}"),
ServeCallError::Send(e) => write!(f, "sending the reply: {e}"),
ServeCallError::Timeout => write!(f, "timed out waiting for an inbound CALL"),
}
}
}
impl std::error::Error for ServeCallError {}
#[derive(Debug)]
pub enum RunSubscriberError {
Subscribe(SendFrameError),
Recv(RecvFrameError),
}
impl std::fmt::Display for RunSubscriberError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
RunSubscriberError::Subscribe(e) => write!(f, "subscribing: {e}"),
RunSubscriberError::Recv(e) => write!(f, "{e}"),
}
}
}
impl std::error::Error for RunSubscriberError {}
const RPC_SENT_TOPIC: &str = "rpc.sent_v1";
const RPC_COMPLETED_TOPIC: &str = "rpc.completed_v1";
const RPC_RECEIVED_TOPIC: &str = "rpc.received_v1";
const RPC_REPLIED_TOPIC: &str = "rpc.replied_v1";
fn request_id_payload(request_id: [u8; 16]) -> Value {
Value::Map(vec![]).with_field("request_id", Value::Bytes(request_id.to_vec()))
}
async fn announce_fact(
session: &mut Session,
realm: [u8; 32],
identity: &KeyPair,
topic: &str,
payload: Value,
) {
let fact = frame::PublishSpec::new(
topic,
realm,
identity.node_id(),
rand::random(),
payload,
now_ms(),
);
let _ = session.publish(&fact, identity).await;
}
async fn announce_rpc_sent(
session: &mut Session,
realm: [u8; 32],
identity: &KeyPair,
request_id: [u8; 16],
) {
announce_fact(
session,
realm,
identity,
RPC_SENT_TOPIC,
request_id_payload(request_id),
)
.await;
}
async fn announce_rpc_completed(
session: &mut Session,
realm: [u8; 32],
identity: &KeyPair,
request_id: [u8; 16],
result: &Result<frame::CallResponse, CallError>,
) {
let payload = request_id_payload(request_id);
let payload = match result {
Err(e) => payload
.with_field("outcome", Value::text("failed"))
.with_field("reason", Value::text(e.to_string())),
Ok(frame::CallResponse::Error { name, .. }) => payload
.with_field("outcome", Value::text("failed"))
.with_field("reason", Value::text(name.clone())),
Ok(frame::CallResponse::Result { .. }) => {
payload.with_field("outcome", Value::text("completed"))
}
};
announce_fact(session, realm, identity, RPC_COMPLETED_TOPIC, payload).await;
}
async fn announce_rpc_received(
session: &mut Session,
realm: [u8; 32],
identity: &KeyPair,
request_id: [u8; 16],
) {
announce_fact(
session,
realm,
identity,
RPC_RECEIVED_TOPIC,
request_id_payload(request_id),
)
.await;
}
async fn announce_rpc_replied(
session: &mut Session,
realm: [u8; 32],
identity: &KeyPair,
request_id: [u8; 16],
handler_err: Option<&str>,
) {
let payload = request_id_payload(request_id);
let payload = match handler_err {
Some(reason) => payload
.with_field("outcome", Value::text("failed"))
.with_field("reason", Value::text(reason)),
None => payload.with_field("outcome", Value::text("replied")),
};
announce_fact(session, realm, identity, RPC_REPLIED_TOPIC, payload).await;
}
#[allow(clippy::needless_option_as_deref)]
async fn build_call_reply<L, P>(
call_info: frame::CallInfo,
lookup: &L,
policy: &P,
identity: &KeyPair,
mut session: Option<&mut Session>,
) -> Value
where
L: Fn(&[u8; 32], &str) -> Option<CallHandler>,
P: Fn(&[u8; 32], &str) -> crate::ucan::Policy,
{
let self_pub = identity.node_id();
if policy(&call_info.realm, &call_info.procedure)
.check(&call_info.ucan_token)
.is_err()
{
return frame::call_error(&frame::CallErrorSpec::new(
call_info.call_id,
bolt4::Code::Unauthorized,
self_pub,
));
}
let Some(handler) = lookup(&call_info.realm, &call_info.procedure) else {
return frame::call_error(&frame::CallErrorSpec::new(
call_info.call_id,
bolt4::Code::UnknownNextPeer,
self_pub,
));
};
let request_id: [u8; 16] = rand::random();
if let Some(s) = session.as_deref_mut() {
announce_rpc_received(s, call_info.realm, identity, request_id).await;
}
let payload = call_info.payload;
let outcome = tokio::spawn(async move { handler(payload).await }).await;
match outcome {
Ok(Ok(value)) => {
if let Some(s) = session.as_deref_mut() {
announce_rpc_replied(s, call_info.realm, identity, request_id, None).await;
}
frame::result(&frame::ResultSpec::new(call_info.call_id, value, self_pub))
}
Ok(Err(reason)) => {
if let Some(s) = session.as_deref_mut() {
announce_rpc_replied(s, call_info.realm, identity, request_id, Some(&reason)).await;
}
let mut spec =
frame::CallErrorSpec::new(call_info.call_id, bolt4::Code::UnknownError, self_pub);
spec.detail = Some(reason);
frame::call_error(&spec)
}
Err(_join_error) => frame::call_error(&frame::CallErrorSpec::new(
call_info.call_id,
bolt4::Code::TemporaryRelayFailure,
self_pub,
)),
}
}
#[cfg(test)]
mod ucan_gating_tests {
use super::*;
use crate::identity::KeyPair;
use crate::ucan::{self, Policy};
fn call_info(ucan_token: Vec<u8>) -> frame::CallInfo {
frame::CallInfo {
call_id: [1; 16],
procedure: "test.proc".into(),
realm: [0; 32],
payload: Value::Null,
deadline_ms: 0,
caller: [2; 32],
ucan_token,
}
}
fn never_called_lookup() -> impl Fn(&[u8; 32], &str) -> Option<CallHandler> {
|_, _| panic!("handler lookup must not run when policy rejects the call")
}
fn echo_lookup() -> impl Fn(&[u8; 32], &str) -> Option<CallHandler> {
|_, _| {
Some(Arc::new(|payload: Value| {
Box::pin(async move { Ok(payload) })
}))
}
}
#[tokio::test]
async fn open_policy_never_gates_dispatch() {
let identity = KeyPair::generate();
let reply = build_call_reply(
call_info(Vec::new()),
&echo_lookup(),
&|_, _| Policy::open(),
&identity,
None,
)
.await;
assert!(matches!(
frame::parse_call_response(&reply),
Ok(frame::CallResponse::Result { .. })
));
}
#[tokio::test]
async fn required_policy_refuses_a_call_with_no_token_before_lookup_runs() {
let id = KeyPair::generate();
let identity = KeyPair::generate();
let reply = build_call_reply(
call_info(Vec::new()),
&never_called_lookup(),
&move |_, _| Policy::required(id.node_id()),
&identity,
None,
)
.await;
match frame::parse_call_response(&reply) {
Ok(frame::CallResponse::Error { code, .. }) => {
assert_eq!(code, bolt4::Code::Unauthorized as u8)
}
other => panic!("expected an Unauthorized ERROR frame, got {other:?}"),
}
}
#[tokio::test]
async fn required_policy_refuses_a_token_from_the_wrong_issuer_before_lookup_runs() {
let required_issuer = KeyPair::generate();
let impostor = KeyPair::generate();
let bad_token = ucan::create(
"did:iss",
"did:aud",
vec![],
&impostor,
ucan::CreateOpts::default(),
)
.unwrap();
let identity = KeyPair::generate();
let reply = build_call_reply(
call_info(bad_token),
&never_called_lookup(),
&move |_, _| Policy::required(required_issuer.node_id()),
&identity,
None,
)
.await;
match frame::parse_call_response(&reply) {
Ok(frame::CallResponse::Error { code, .. }) => {
assert_eq!(code, bolt4::Code::Unauthorized as u8)
}
other => panic!("expected an Unauthorized ERROR frame, got {other:?}"),
}
}
#[tokio::test]
async fn required_policy_lets_a_valid_token_reach_the_handler() {
let id = KeyPair::generate();
let good_token = ucan::create(
"did:iss",
"did:aud",
vec![],
&id,
ucan::CreateOpts::default(),
)
.unwrap();
let identity = KeyPair::generate();
let reply = build_call_reply(
call_info(good_token),
&echo_lookup(),
&move |_, _| Policy::required(id.node_id()),
&identity,
None,
)
.await;
assert!(matches!(
frame::parse_call_response(&reply),
Ok(frame::CallResponse::Result { .. })
));
}
}