Skip to main content

hara_native/task/
promise.rs

1use std::cell::{Cell, RefCell};
2use std::collections::{HashSet, VecDeque};
3use std::rc::{Rc, Weak};
4use std::time::{Duration, Instant};
5
6use crate::core::Value;
7
8#[derive(Debug, Clone, PartialEq)]
9pub enum PromiseRejection {
10    Message(String),
11    Value(Value),
12    Cancelled(Value),
13}
14
15impl PromiseRejection {
16    pub fn value(&self) -> Value {
17        match self {
18            Self::Message(message) => Value::String(message.clone()),
19            Self::Value(value) | Self::Cancelled(value) => value.clone(),
20        }
21    }
22
23    pub fn message(&self) -> String {
24        match self {
25            Self::Message(message) => message.clone(),
26            Self::Value(value) | Self::Cancelled(value) => value.display(),
27        }
28    }
29
30    pub fn is_cancelled(&self) -> bool {
31        match self {
32            Self::Cancelled(_) => true,
33            Self::Message(message) => message == "cancelled",
34            Self::Value(_) => false,
35        }
36    }
37
38    pub fn cancelled() -> Self {
39        Self::Cancelled(Value::Map(
40            [
41                (
42                    Value::Keyword("code".into()),
43                    Value::Keyword("task/cancelled".into()),
44                ),
45                (
46                    Value::Keyword("message".into()),
47                    Value::String("cancelled".into()),
48                ),
49                (
50                    Value::Keyword("origin".into()),
51                    Value::Keyword("runtime".into()),
52                ),
53                (Value::Keyword("retryable".into()), Value::Bool(false)),
54            ]
55            .into_iter()
56            .collect(),
57        ))
58    }
59}
60
61impl From<String> for PromiseRejection {
62    fn from(value: String) -> Self {
63        Self::Message(value)
64    }
65}
66
67impl From<&str> for PromiseRejection {
68    fn from(value: &str) -> Self {
69        Self::Message(value.into())
70    }
71}
72
73#[derive(Debug, Clone, PartialEq)]
74pub enum PromiseState {
75    Pending,
76    Fulfilled(Value),
77    Rejected(PromiseRejection),
78}
79
80#[derive(Default)]
81struct PromiseHooks {
82    poller: Option<Rc<dyn Fn()>>,
83    waiter: Option<Rc<dyn Fn()>>,
84    cancel: Option<Rc<dyn Fn()>>,
85}
86
87struct PromiseInner {
88    state: PromiseState,
89    continuations: Vec<Rc<dyn Fn(PromiseState)>>,
90    deferred: Option<(Instant, Rc<dyn Fn() -> Result<Value, String>>)>,
91    hooks: PromiseHooks,
92    adopted_from: Option<Weak<RefCell<PromiseInner>>>,
93}
94
95type ContinuationJob = (Rc<dyn Fn(PromiseState)>, PromiseState);
96
97thread_local! {
98    static CONTINUATION_QUEUE: RefCell<VecDeque<ContinuationJob>> = RefCell::new(VecDeque::new());
99    static DRAINING_CONTINUATIONS: Cell<bool> = const { Cell::new(false) };
100}
101
102fn enqueue_continuation(continuation: Rc<dyn Fn(PromiseState)>, state: PromiseState) {
103    CONTINUATION_QUEUE.with(|queue| queue.borrow_mut().push_back((continuation, state)));
104    DRAINING_CONTINUATIONS.with(|draining| {
105        if draining.replace(true) {
106            return;
107        }
108        loop {
109            let job = CONTINUATION_QUEUE.with(|queue| queue.borrow_mut().pop_front());
110            let Some((continuation, state)) = job else {
111                break;
112            };
113            continuation(state);
114        }
115        draining.set(false);
116    });
117}
118
119#[derive(Clone)]
120pub struct Promise {
121    inner: Rc<RefCell<PromiseInner>>,
122}
123
124#[derive(Clone)]
125pub(crate) struct WeakPromise {
126    inner: Weak<RefCell<PromiseInner>>,
127}
128
129impl WeakPromise {
130    pub(crate) fn upgrade(&self) -> Option<Promise> {
131        self.inner.upgrade().map(|inner| Promise { inner })
132    }
133}
134
135impl std::fmt::Debug for Promise {
136    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
137        formatter
138            .debug_struct("Promise")
139            .field("state", &self.state())
140            .finish()
141    }
142}
143
144impl Default for Promise {
145    fn default() -> Self {
146        Self::new()
147    }
148}
149
150impl Promise {
151    pub fn new() -> Self {
152        Self {
153            inner: Rc::new(RefCell::new(PromiseInner {
154                state: PromiseState::Pending,
155                continuations: Vec::new(),
156                deferred: None,
157                hooks: PromiseHooks::default(),
158                adopted_from: None,
159            })),
160        }
161    }
162
163    pub fn state(&self) -> PromiseState {
164        self.run_deferred_if_ready();
165        let poller = self.inner.borrow().hooks.poller.clone();
166        if let Some(poller) = poller {
167            poller();
168        }
169        self.inner.borrow().state.clone()
170    }
171
172    fn run_deferred_if_ready(&self) {
173        let task = {
174            let mut inner = self.inner.borrow_mut();
175            if !inner
176                .deferred
177                .as_ref()
178                .is_some_and(|(at, _)| Instant::now() >= *at)
179            {
180                return;
181            }
182            inner.deferred.take().map(|(_, task)| task)
183        };
184        if let Some(task) = task {
185            settle_result(self, task());
186        }
187    }
188
189    pub fn set_poller(&self, poller: Rc<dyn Fn()>) {
190        self.inner.borrow_mut().hooks.poller = Some(poller);
191    }
192
193    pub fn set_waiter(&self, waiter: Rc<dyn Fn()>) {
194        self.inner.borrow_mut().hooks.waiter = Some(waiter);
195    }
196
197    pub fn set_cancel_hook(&self, cancel: Rc<dyn Fn()>) {
198        self.inner.borrow_mut().hooks.cancel = Some(cancel);
199    }
200
201    pub fn wait_state(&self) -> PromiseState {
202        let waiter = self.inner.borrow().hooks.waiter.clone();
203        if let Some(waiter) = waiter {
204            waiter();
205        } else {
206            #[cfg(not(target_arch = "wasm32"))]
207            if let Some(deadline) = self
208                .inner
209                .borrow()
210                .deferred
211                .as_ref()
212                .map(|(deadline, _)| *deadline)
213            {
214                if let Some(delay) = deadline.checked_duration_since(Instant::now()) {
215                    std::thread::sleep(delay);
216                }
217            }
218        }
219        loop {
220            let state = self.state();
221            if !matches!(state, PromiseState::Pending) {
222                return state;
223            }
224            let delay = {
225                let inner = self.inner.borrow();
226                inner
227                    .deferred
228                    .as_ref()
229                    .map(|(at, _)| at.saturating_duration_since(Instant::now()))
230            };
231            let Some(delay) = delay else {
232                return state;
233            };
234            if !delay.is_zero() {
235                std::thread::sleep(delay);
236            } else {
237                std::thread::yield_now();
238            }
239        }
240    }
241
242    pub fn wait_state_timeout(&self, timeout: Duration) -> PromiseState {
243        #[cfg(target_arch = "wasm32")]
244        {
245            let _ = timeout;
246            return self.state();
247        }
248        #[cfg(not(target_arch = "wasm32"))]
249        {
250            let deadline = Instant::now() + timeout;
251            loop {
252                let state = self.state();
253                if !matches!(state, PromiseState::Pending) || Instant::now() >= deadline {
254                    return state;
255                }
256                let remaining = deadline.saturating_duration_since(Instant::now());
257                std::thread::sleep(remaining.min(Duration::from_millis(1)));
258            }
259        }
260    }
261
262    pub(crate) fn notify_cancel(&self) {
263        let (cancel, adopted_from) = {
264            let inner = self.inner.borrow();
265            (inner.hooks.cancel.clone(), inner.adopted_from.clone())
266        };
267        if let Some(cancel) = cancel {
268            cancel();
269        }
270        if let Some(source) = adopted_from.and_then(|source| source.upgrade()) {
271            Promise { inner: source }.cancel();
272        }
273    }
274
275    pub fn cancel(&self) -> bool {
276        if !matches!(self.inner.borrow().state, PromiseState::Pending) {
277            return false;
278        }
279        self.notify_cancel();
280        self.reject_rejection(PromiseRejection::cancelled())
281    }
282
283    pub fn schedule(&self, delay: Duration, task: Rc<dyn Fn() -> Result<Value, String>>) {
284        if delay.is_zero() {
285            settle_result(self, task());
286        } else {
287            self.inner.borrow_mut().deferred = Some((Instant::now() + delay, task));
288        }
289    }
290
291    pub fn resolve(&self, value: Value) -> bool {
292        self.settle(PromiseState::Fulfilled(value))
293    }
294
295    pub fn reject(&self, error: impl Into<String>) -> bool {
296        self.reject_rejection(PromiseRejection::Message(error.into()))
297    }
298
299    pub fn reject_value(&self, error: Value) -> bool {
300        self.reject_rejection(PromiseRejection::Value(error))
301    }
302
303    pub fn reject_rejection(&self, error: PromiseRejection) -> bool {
304        self.settle(PromiseState::Rejected(error))
305    }
306
307    fn settle(&self, next: PromiseState) -> bool {
308        let continuations = {
309            let mut inner = self.inner.borrow_mut();
310            if !matches!(inner.state, PromiseState::Pending) {
311                return false;
312            }
313            inner.state = next.clone();
314            inner.deferred = None;
315            inner.hooks = PromiseHooks::default();
316            inner.adopted_from = None;
317            std::mem::take(&mut inner.continuations)
318        };
319        for continuation in continuations {
320            enqueue_continuation(continuation, next.clone());
321        }
322        true
323    }
324
325    pub fn on_settle(&self, continuation: Rc<dyn Fn(PromiseState)>) {
326        let state = self.state();
327        if matches!(state, PromiseState::Pending) {
328            self.inner.borrow_mut().continuations.push(continuation);
329        } else {
330            enqueue_continuation(continuation, state);
331        }
332    }
333
334    pub fn adopt(&self, other: &Promise) -> bool {
335        if self.same_identity(other) || self.adoption_would_cycle(other) {
336            return self.reject("promise adoption cycle");
337        }
338        match other.state() {
339            PromiseState::Pending => {
340                if !matches!(self.state(), PromiseState::Pending) {
341                    return false;
342                }
343                self.inner.borrow_mut().adopted_from = Some(Rc::downgrade(&other.inner));
344                let source = other.clone();
345                self.set_poller(Rc::new(move || {
346                    source.state();
347                }));
348                let source = other.clone();
349                self.set_waiter(Rc::new(move || {
350                    source.wait_state();
351                }));
352                let destination = self.clone();
353                other.on_settle(Rc::new(move |state| match state {
354                    PromiseState::Fulfilled(value) => {
355                        destination.resolve(value);
356                    }
357                    PromiseState::Rejected(error) => {
358                        destination.reject_rejection(error);
359                    }
360                    PromiseState::Pending => {}
361                }));
362                true
363            }
364            PromiseState::Fulfilled(value) => self.resolve(value),
365            PromiseState::Rejected(error) => self.reject_rejection(error),
366        }
367    }
368
369    fn adoption_would_cycle(&self, source: &Promise) -> bool {
370        let mut current = Some(source.inner.clone());
371        let mut seen = HashSet::new();
372        while let Some(inner) = current {
373            if Rc::ptr_eq(&self.inner, &inner) {
374                return true;
375            }
376            let address = Rc::as_ptr(&inner) as usize;
377            if !seen.insert(address) {
378                return true;
379            }
380            current = inner.borrow().adopted_from.as_ref().and_then(Weak::upgrade);
381        }
382        false
383    }
384
385    pub fn same_identity(&self, other: &Self) -> bool {
386        Rc::ptr_eq(&self.inner, &other.inner)
387    }
388
389    pub(crate) fn downgrade(&self) -> WeakPromise {
390        WeakPromise {
391            inner: Rc::downgrade(&self.inner),
392        }
393    }
394
395    pub fn identity_address(&self) -> usize {
396        Rc::as_ptr(&self.inner) as usize
397    }
398}
399
400pub fn settle_result(destination: &Promise, result: Result<Value, String>) {
401    match result {
402        Ok(Value::Promise(source)) => {
403            destination.adopt(&source);
404        }
405        Ok(value) => {
406            destination.resolve(value);
407        }
408        Err(error) => {
409            destination.reject(error);
410        }
411    }
412}
413
414pub trait PromiseProvider {
415    fn native(&self) -> bool;
416    fn run(&self, task: Rc<dyn Fn() -> Result<Value, String>>) -> Promise;
417    fn delay(&self, duration: Duration, task: Rc<dyn Fn() -> Result<Value, String>>) -> Promise;
418}
419
420#[derive(Debug, Clone, Copy, Default)]
421pub struct LocalPromiseProvider;
422
423impl PromiseProvider for LocalPromiseProvider {
424    fn native(&self) -> bool {
425        true
426    }
427
428    fn run(&self, task: Rc<dyn Fn() -> Result<Value, String>>) -> Promise {
429        let promise = Promise::new();
430        settle_result(&promise, task());
431        promise
432    }
433
434    fn delay(&self, duration: Duration, task: Rc<dyn Fn() -> Result<Value, String>>) -> Promise {
435        let promise = Promise::new();
436        promise.schedule(duration, task);
437        promise
438    }
439}
440
441#[cfg(test)]
442mod tests {
443    use super::*;
444
445    #[test]
446    fn cancellation_is_a_structured_rejection() {
447        let promise = Promise::new();
448        assert!(promise.cancel());
449        let PromiseState::Rejected(rejection) = promise.state() else {
450            panic!("cancelled promise was not rejected");
451        };
452        assert!(rejection.is_cancelled());
453        let Value::Map(fields) = rejection.value() else {
454            panic!("cancellation was not represented by a map");
455        };
456        assert_eq!(
457            fields.get(&Value::Keyword("code".into())),
458            Some(&Value::Keyword("task/cancelled".into()))
459        );
460        assert_eq!(
461            fields.get(&Value::Keyword("retryable".into())),
462            Some(&Value::Bool(false))
463        );
464    }
465
466    #[test]
467    fn adoption_cycles_are_rejected() {
468        let direct = Promise::new();
469        assert!(direct.adopt(&direct));
470        assert!(matches!(direct.state(), PromiseState::Rejected(_)));
471
472        let first = Promise::new();
473        let second = Promise::new();
474        assert!(first.adopt(&second));
475        assert!(second.adopt(&first));
476        assert!(matches!(second.state(), PromiseState::Rejected(_)));
477        assert!(matches!(first.state(), PromiseState::Rejected(_)));
478    }
479
480    #[test]
481    fn continuation_chains_are_trampolined() {
482        let promises: Vec<_> = (0..20_000).map(|_| Promise::new()).collect();
483        for pair in promises.windows(2) {
484            let next = pair[1].clone();
485            pair[0].on_settle(Rc::new(move |state| match state {
486                PromiseState::Fulfilled(value) => {
487                    next.resolve(value);
488                }
489                PromiseState::Rejected(error) => {
490                    next.reject_rejection(error);
491                }
492                PromiseState::Pending => {}
493            }));
494        }
495        promises[0].resolve(Value::Number(42));
496        assert_eq!(
497            promises.last().unwrap().state(),
498            PromiseState::Fulfilled(Value::Number(42))
499        );
500    }
501}