1use std::sync::{Arc, Mutex};
4use std::time::Duration;
5
6use anyhow::{anyhow, bail, Result};
7use async_trait::async_trait;
8use openrtc::native::{AssertionProvider, IdentityAssertion};
9use serde::Serialize;
10use tokio::sync::{oneshot, watch};
11
12const ASSERTION_TIMEOUT: Duration = Duration::from_secs(30);
13type Sink = Arc<dyn Fn(NativeAssertionRequest) -> Result<(), String> + Send + Sync>;
14type Reply = oneshot::Sender<Result<IdentityAssertion>>;
15
16#[derive(Clone, Debug, Serialize)]
17#[serde(rename_all = "camelCase")]
18pub struct NativeAssertionRequest {
19 pub request_id: String,
20 pub device_id: String,
21 pub force_refresh: bool,
22 #[serde(skip_serializing_if = "Option::is_none")]
23 pub recovery_reason: Option<&'static str>,
24}
25
26struct State {
27 sink: Sink,
28 pending: Option<(String, Reply)>,
29 retired: bool,
30}
31
32pub struct NativeAssertionBridge {
36 session_key: String,
37 state: Mutex<State>,
38 epoch: watch::Sender<u64>,
39}
40
41impl NativeAssertionBridge {
42 pub fn new(
43 session_key: String,
44 sink: impl Fn(NativeAssertionRequest) -> Result<(), String> + Send + Sync + 'static,
45 ) -> Result<Self> {
46 if session_key.trim().is_empty() || session_key.len() > 512 {
47 bail!("native assertion session key is invalid");
48 }
49 Ok(Self {
50 session_key,
51 state: Mutex::new(State {
52 sink: Arc::new(sink),
53 pending: None,
54 retired: false,
55 }),
56 epoch: watch::channel(0).0,
57 })
58 }
59
60 pub fn rebind(
61 &self,
62 sink: impl Fn(NativeAssertionRequest) -> Result<(), String> + Send + Sync + 'static,
63 ) -> Result<()> {
64 let mut state = self
65 .state
66 .lock()
67 .map_err(|_| anyhow!("native assertion lock poisoned"))?;
68 if state.retired {
69 bail!("native assertion identity is retired");
70 }
71 if let Some((_, reply)) = state.pending.take() {
72 let _ = reply.send(Err(anyhow!("native assertion delivery was replaced")));
73 }
74 state.sink = Arc::new(sink);
75 Ok(())
76 }
77
78 pub fn complete(
81 &self,
82 request_id: &str,
83 result: Result<IdentityAssertion, String>,
84 ) -> Result<bool> {
85 let mut state = self
86 .state
87 .lock()
88 .map_err(|_| anyhow!("native assertion lock poisoned"))?;
89 if state.retired
90 || !state
91 .pending
92 .as_ref()
93 .is_some_and(|(id, _)| id == request_id)
94 {
95 return Ok(false);
96 }
97 let (_, reply) = state.pending.take().expect("matching pending request");
98 let result = result
99 .map_err(|error| anyhow!(error))
100 .and_then(|assertion| {
101 if assertion.token.trim().is_empty() || assertion.token.len() > 64 * 1024 {
102 bail!("host returned an invalid identity assertion");
103 }
104 Ok(assertion)
105 });
106 Ok(reply.send(result).is_ok())
107 }
108
109 pub fn retire(&self) -> Result<()> {
110 let mut state = self
111 .state
112 .lock()
113 .map_err(|_| anyhow!("native assertion lock poisoned"))?;
114 if !state.retired {
115 state.retired = true;
116 self.epoch.send_replace(1);
117 if let Some((_, reply)) = state.pending.take() {
118 let _ = reply.send(Err(anyhow!("native assertion identity is retired")));
119 }
120 }
121 Ok(())
122 }
123
124 async fn request(
125 &self,
126 device: &str,
127 refresh: bool,
128 recovery: bool,
129 ) -> Result<IdentityAssertion> {
130 if device.trim().is_empty() || device.len() > 192 {
131 bail!("native assertion requires a durable device id");
132 }
133 let id = uuid::Uuid::new_v4().to_string();
134 let (reply, receiver) = oneshot::channel();
135 let guard = PendingRequest {
136 bridge: self,
137 id: &id,
138 };
139 let sink = {
140 let mut state = self
141 .state
142 .lock()
143 .map_err(|_| anyhow!("native assertion lock poisoned"))?;
144 if state.retired {
145 bail!("native assertion identity is retired");
146 }
147 if state.pending.is_some() {
148 bail!("native assertion request already pending");
149 }
150 state.pending = Some((id.clone(), reply));
151 state.sink.clone()
152 };
153 sink(NativeAssertionRequest {
154 request_id: id.clone(),
155 device_id: device.into(),
156 force_refresh: refresh,
157 recovery_reason: recovery.then_some("device-key-rotation"),
158 })
159 .map_err(|error| anyhow!(error))?;
160 let result = tokio::time::timeout(ASSERTION_TIMEOUT, receiver)
161 .await
162 .map_err(|_| anyhow!("native assertion request timed out"))?
163 .map_err(|_| anyhow!("native assertion response channel closed"))?;
164 drop(guard);
165 result
166 }
167}
168
169struct PendingRequest<'a> {
170 bridge: &'a NativeAssertionBridge,
171 id: &'a str,
172}
173
174impl Drop for PendingRequest<'_> {
175 fn drop(&mut self) {
176 if let Ok(mut state) = self.bridge.state.lock() {
177 if state.pending.as_ref().is_some_and(|(id, _)| id == self.id) {
178 state.pending = None;
179 }
180 }
181 }
182}
183
184#[async_trait]
185impl AssertionProvider for NativeAssertionBridge {
186 async fn assertion(&self, _: bool) -> Result<IdentityAssertion> {
187 bail!("native assertion requires a device-bound request")
188 }
189 async fn assertion_for_device(&self, refresh: bool, device: &str) -> Result<IdentityAssertion> {
190 self.request(device, refresh, false).await
191 }
192 async fn assertion_for_device_recovery(&self, device: &str) -> Result<IdentityAssertion> {
193 self.request(device, true, true).await
194 }
195 fn session_key(&self) -> Result<Option<String>> {
196 let state = self
197 .state
198 .lock()
199 .map_err(|_| anyhow!("native assertion lock poisoned"))?;
200 Ok((!state.retired).then(|| self.session_key.clone()))
201 }
202 fn identity_epoch(&self) -> u64 {
203 *self.epoch.borrow()
204 }
205 fn subscribe_identity_epoch(&self) -> Option<watch::Receiver<u64>> {
206 Some(self.epoch.subscribe())
207 }
208}
209
210#[cfg(test)]
211mod tests {
212 use super::*;
213 use tokio::sync::mpsc;
214
215 fn fixture() -> (
216 Arc<NativeAssertionBridge>,
217 mpsc::UnboundedReceiver<NativeAssertionRequest>,
218 ) {
219 let (sender, receiver) = mpsc::unbounded_channel();
220 let bridge = NativeAssertionBridge::new("user-session".into(), move |request| {
221 sender.send(request).map_err(|_| "fixture closed".into())
222 })
223 .unwrap();
224 (Arc::new(bridge), receiver)
225 }
226
227 fn assertion() -> Result<IdentityAssertion, String> {
228 Ok(IdentityAssertion {
229 token: "fixture-assertion".into(),
230 provider_id: Some("fixture".into()),
231 })
232 }
233
234 async fn next(
235 receiver: &mut mpsc::UnboundedReceiver<NativeAssertionRequest>,
236 ) -> NativeAssertionRequest {
237 tokio::time::timeout(Duration::from_secs(2), receiver.recv())
238 .await
239 .unwrap()
240 .unwrap()
241 }
242
243 #[tokio::test]
244 async fn device_bound_assertion_is_single_flight_and_exactly_once() {
245 let (bridge, mut requests) = fixture();
246 let pending = tokio::spawn({
247 let bridge = bridge.clone();
248 async move { bridge.assertion_for_device(false, "device").await }
249 });
250 let request = next(&mut requests).await;
251 assert_eq!(request.device_id, "device");
252 assert!(!request.force_refresh);
253 assert_eq!(request.recovery_reason, None);
254 assert!(bridge.assertion_for_device(false, "device").await.is_err());
255 assert!(
256 requests.try_recv().is_err(),
257 "no queued retry or second host call"
258 );
259 assert!(!bridge.complete("old-request", assertion()).unwrap());
260 assert!(bridge.complete(&request.request_id, assertion()).unwrap());
261 assert!(!bridge.complete(&request.request_id, assertion()).unwrap());
262 assert_eq!(pending.await.unwrap().unwrap().token, "fixture-assertion");
263 assert_eq!(bridge.identity_epoch(), 0);
264 }
265
266 #[tokio::test]
267 async fn reload_rebinds_delivery_without_changing_identity_and_fences_old_reply() {
268 let (bridge, mut old_requests) = fixture();
269 let pending = tokio::spawn({
270 let bridge = bridge.clone();
271 async move { bridge.assertion_for_device(false, "device").await }
272 });
273 let old = next(&mut old_requests).await;
274 let (sender, mut requests) = mpsc::unbounded_channel();
275 bridge
276 .rebind(move |request| sender.send(request).map_err(|_| "fixture closed".into()))
277 .unwrap();
278 assert!(pending.await.unwrap().is_err());
279 assert_eq!(bridge.identity_epoch(), 0);
280 assert_eq!(
281 bridge.session_key().unwrap().as_deref(),
282 Some("user-session")
283 );
284 let recovery = tokio::spawn({
285 let bridge = bridge.clone();
286 async move { bridge.assertion_for_device_recovery("device").await }
287 });
288 let request = next(&mut requests).await;
289 assert!(request.force_refresh);
290 assert_eq!(request.recovery_reason, Some("device-key-rotation"));
291 assert!(!bridge.complete(&old.request_id, assertion()).unwrap());
292 assert!(bridge.complete(&request.request_id, assertion()).unwrap());
293 recovery.await.unwrap().unwrap();
294 }
295
296 #[tokio::test]
297 async fn logout_retires_pending_and_future_assertions_and_notifies_rust() {
298 let (bridge, mut requests) = fixture();
299 let mut epoch = bridge.subscribe_identity_epoch().unwrap();
300 let pending = tokio::spawn({
301 let bridge = bridge.clone();
302 async move { bridge.assertion_for_device(false, "device").await }
303 });
304 let request = next(&mut requests).await;
305 bridge.retire().unwrap();
306 epoch.changed().await.unwrap();
307 assert_eq!(*epoch.borrow(), 1);
308 assert!(pending.await.unwrap().is_err());
309 assert!(!bridge.complete(&request.request_id, assertion()).unwrap());
310 assert!(bridge.assertion_for_device(false, "device").await.is_err());
311 assert!(bridge.rebind(|_| Ok(())).is_err());
312 assert!(bridge.session_key().unwrap().is_none());
313 bridge.retire().unwrap();
314 assert_eq!(bridge.identity_epoch(), 1);
315 assert!(requests.try_recv().is_err());
316 }
317
318 #[tokio::test]
319 async fn cancelled_or_failed_delivery_does_not_leave_a_pending_owner() {
320 let (bridge, mut requests) = fixture();
321 let pending = tokio::spawn({
322 let bridge = bridge.clone();
323 async move { bridge.assertion_for_device(false, "device").await }
324 });
325 let old = next(&mut requests).await;
326 pending.abort();
327 assert!(pending.await.unwrap_err().is_cancelled());
328 assert!(bridge.state.lock().unwrap().pending.is_none());
329 bridge.rebind(|_| Err("host closed".into())).unwrap();
330 assert!(bridge.assertion_for_device(false, "device").await.is_err());
331 assert!(bridge.state.lock().unwrap().pending.is_none());
332 assert!(!bridge.complete(&old.request_id, assertion()).unwrap());
333 }
334
335 #[tokio::test]
336 async fn missing_host_response_expires_without_retrying() {
337 let (bridge, mut requests) = fixture();
338 let pending = tokio::spawn({
339 let bridge = bridge.clone();
340 async move { bridge.assertion_for_device(false, "device").await }
341 });
342 let request = next(&mut requests).await;
343 let result = tokio::time::timeout(ASSERTION_TIMEOUT + Duration::from_secs(2), pending)
344 .await
345 .unwrap()
346 .unwrap();
347 assert_eq!(
348 result.unwrap_err().to_string(),
349 "native assertion request timed out"
350 );
351 assert!(bridge.state.lock().unwrap().pending.is_none());
352 assert!(!bridge.complete(&request.request_id, assertion()).unwrap());
353 assert!(requests.try_recv().is_err());
354 }
355}