use std::sync::Arc;
use mshr::{Endpoint, EndpointAddr, NodeId};
use crate::lan_pair::enroll::EnrollmentSink;
use crate::lan_pair::{
AccepterMsg, ConfirmationStrategy, LanPairError, PairAccept, PairOffer, ALPN, MAX_FRAME,
};
pub struct Accepter {
endpoint: Endpoint,
enrollment_sink: Option<Arc<dyn EnrollmentSink>>,
}
impl Accepter {
pub fn new(endpoint: Endpoint) -> Self {
Self { endpoint, enrollment_sink: None }
}
pub fn with_enrollment_sink(mut self, sink: Arc<dyn EnrollmentSink>) -> Self {
self.enrollment_sink = Some(sink);
self
}
pub fn node_id(&self) -> NodeId {
self.endpoint.node_id()
}
pub async fn pair<S, A>(
&self,
addr: A,
strategy: &S,
) -> Result<PairAccept, LanPairError>
where
S: ConfirmationStrategy,
A: Into<EndpointAddr>,
{
let conn = self
.endpoint
.connect_alpn(addr, ALPN)
.await
.map_err(|e| LanPairError::Transport(e.to_string()))?;
let (mut send, mut recv) = conn
.accept_bi()
.await
.map_err(|e| LanPairError::Transport(e.to_string()))?;
let offer_bytes = recv
.read_to_end(MAX_FRAME)
.await
.map_err(|e| LanPairError::Transport(e.to_string()))?;
let offer: PairOffer = serde_json::from_slice(&offer_bytes)
.map_err(|e| LanPairError::Codec(e.to_string()))?;
let authenticated_node_id = conn.remote_id();
if offer.node_id != *authenticated_node_id.as_bytes() {
return Err(LanPairError::NodeIdMismatch);
}
let decision = strategy.confirm(&offer).await?;
let msg = match decision.as_ref() {
Some(accept) => AccepterMsg::Accept(accept.clone()),
None => AccepterMsg::Reject { reason: "declined".into() },
};
let resp_bytes =
serde_json::to_vec(&msg).map_err(|e| LanPairError::Codec(e.to_string()))?;
send.write_all(&resp_bytes)
.await
.map_err(|e| LanPairError::Transport(e.to_string()))?;
send.finish()
.map_err(|e| LanPairError::Transport(e.to_string()))?;
let _ = conn.closed().await;
if let (Some(accept), Some(sink)) = (decision.as_ref(), self.enrollment_sink.as_ref()) {
sink.enroll(&accept.user_id, authenticated_node_id).await?;
}
match decision {
Some(accept) => Ok(accept),
None => Err(LanPairError::Rejected),
}
}
}