use std::sync::{Arc, Mutex};
use std::time::Duration;
use anyhow::{anyhow, bail, Result};
use async_trait::async_trait;
use openrtc::native::{AssertionProvider, IdentityAssertion};
use serde::Serialize;
use tokio::sync::{oneshot, watch};
const ASSERTION_TIMEOUT: Duration = Duration::from_secs(30);
type Sink = Arc<dyn Fn(NativeAssertionRequest) -> Result<(), String> + Send + Sync>;
type Reply = oneshot::Sender<Result<IdentityAssertion>>;
#[derive(Clone, Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct NativeAssertionRequest {
pub request_id: String,
pub device_id: String,
pub force_refresh: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub recovery_reason: Option<&'static str>,
}
struct State {
sink: Sink,
pending: Option<(String, Reply)>,
retired: bool,
}
pub struct NativeAssertionBridge {
session_key: String,
state: Mutex<State>,
epoch: watch::Sender<u64>,
}
impl NativeAssertionBridge {
pub fn new(
session_key: String,
sink: impl Fn(NativeAssertionRequest) -> Result<(), String> + Send + Sync + 'static,
) -> Result<Self> {
if session_key.trim().is_empty() || session_key.len() > 512 {
bail!("native assertion session key is invalid");
}
Ok(Self {
session_key,
state: Mutex::new(State {
sink: Arc::new(sink),
pending: None,
retired: false,
}),
epoch: watch::channel(0).0,
})
}
pub fn rebind(
&self,
sink: impl Fn(NativeAssertionRequest) -> Result<(), String> + Send + Sync + 'static,
) -> Result<()> {
let mut state = self
.state
.lock()
.map_err(|_| anyhow!("native assertion lock poisoned"))?;
if state.retired {
bail!("native assertion identity is retired");
}
if let Some((_, reply)) = state.pending.take() {
let _ = reply.send(Err(anyhow!("native assertion delivery was replaced")));
}
state.sink = Arc::new(sink);
Ok(())
}
pub fn complete(
&self,
request_id: &str,
result: Result<IdentityAssertion, String>,
) -> Result<bool> {
let mut state = self
.state
.lock()
.map_err(|_| anyhow!("native assertion lock poisoned"))?;
if state.retired
|| !state
.pending
.as_ref()
.is_some_and(|(id, _)| id == request_id)
{
return Ok(false);
}
let (_, reply) = state.pending.take().expect("matching pending request");
let result = result
.map_err(|error| anyhow!(error))
.and_then(|assertion| {
if assertion.token.trim().is_empty() || assertion.token.len() > 64 * 1024 {
bail!("host returned an invalid identity assertion");
}
Ok(assertion)
});
Ok(reply.send(result).is_ok())
}
pub fn retire(&self) -> Result<()> {
let mut state = self
.state
.lock()
.map_err(|_| anyhow!("native assertion lock poisoned"))?;
if !state.retired {
state.retired = true;
self.epoch.send_replace(1);
if let Some((_, reply)) = state.pending.take() {
let _ = reply.send(Err(anyhow!("native assertion identity is retired")));
}
}
Ok(())
}
async fn request(
&self,
device: &str,
refresh: bool,
recovery: bool,
) -> Result<IdentityAssertion> {
if device.trim().is_empty() || device.len() > 192 {
bail!("native assertion requires a durable device id");
}
let id = uuid::Uuid::new_v4().to_string();
let (reply, receiver) = oneshot::channel();
let guard = PendingRequest {
bridge: self,
id: &id,
};
let sink = {
let mut state = self
.state
.lock()
.map_err(|_| anyhow!("native assertion lock poisoned"))?;
if state.retired {
bail!("native assertion identity is retired");
}
if state.pending.is_some() {
bail!("native assertion request already pending");
}
state.pending = Some((id.clone(), reply));
state.sink.clone()
};
sink(NativeAssertionRequest {
request_id: id.clone(),
device_id: device.into(),
force_refresh: refresh,
recovery_reason: recovery.then_some("device-key-rotation"),
})
.map_err(|error| anyhow!(error))?;
let result = tokio::time::timeout(ASSERTION_TIMEOUT, receiver)
.await
.map_err(|_| anyhow!("native assertion request timed out"))?
.map_err(|_| anyhow!("native assertion response channel closed"))?;
drop(guard);
result
}
}
struct PendingRequest<'a> {
bridge: &'a NativeAssertionBridge,
id: &'a str,
}
impl Drop for PendingRequest<'_> {
fn drop(&mut self) {
if let Ok(mut state) = self.bridge.state.lock() {
if state.pending.as_ref().is_some_and(|(id, _)| id == self.id) {
state.pending = None;
}
}
}
}
#[async_trait]
impl AssertionProvider for NativeAssertionBridge {
async fn assertion(&self, _: bool) -> Result<IdentityAssertion> {
bail!("native assertion requires a device-bound request")
}
async fn assertion_for_device(&self, refresh: bool, device: &str) -> Result<IdentityAssertion> {
self.request(device, refresh, false).await
}
async fn assertion_for_device_recovery(&self, device: &str) -> Result<IdentityAssertion> {
self.request(device, true, true).await
}
fn session_key(&self) -> Result<Option<String>> {
let state = self
.state
.lock()
.map_err(|_| anyhow!("native assertion lock poisoned"))?;
Ok((!state.retired).then(|| self.session_key.clone()))
}
fn identity_epoch(&self) -> u64 {
*self.epoch.borrow()
}
fn subscribe_identity_epoch(&self) -> Option<watch::Receiver<u64>> {
Some(self.epoch.subscribe())
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::sync::mpsc;
fn fixture() -> (
Arc<NativeAssertionBridge>,
mpsc::UnboundedReceiver<NativeAssertionRequest>,
) {
let (sender, receiver) = mpsc::unbounded_channel();
let bridge = NativeAssertionBridge::new("user-session".into(), move |request| {
sender.send(request).map_err(|_| "fixture closed".into())
})
.unwrap();
(Arc::new(bridge), receiver)
}
fn assertion() -> Result<IdentityAssertion, String> {
Ok(IdentityAssertion {
token: "fixture-assertion".into(),
provider_id: Some("fixture".into()),
})
}
async fn next(
receiver: &mut mpsc::UnboundedReceiver<NativeAssertionRequest>,
) -> NativeAssertionRequest {
tokio::time::timeout(Duration::from_secs(2), receiver.recv())
.await
.unwrap()
.unwrap()
}
#[tokio::test]
async fn device_bound_assertion_is_single_flight_and_exactly_once() {
let (bridge, mut requests) = fixture();
let pending = tokio::spawn({
let bridge = bridge.clone();
async move { bridge.assertion_for_device(false, "device").await }
});
let request = next(&mut requests).await;
assert_eq!(request.device_id, "device");
assert!(!request.force_refresh);
assert_eq!(request.recovery_reason, None);
assert!(bridge.assertion_for_device(false, "device").await.is_err());
assert!(
requests.try_recv().is_err(),
"no queued retry or second host call"
);
assert!(!bridge.complete("old-request", assertion()).unwrap());
assert!(bridge.complete(&request.request_id, assertion()).unwrap());
assert!(!bridge.complete(&request.request_id, assertion()).unwrap());
assert_eq!(pending.await.unwrap().unwrap().token, "fixture-assertion");
assert_eq!(bridge.identity_epoch(), 0);
}
#[tokio::test]
async fn reload_rebinds_delivery_without_changing_identity_and_fences_old_reply() {
let (bridge, mut old_requests) = fixture();
let pending = tokio::spawn({
let bridge = bridge.clone();
async move { bridge.assertion_for_device(false, "device").await }
});
let old = next(&mut old_requests).await;
let (sender, mut requests) = mpsc::unbounded_channel();
bridge
.rebind(move |request| sender.send(request).map_err(|_| "fixture closed".into()))
.unwrap();
assert!(pending.await.unwrap().is_err());
assert_eq!(bridge.identity_epoch(), 0);
assert_eq!(
bridge.session_key().unwrap().as_deref(),
Some("user-session")
);
let recovery = tokio::spawn({
let bridge = bridge.clone();
async move { bridge.assertion_for_device_recovery("device").await }
});
let request = next(&mut requests).await;
assert!(request.force_refresh);
assert_eq!(request.recovery_reason, Some("device-key-rotation"));
assert!(!bridge.complete(&old.request_id, assertion()).unwrap());
assert!(bridge.complete(&request.request_id, assertion()).unwrap());
recovery.await.unwrap().unwrap();
}
#[tokio::test]
async fn logout_retires_pending_and_future_assertions_and_notifies_rust() {
let (bridge, mut requests) = fixture();
let mut epoch = bridge.subscribe_identity_epoch().unwrap();
let pending = tokio::spawn({
let bridge = bridge.clone();
async move { bridge.assertion_for_device(false, "device").await }
});
let request = next(&mut requests).await;
bridge.retire().unwrap();
epoch.changed().await.unwrap();
assert_eq!(*epoch.borrow(), 1);
assert!(pending.await.unwrap().is_err());
assert!(!bridge.complete(&request.request_id, assertion()).unwrap());
assert!(bridge.assertion_for_device(false, "device").await.is_err());
assert!(bridge.rebind(|_| Ok(())).is_err());
assert!(bridge.session_key().unwrap().is_none());
bridge.retire().unwrap();
assert_eq!(bridge.identity_epoch(), 1);
assert!(requests.try_recv().is_err());
}
#[tokio::test]
async fn cancelled_or_failed_delivery_does_not_leave_a_pending_owner() {
let (bridge, mut requests) = fixture();
let pending = tokio::spawn({
let bridge = bridge.clone();
async move { bridge.assertion_for_device(false, "device").await }
});
let old = next(&mut requests).await;
pending.abort();
assert!(pending.await.unwrap_err().is_cancelled());
assert!(bridge.state.lock().unwrap().pending.is_none());
bridge.rebind(|_| Err("host closed".into())).unwrap();
assert!(bridge.assertion_for_device(false, "device").await.is_err());
assert!(bridge.state.lock().unwrap().pending.is_none());
assert!(!bridge.complete(&old.request_id, assertion()).unwrap());
}
#[tokio::test]
async fn missing_host_response_expires_without_retrying() {
let (bridge, mut requests) = fixture();
let pending = tokio::spawn({
let bridge = bridge.clone();
async move { bridge.assertion_for_device(false, "device").await }
});
let request = next(&mut requests).await;
let result = tokio::time::timeout(ASSERTION_TIMEOUT + Duration::from_secs(2), pending)
.await
.unwrap()
.unwrap();
assert_eq!(
result.unwrap_err().to_string(),
"native assertion request timed out"
);
assert!(bridge.state.lock().unwrap().pending.is_none());
assert!(!bridge.complete(&request.request_id, assertion()).unwrap());
assert!(requests.try_recv().is_err());
}
}