use async_trait::async_trait;
use crate::lan_pair::LanPairError;
pub const NODE_RESOURCE_KIND: &str = "node";
pub const OWNS_RELATIONSHIP: &str = "owns";
#[async_trait]
pub trait EnrollmentSink: Send + Sync {
async fn enroll(&self, user_id: &str, node_id: mshr::NodeId) -> Result<(), LanPairError>;
}
pub struct HttpEnrollmentSink {
client: reqwest::Client,
endpoint: String,
session_bearer: String,
}
impl HttpEnrollmentSink {
pub fn new(
client: reqwest::Client,
endpoint: impl Into<String>,
session_bearer: impl Into<String>,
) -> Self {
Self {
client,
endpoint: endpoint.into(),
session_bearer: session_bearer.into(),
}
}
}
impl std::fmt::Debug for HttpEnrollmentSink {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("HttpEnrollmentSink")
.field("endpoint", &self.endpoint)
.finish_non_exhaustive()
}
}
#[derive(serde::Serialize)]
struct EnrollNodeRequestBody<'a> {
node_id: &'a str,
}
#[async_trait]
impl EnrollmentSink for HttpEnrollmentSink {
async fn enroll(&self, _user_id: &str, node_id: mshr::NodeId) -> Result<(), LanPairError> {
let node_id_hex = node_id.to_string();
let resp = self
.client
.post(&self.endpoint)
.bearer_auth(&self.session_bearer)
.json(&EnrollNodeRequestBody {
node_id: &node_id_hex,
})
.send()
.await
.map_err(|e| LanPairError::Enrollment(format!("enrollment http request: {e}")))?;
if resp.status().is_success() {
return Ok(());
}
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
Err(LanPairError::Enrollment(format!(
"enrollment route returned {status}: {body}"
)))
}
}
#[cfg(test)]
pub(crate) mod test_support {
use std::sync::Mutex;
use super::*;
#[derive(Default)]
pub(crate) struct RecordingSink {
calls: Mutex<Vec<(String, mshr::NodeId)>>,
fail: bool,
}
impl RecordingSink {
pub(crate) fn new() -> Self {
Self::default()
}
pub(crate) fn failing() -> Self {
Self { calls: Mutex::new(Vec::new()), fail: true }
}
pub(crate) fn calls(&self) -> Vec<(String, mshr::NodeId)> {
self.calls.lock().unwrap().clone()
}
}
#[async_trait]
impl EnrollmentSink for RecordingSink {
async fn enroll(&self, user_id: &str, node_id: mshr::NodeId) -> Result<(), LanPairError> {
self.calls.lock().unwrap().push((user_id.to_string(), node_id));
if self.fail {
Err(LanPairError::Enrollment("test-forced enrollment failure".into()))
} else {
Ok(())
}
}
}
}
#[cfg(test)]
mod http_tests {
use std::collections::HashMap;
use std::sync::Arc;
use wiremock::matchers::{body_json, header, method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
use super::*;
use crate::lan_pair::confirm::AutoTrust;
use crate::lan_pair::{Accepter, Offerer, ALPN};
use mshr::{Endpoint, Keypair};
const TEST_BEARER: &str = "test-session-bearer-abc123";
fn test_accept() -> crate::lan_pair::PairAccept {
crate::lan_pair::PairAccept {
user_id: "test-user-001".to_string(),
device_id: "test-device-rpi".to_string(),
attrs: HashMap::new(),
expires_at: i64::MAX,
}
}
#[tokio::test]
async fn enroll_posts_only_node_id_with_the_session_bearer_no_other_fields() {
let server = MockServer::start().await;
let real_node_id: mshr::NodeId = Keypair::generate().node_id();
let node_id = real_node_id.to_string();
Mock::given(method("POST"))
.and(path("/enrollment/node"))
.and(header("authorization", format!("Bearer {TEST_BEARER}").as_str()))
.and(body_json(serde_json::json!({ "node_id": node_id })))
.respond_with(ResponseTemplate::new(201).set_body_json(serde_json::json!({
"id": "own-1",
"principal_id": "user:test-user-001",
"resource_kind": "node",
"resource_id": node_id,
"relationship": "owns",
"granted_by": "svc:cheers-enrollment",
"on_behalf_of": "user:test-user-001",
"granted_at": 1_000,
"revoked_at": null,
})))
.expect(1)
.mount(&server)
.await;
let sink = HttpEnrollmentSink::new(
reqwest::Client::new(),
format!("{}/enrollment/node", server.uri()),
TEST_BEARER,
);
sink.enroll("test-user-001", real_node_id)
.await
.expect("enroll succeeds against the mock 201");
server.verify().await;
}
#[tokio::test]
async fn enroll_surfaces_non_success_status_as_enrollment_error() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/enrollment/node"))
.respond_with(ResponseTemplate::new(401).set_body_string("unauthorized"))
.mount(&server)
.await;
let sink = HttpEnrollmentSink::new(
reqwest::Client::new(),
format!("{}/enrollment/node", server.uri()),
"an-expired-or-invalid-bearer",
);
let node_id = Keypair::generate().node_id();
let err = sink
.enroll("test-user-001", node_id)
.await
.expect_err("401 must surface as an error, not Ok");
assert!(matches!(err, LanPairError::Enrollment(_)), "got {err:?}");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn completed_pair_posts_the_authenticated_node_id_via_the_enrollment_route() {
let server = MockServer::start().await;
let offerer_ep = Endpoint::builder()
.keypair(Keypair::generate())
.alpns([ALPN])
.bind()
.await
.expect("offerer bind");
let accepter_ep = Endpoint::builder()
.keypair(Keypair::generate())
.alpns([ALPN])
.bind()
.await
.expect("accepter bind");
let addr = offerer_ep.endpoint_addr();
let offerer_node_id_hex = offerer_ep.node_id().to_string();
Mock::given(method("POST"))
.and(path("/enrollment/node"))
.and(header("authorization", format!("Bearer {TEST_BEARER}").as_str()))
.and(body_json(
serde_json::json!({ "node_id": offerer_node_id_hex }),
))
.respond_with(ResponseTemplate::new(201).set_body_json(serde_json::json!({
"id": "own-1",
"principal_id": "user:test-user-001",
"resource_kind": "node",
"resource_id": offerer_node_id_hex,
"relationship": "owns",
"granted_by": "svc:cheers-enrollment",
"on_behalf_of": "user:test-user-001",
"granted_at": 1_000,
"revoked_at": null,
})))
.expect(1)
.mount(&server)
.await;
let sink = Arc::new(HttpEnrollmentSink::new(
reqwest::Client::new(),
format!("{}/enrollment/node", server.uri()),
TEST_BEARER,
));
let accepter = Accepter::new(accepter_ep).with_enrollment_sink(sink);
let offerer = Offerer::new(offerer_ep);
let strategy = AutoTrust { accept: test_accept() };
let (o, a) = tokio::join!(offerer.wait_for_pair(), accepter.pair(addr, &strategy));
o.expect("offerer completes pair");
a.expect("accepter completes pair and the HTTP enrollment POST succeeds");
server.verify().await;
}
}