use std::sync::Arc;
use cheers_core::{Credential, CredentialStore, DeviceBinding, DeviceId, UserId};
use mshr::{Endpoint, EndpointAddr, NodeId};
use crate::lan_pair::{AccepterMsg, LanPairError, PairAccept, PairOffer, MAX_FRAME};
pub struct Offerer {
endpoint: Endpoint,
code: Option<String>,
credential_store: Option<Arc<dyn CredentialStore>>,
}
impl Offerer {
pub fn new(endpoint: Endpoint) -> Self {
Self { endpoint, code: None, credential_store: None }
}
pub fn with_credential_store(mut self, store: Arc<dyn CredentialStore>) -> Self {
self.credential_store = Some(store);
self
}
pub fn with_code(mut self, code: impl Into<String>) -> Self {
self.code = Some(code.into());
self
}
pub fn with_random_code(mut self) -> Result<Self, LanPairError> {
let mut buf = [0u8; 3];
getrandom::fill(&mut buf).map_err(|e| LanPairError::Transport(e.to_string()))?;
let n = ((buf[0] as u32) << 16 | (buf[1] as u32) << 8 | buf[2] as u32) % 1_000_000;
self.code = Some(format!("{n:06}"));
Ok(self)
}
pub fn code(&self) -> Option<&str> {
self.code.as_deref()
}
pub fn node_id(&self) -> NodeId {
self.endpoint.node_id()
}
pub fn endpoint_addr(&self) -> EndpointAddr {
self.endpoint.endpoint_addr()
}
pub async fn wait_for_pair(&self) -> Result<PairAccept, LanPairError> {
let incoming = self
.endpoint
.accept()
.await
.ok_or(LanPairError::ConnectionClosed)?;
let conn = incoming
.await
.map_err(|e| LanPairError::Transport(e.to_string()))?;
let offer = PairOffer {
node_id: *self.endpoint.node_id().as_bytes(),
capabilities: vec!["cheers/lan-pair/v1".to_string()],
code: self.code.clone(),
};
let (mut send, mut recv) = conn
.open_bi()
.await
.map_err(|e| LanPairError::Transport(e.to_string()))?;
let offer_bytes =
serde_json::to_vec(&offer).map_err(|e| LanPairError::Codec(e.to_string()))?;
send.write_all(&offer_bytes)
.await
.map_err(|e| LanPairError::Transport(e.to_string()))?;
send.finish()
.map_err(|e| LanPairError::Transport(e.to_string()))?;
let resp_bytes = recv
.read_to_end(MAX_FRAME)
.await
.map_err(|e| LanPairError::Transport(e.to_string()))?;
let msg: AccepterMsg = serde_json::from_slice(&resp_bytes)
.map_err(|e| LanPairError::Codec(e.to_string()))?;
match msg {
AccepterMsg::Accept(accept) => {
if let Some(store) = &self.credential_store {
self.park_credential(store.as_ref(), &accept).await?;
}
Ok(accept)
}
AccepterMsg::Reject { .. } => Err(LanPairError::Rejected),
}
}
async fn park_credential(
&self,
store: &dyn CredentialStore,
accept: &PairAccept,
) -> Result<(), LanPairError> {
let material = serde_json::to_vec(accept)
.map_err(|e| LanPairError::Enrollment(format!("serialize PairAccept: {e}")))?;
let credential = Credential::new(
UserId::new(accept.user_id.clone()),
DeviceId::new(accept.device_id.clone()),
DeviceBinding::LanPair,
material,
);
store
.put(&accept.device_id, &credential)
.await
.map_err(|e| LanPairError::Enrollment(format!("CredentialStore::put: {e}")))
}
}