1use std::future::Future;
2use std::pin::Pin;
3use std::sync::Arc;
4
5use unb_core::NodeIdentity;
6
7use crate::handler::HandlerError;
8
9#[derive(Clone, Debug, PartialEq, Eq)]
10pub struct VerifiedPeer {
11 node_id: String,
12 instance_id: String,
13 epoch: u64,
14 declared_only: bool,
15}
16
17impl VerifiedPeer {
18 pub(crate) fn from_identity(identity: &NodeIdentity) -> Self {
19 Self {
20 node_id: identity.node_id.clone(),
21 instance_id: identity.instance_id.clone(),
22 epoch: identity.epoch,
23 declared_only: false,
24 }
25 }
26
27 pub fn node_id(&self) -> &str {
28 &self.node_id
29 }
30
31 pub fn instance_id(&self) -> &str {
32 &self.instance_id
33 }
34
35 pub fn epoch(&self) -> u64 {
36 self.epoch
37 }
38
39 pub fn declared_only(&self) -> bool {
40 self.declared_only
41 }
42}
43
44pub struct PeerRequest {
45 local: NodeIdentity,
46 remote: NodeIdentity,
47 extensions: http::Extensions,
48}
49
50impl PeerRequest {
51 pub(crate) fn new(local: NodeIdentity, remote: NodeIdentity) -> PeerRequest {
52 PeerRequest {
53 local,
54 remote,
55 extensions: http::Extensions::new(),
56 }
57 }
58
59 pub fn local(&self) -> &NodeIdentity {
60 &self.local
61 }
62
63 pub fn remote(&self) -> &NodeIdentity {
64 &self.remote
65 }
66
67 pub fn extensions(&self) -> &http::Extensions {
68 &self.extensions
69 }
70
71 pub fn extensions_mut(&mut self) -> &mut http::Extensions {
72 &mut self.extensions
73 }
74
75 pub fn accept(&mut self) {
76 self.insert_verified(false);
77 }
78
79 pub fn accept_declared(&mut self) {
80 self.insert_verified(true);
81 }
82
83 pub(crate) fn verified(&self) -> Option<VerifiedPeer> {
84 self.extensions.get::<VerifiedPeer>().cloned()
85 }
86
87 fn insert_verified(&mut self, declared_only: bool) {
88 self.extensions.insert(VerifiedPeer {
89 node_id: self.remote.node_id.clone(),
90 instance_id: self.remote.instance_id.clone(),
91 epoch: self.remote.epoch,
92 declared_only,
93 });
94 }
95}
96
97pub trait PeerLayer: Send + Sync + 'static {
98 fn admit(
99 &self,
100 request: PeerRequest,
101 next: PeerNext,
102 ) -> Pin<Box<dyn Future<Output = Result<PeerRequest, HandlerError>> + Send + '_>>;
103}
104
105pub struct PeerNext {
106 layers: Arc<[Arc<dyn PeerLayer>]>,
107 index: usize,
108}
109
110impl PeerNext {
111 pub(crate) fn root(layers: Arc<[Arc<dyn PeerLayer>]>) -> PeerNext {
112 PeerNext { layers, index: 0 }
113 }
114
115 pub async fn admit(mut self, request: PeerRequest) -> Result<PeerRequest, HandlerError> {
116 if self.index < self.layers.len() {
117 let layer = self.layers[self.index].clone();
118 self.index += 1;
119 layer.admit(request, self).await
120 } else {
121 Ok(request)
122 }
123 }
124}
125
126pub struct InsecureAcceptDeclaredPeerIdentities;
127
128impl PeerLayer for InsecureAcceptDeclaredPeerIdentities {
129 fn admit(
130 &self,
131 mut request: PeerRequest,
132 next: PeerNext,
133 ) -> Pin<Box<dyn Future<Output = Result<PeerRequest, HandlerError>> + Send + '_>> {
134 request.accept_declared();
135 Box::pin(next.admit(request))
136 }
137}
138
139pub struct PeerLayerFn<F>(F);
140
141impl<F, Fut> PeerLayer for PeerLayerFn<F>
142where
143 F: Fn(PeerRequest, PeerNext) -> Fut + Send + Sync + 'static,
144 Fut: Future<Output = Result<PeerRequest, HandlerError>> + Send + 'static,
145{
146 fn admit(
147 &self,
148 request: PeerRequest,
149 next: PeerNext,
150 ) -> Pin<Box<dyn Future<Output = Result<PeerRequest, HandlerError>> + Send + '_>> {
151 Box::pin((self.0)(request, next))
152 }
153}
154
155pub fn peer_layer_fn<F, Fut>(f: F) -> PeerLayerFn<F>
156where
157 F: Fn(PeerRequest, PeerNext) -> Fut + Send + Sync + 'static,
158 Fut: Future<Output = Result<PeerRequest, HandlerError>> + Send + 'static,
159{
160 PeerLayerFn(f)
161}
162
163#[cfg(test)]
164mod tests {
165 use super::*;
166 use serde_json::Value;
167 use unb_core::NodeIdentity;
168
169 fn identity(node: &str) -> NodeIdentity {
170 NodeIdentity {
171 node_id: node.into(),
172 instance_id: format!("{node}-i"),
173 epoch: 1,
174 proof: Value::Null,
175 }
176 }
177
178 fn chain(layer: impl PeerLayer) -> PeerNext {
179 let layers: Arc<[Arc<dyn PeerLayer>]> =
180 Arc::from(vec![Arc::new(layer) as Arc<dyn PeerLayer>]);
181 PeerNext::root(layers)
182 }
183
184 #[tokio::test]
185 async fn a_closure_that_accepts_and_continues_admits() {
186 let next = chain(peer_layer_fn(
187 |mut request: PeerRequest, next: PeerNext| async move {
188 request.accept();
189 next.admit(request).await
190 },
191 ));
192 let admitted = next
193 .admit(PeerRequest::new(identity("local"), identity("remote")))
194 .await
195 .unwrap();
196 let verified = admitted.verified().unwrap();
197 assert!(!verified.declared_only());
198 }
199
200 #[tokio::test]
201 async fn a_closure_that_continues_without_accepting_produces_no_verified_peer() {
202 let next = chain(peer_layer_fn(
203 |request: PeerRequest, next: PeerNext| async move { next.admit(request).await },
204 ));
205 let admitted = next
206 .admit(PeerRequest::new(identity("local"), identity("remote")))
207 .await
208 .unwrap();
209 assert!(admitted.verified().is_none());
210 }
211
212 #[tokio::test]
213 async fn a_closure_that_errors_rejects() {
214 let next = chain(peer_layer_fn(
215 |_request: PeerRequest, _next: PeerNext| async move {
216 Err(HandlerError::new(crate::ErrorCode::Unauthorized, "denied"))
217 },
218 ));
219 let result = next
220 .admit(PeerRequest::new(identity("local"), identity("remote")))
221 .await;
222 assert!(result.is_err());
223 }
224}