1//! A single-threaded executor whose only blocking point is a Postgres
2//! `WaitEventSet` (contract §2): the latch, postmaster death, and every
3//! socket a pending future is waiting on. A cancel or terminate sets the
4//! latch, so `CHECK_FOR_INTERRUPTS` runs promptly and its ERROR unwinds
5//! through here, dropping every future on the way (cancellation = Drop).
6
7use std::cell::{Cell, RefCell};
8use std::cmp::Reverse;
9use std::collections::{BinaryHeap, HashMap, VecDeque};
10use std::future::Future;
11use std::os::fd::RawFd;
12use std::pin::{Pin, pin};
13use std::sync::{Arc, Mutex};
14use std::task::{Context, Poll, Wake, Waker};
15use std::time::{Duration, Instant};
16
17use pgrx::pg_sys;
18
19type TaskId = usize;
20type ScopeId = u64;
21
22/// A spawned future, and the scope that owns it if any.
23struct Task {
24    future: Pin<Box<dyn Future<Output = ()>>>,
25    scope: Option<ScopeId>,
26}
27
28/// The future passed to `block_on`, which lives on its caller's stack.
29const ROOT: TaskId = usize::MAX;
30
31thread_local! {
32    static RUNTIME: Runtime = Runtime::default();
33}
34
35/// A task's place in the runtime. `Running` is distinct from `Vacant`:
36/// a task spawned while another is being polled must not take its slot.
37enum Slot {
38    Vacant,
39    Running,
40    Waiting(Task),
41}
42
43#[derive(Default)]
44struct Runtime {
45    tasks: RefCell<Vec<Slot>>,
46    /// One waker per slot, made once: the same task always presents the
47    /// same waker, so `Waker::will_wake` means what it says. A slot's
48    /// waker outlives its task; a stale wake is only a spurious poll.
49    wakers: RefCell<Vec<Waker>>,
50    ready: Arc<ReadyQueue>,
51    reactor: RefCell<Reactor>,
52    running: Cell<bool>,
53    /// Live scopes, by the subtransaction each was opened in.
54    scopes: RefCell<HashMap<ScopeId, pg_sys::SubTransactionId>>,
55    next_scope: Cell<ScopeId>,
56}
57
58/// Wakers must be `Send + Sync`; nothing here ever runs on another
59/// thread, so the mutex is never contended.
60#[derive(Default)]
61struct ReadyQueue(Mutex<VecDeque<TaskId>>);
62
63impl ReadyQueue {
64    fn push(&self, id: TaskId) {
65        self.0.lock().unwrap().push_back(id);
66    }
67
68    fn pop(&self) -> Option<TaskId> {
69        self.0.lock().unwrap().pop_front()
70    }
71
72    fn len(&self) -> usize {
73        self.0.lock().unwrap().len()
74    }
75
76    fn is_empty(&self) -> bool {
77        self.len() == 0
78    }
79}
80
81struct TaskWaker {
82    id: TaskId,
83    ready: Arc<ReadyQueue>,
84}
85
86impl Wake for TaskWaker {
87    fn wake(self: Arc<Self>) {
88        self.ready.push(self.id);
89    }
90
91    fn wake_by_ref(self: &Arc<Self>) {
92        self.ready.push(self.id);
93    }
94}
95
96fn waker(rt: &Runtime, id: TaskId) -> Waker {
97    Waker::from(Arc::new(TaskWaker { id, ready: rt.ready.clone() }))
98}
99
100/// The backend thread. Sockets and timers are `Send + Sync` because
101/// hickory's traits demand it, but readiness lives in this thread's
102/// reactor, so registering from any other thread would never be woken.
103static BACKEND: std::sync::OnceLock<std::thread::ThreadId> = std::sync::OnceLock::new();
104
105fn assert_backend_thread() {
106    let here = std::thread::current().id();
107    debug_assert_eq!(*BACKEND.get_or_init(|| here), here, "executor used off the backend thread");
108}
109
110/// Runs `future` to completion on this backend, waiting on Postgres.
111pub fn block_on<F: Future>(future: F) -> F::Output {
112    assert_backend_thread();
113    register_abort_flush();
114    RUNTIME.with(|rt| {
115        assert!(!rt.running.replace(true), "block_on is not reentrant");
116        let _running = Reset(&rt.running);
117        let mut future = pin!(future);
118        let root = waker(rt, ROOT);
119        rt.ready.push(ROOT);
120        loop {
121            // Poll only what was ready when the pass began: a task that
122            // wakes itself runs again next pass, after the reactor has
123            // fired timers and checked for interrupts, so it cannot starve
124            // a cancel or a timeout.
125            for _ in 0..rt.ready.len() {
126                let Some(id) = rt.ready.pop() else { break };
127                if id == ROOT {
128                    if let Poll::Ready(out) = future.as_mut().poll(&mut Context::from_waker(&root)) {
129                        return out;
130                    }
131                } else {
132                    poll_task(rt, id);
133                }
134            }
135            let block = rt.ready.is_empty();
136            rt.reactor.borrow_mut().wait(block);
137        }
138    })
139}
140
141struct Reset<'a>(&'a Cell<bool>);
142
143impl Drop for Reset<'_> {
144    fn drop(&mut self) {
145        self.0.set(false);
146    }
147}
148
149/// Spawns a task that lives on this backend until it completes. It is
150/// polled only while some `block_on` runs.
151pub fn spawn_local(future: impl Future<Output = ()> + 'static) {
152    spawn(Box::pin(future), None);
153}
154
155fn spawn(future: Pin<Box<dyn Future<Output = ()>>>, scope: Option<ScopeId>) {
156    assert_backend_thread();
157    RUNTIME.with(|rt| {
158        let mut tasks = rt.tasks.borrow_mut();
159        let task = Slot::Waiting(Task { future, scope });
160        let id = match tasks.iter().position(|s| matches!(s, Slot::Vacant)) {
161            Some(free) => {
162                tasks[free] = task;
163                free
164            }
165            None => {
166                tasks.push(task);
167                let id = tasks.len() - 1;
168                rt.wakers.borrow_mut().push(waker(rt, id));
169                id
170            }
171        };
172        rt.ready.push(id);
173    });
174}
175
176fn poll_task(rt: &Runtime, id: TaskId) {
177    // Taken out while polled, because the task may spawn others.
178    let mut task = {
179        let mut tasks = rt.tasks.borrow_mut();
180        match std::mem::replace(&mut tasks[id], Slot::Running) {
181            Slot::Waiting(task) => task,
182            other => {
183                tasks[id] = other;
184                return; // finished, or a stale wake
185            }
186        }
187    };
188    // If the task panics (or an ERROR unwinds through it), its slot must
189    // not stay `Running` forever; the task is dropped with the unwind.
190    struct Vacate<'a>(&'a Runtime, TaskId);
191    impl Drop for Vacate<'_> {
192        fn drop(&mut self) {
193            if let Ok(mut tasks) = self.0.tasks.try_borrow_mut() {
194                tasks[self.1] = Slot::Vacant;
195            }
196        }
197    }
198    let vacate = Vacate(rt, id);
199    let task_waker = rt.wakers.borrow()[id].clone();
200    let poll = task.future.as_mut().poll(&mut Context::from_waker(&task_waker));
201    std::mem::forget(vacate);
202    rt.tasks.borrow_mut()[id] = match poll {
203        Poll::Pending => Slot::Waiting(task),
204        Poll::Ready(()) => Slot::Vacant,
205    };
206}
207
208#[derive(Clone, Copy)]
209pub enum Interest {
210    Read,
211    Write,
212}
213
214/// Parks the current task until `fd` is ready for `interest`.
215pub fn register(fd: RawFd, interest: Interest, waker: &Waker) {
216    assert_backend_thread();
217    RUNTIME.with(|rt| {
218        let mut reactor = rt.reactor.borrow_mut();
219        let slot = reactor.sockets.entry(fd).or_default();
220        let w = match interest {
221            Interest::Read => &mut slot.read,
222            Interest::Write => &mut slot.write,
223        };
224        // Replacing is right: every fd has one owner (`TcpIo` and `UdpIo`
225        // are not `Clone`), so the newest registration is the only live
226        // waiter. An older waker belongs to a previous owner of the same
227        // socket, e.g. the root future that ran the TLS handshake before
228        // hyper's h2 driver task took the stream over.
229        *w = Some(waker.clone());
230    });
231}
232
233/// Forgets `fd`; its owner calls this before closing it, so no closed
234/// descriptor ever reaches `AddWaitEventToSet`.
235pub fn deregister(fd: RawFd) {
236    // try_with: sockets may be dropped during thread-local teardown.
237    let _ = RUNTIME.try_with(|rt| rt.reactor.borrow_mut().sockets.remove(&fd));
238}
239
240#[derive(Default)]
241struct Reactor {
242    sockets: HashMap<RawFd, Wakers>,
243    /// Deadlines, earliest first. An entry whose id is no longer in
244    /// `timer_wakers` belongs to a dropped `Sleep` and is skipped.
245    timers: BinaryHeap<Reverse<(Instant, u64)>>,
246    timer_wakers: HashMap<u64, Waker>,
247    next_timer: u64,
248}
249
250#[derive(Default)]
251struct Wakers {
252    read: Option<Waker>,
253    write: Option<Waker>,
254}
255
256impl Reactor {
257    /// One wait on a set built for it: there is no `RemoveWaitEvent`, so a
258    /// long-lived set could never drop a closed socket (contract §2).
259    /// Milliseconds until the earliest live deadline, or -1 for none.
260    fn timeout_ms(&mut self) -> i64 {
261        while let Some(Reverse((at, id))) = self.timers.peek().copied() {
262            if !self.timer_wakers.contains_key(&id) {
263                self.timers.pop();
264                continue;
265            }
266            let left = at.saturating_duration_since(Instant::now());
267            // Round up, so a wake never comes before the deadline.
268            return left.as_micros().div_ceil(1000) as i64;
269        }
270        -1
271    }
272
273    fn fire_timers(&mut self) {
274        let now = Instant::now();
275        while let Some(Reverse((at, id))) = self.timers.peek().copied() {
276            if at > now {
277                break;
278            }
279            self.timers.pop();
280            if let Some(w) = self.timer_wakers.remove(&id) {
281                w.wake();
282            }
283        }
284    }
285
286    /// Waits for the next event, or only looks (timeout 0) when `block`
287    /// is false because tasks are still ready.
288    fn wait(&mut self, block: bool) {
289        let waiting: Vec<(RawFd, u32)> = self
290            .sockets
291            .iter()
292            .filter_map(|(&fd, w)| {
293                let events = w.read.as_ref().map_or(0, |_| pg_sys::WL_SOCKET_READABLE)
294                    | w.write.as_ref().map_or(0, |_| pg_sys::WL_SOCKET_WRITEABLE);
295                (events != 0).then_some((fd, events))
296            })
297            .collect();
298
299        let set = WaitSet::new(waiting.len() + 2);
300        set.add(pg_sys::WL_LATCH_SET, None);
301        set.add(pg_sys::WL_EXIT_ON_PM_DEATH, None);
302        for &(fd, events) in &waiting {
303            set.add(events, Some(fd));
304        }
305
306        let mut occurred = vec![pg_sys::WaitEvent::default(); waiting.len() + 2];
307        let timeout = if block { self.timeout_ms() } else { 0 };
308        let n = set.wait(&mut occurred, timeout);
309        drop(set);
310        self.fire_timers();
311
312        let mut latch = false;
313        for event in &occurred[..n] {
314            if event.events & pg_sys::WL_LATCH_SET != 0 {
315                latch = true;
316            }
317            let Some(w) = self.sockets.get_mut(&event.fd) else { continue };
318            // Closed and errored sockets report readable/writeable; the
319            // woken read or write then sees the error.
320            if event.events & pg_sys::WL_SOCKET_READABLE != 0
321                && let Some(waker) = w.read.take()
322            {
323                waker.wake();
324            }
325            if event.events & pg_sys::WL_SOCKET_WRITEABLE != 0
326                && let Some(waker) = w.write.take()
327            {
328                waker.wake();
329            }
330        }
331        if latch {
332            // SAFETY: MyLatch is set for every backend before any
333            // extension code runs, and only this thread touches it.
334            unsafe { pg_sys::ResetLatch(pg_sys::MyLatch) };
335        }
336        // After ResetLatch, so a cancel arriving now sets the latch again
337        // and the next wait returns at once instead of sleeping on it.
338        pgrx::check_for_interrupts!();
339    }
340}
341
342/// A `WaitEventSet` freed on drop, including while an ERROR unwinds.
343struct WaitSet(*mut pg_sys::WaitEventSet);
344
345impl WaitSet {
346    fn new(capacity: usize) -> Self {
347        // SAFETY: a NULL resource owner is allowed (latch.c) and makes the
348        // set ours to free, which Drop does exactly once.
349        WaitSet(unsafe { pg_sys::CreateWaitEventSet(std::ptr::null_mut(), capacity as i32) })
350    }
351
352    fn add(&self, events: u32, fd: Option<RawFd>) {
353        let latch = if events & pg_sys::WL_LATCH_SET != 0 {
354            // SAFETY: read of a backend-global set at startup.
355            unsafe { pg_sys::MyLatch }
356        } else {
357            std::ptr::null_mut()
358        };
359        // SAFETY: the set was sized for every event added (Reactor::wait
360        // counts them), and each fd is open: owners deregister before
361        // closing (net::TcpIo's Drop).
362        unsafe {
363            pg_sys::AddWaitEventToSet(self.0, events, fd.unwrap_or(-1), latch, std::ptr::null_mut());
364        }
365    }
366
367    fn wait(&self, occurred: &mut [pg_sys::WaitEvent], timeout_ms: i64) -> usize {
368        // SAFETY: `occurred` has room for `len` events; timeout -1 waits
369        // until an event, which always includes the latch.
370        let n = unsafe {
371            pg_sys::WaitEventSetWait(
372                self.0,
373                timeout_ms as _,
374                occurred.as_mut_ptr(),
375                occurred.len() as i32,
376                pg_sys::PG_WAIT_EXTENSION,
377            )
378        };
379        n as usize
380    }
381}
382
383impl Drop for WaitSet {
384    fn drop(&mut self) {
385        // SAFETY: created by CreateWaitEventSet with no resource owner, so
386        // nothing else frees it.
387        unsafe { pg_sys::FreeWaitEventSet(self.0) };
388    }
389}
390
391/// Completes once `duration` has passed. Waiting happens in the same
392/// `WaitEventSet`, so it stays cancellable.
393pub fn sleep(duration: Duration) -> Sleep {
394    sleep_until(Instant::now() + duration)
395}
396
397/// Completes at `deadline`, as [`sleep`].
398pub fn sleep_until(deadline: Instant) -> Sleep {
399    Sleep { deadline, id: None }
400}
401
402pub struct Sleep {
403    deadline: Instant,
404    id: Option<u64>,
405}
406
407impl Future for Sleep {
408    type Output = ();
409
410    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<()> {
411        if Instant::now() >= self.deadline {
412            return Poll::Ready(());
413        }
414        let deadline = self.deadline;
415        let id = RUNTIME.with(|rt| {
416            let mut reactor = rt.reactor.borrow_mut();
417            let id = match self.id {
418                Some(id) => id,
419                None => {
420                    let id = reactor.next_timer;
421                    reactor.next_timer += 1;
422                    reactor.timers.push(Reverse((deadline, id)));
423                    id
424                }
425            };
426            reactor.timer_wakers.insert(id, cx.waker().clone());
427            id
428        });
429        self.id = Some(id);
430        Poll::Pending
431    }
432}
433
434impl Drop for Sleep {
435    fn drop(&mut self) {
436        if let Some(id) = self.id {
437            let _ = RUNTIME.try_with(|rt| rt.reactor.borrow_mut().timer_wakers.remove(&id));
438        }
439    }
440}
441
442/// `future`, or `None` if `duration` passes first (dropping `future`).
443pub async fn timeout<F: Future>(duration: Duration, future: F) -> Option<F::Output> {
444    let mut future = pin!(future);
445    let mut sleep = pin!(sleep(duration));
446    std::future::poll_fn(|cx| {
447        if let Poll::Ready(out) = future.as_mut().poll(cx) {
448            return Poll::Ready(Some(out));
449        }
450        sleep.as_mut().poll(cx).map(|()| None)
451    })
452    .await
453}
454
455/// Uniform in [0, 1) from Postgres's own PRNG, seeded per backend.
456pub fn random_unit() -> f64 {
457    // SAFETY: pg_global_prng_state is seeded at backend start and used
458    // only from this thread.
459    unsafe { pg_sys::pg_prng_double(&raw mut pg_sys::pg_global_prng_state) }
460}
461
462/// A cancel unwinds out of `block_on`, dropping the in-flight futures.
463/// Dropping an h2 request only queues its RST_STREAM and wakes the
464/// connection task; with no `block_on` running, the reset would wait for
465/// the next statement while the server holds the stream open. So on
466/// every (sub)transaction abort the ready tasks are polled once more.
467fn register_abort_flush() {
468    thread_local! {
469        static REGISTERED: Cell<bool> = const { Cell::new(false) };
470    }
471    if REGISTERED.replace(true) {
472        return;
473    }
474    // SAFETY: registered once per backend and never unregistered; the
475    // callbacks are `extern "C-unwind"` and never unwind (see below).
476    unsafe {
477        pg_sys::RegisterXactCallback(Some(on_xact), std::ptr::null_mut());
478        pg_sys::RegisterSubXactCallback(Some(on_subxact), std::ptr::null_mut());
479    }
480}
481
482/// xact.h's `TopSubTransactionId`, a macro bindgen does not carry. Every
483/// scope was opened at or above it, so a top-level abort closes them all.
484const TOP_SUBTRANSACTION: pg_sys::SubTransactionId = 1;
485
486unsafe extern "C-unwind" fn on_xact(event: pg_sys::XactEvent::Type, _arg: *mut std::ffi::c_void) {
487    if matches!(event, pg_sys::XactEvent::XACT_EVENT_ABORT | pg_sys::XactEvent::XACT_EVENT_PARALLEL_ABORT) {
488        flush_after_abort(TOP_SUBTRANSACTION);
489    }
490}
491
492unsafe extern "C-unwind" fn on_subxact(
493    event: pg_sys::SubXactEvent::Type,
494    my: pg_sys::SubTransactionId,
495    _parent: pg_sys::SubTransactionId,
496    _arg: *mut std::ffi::c_void,
497) {
498    if event == pg_sys::SubXactEvent::SUBXACT_EVENT_ABORT_SUB {
499        flush_after_abort(my);
500    }
501}
502
503/// Post-abort callbacks may not fail: an ERROR or a panic here takes the
504/// cluster down. Neither step makes a Postgres call, and any panic is
505/// swallowed; the worst case is a reset sent at the next statement.
506///
507/// The aborted (sub)transaction's scopes go first: their owners are
508/// freed only later, with the executor's memory, and their requests'
509/// resets must be in the flush.
510fn flush_after_abort(aborted: pg_sys::SubTransactionId) {
511    let _ = std::panic::catch_unwind(|| {
512        let _ = RUNTIME.try_with(|rt| {
513            let doomed: Vec<ScopeId> = rt
514                .scopes
515                .borrow()
516                .iter()
517                .filter(|&(_, &opened)| opened >= aborted)
518                .map(|(&id, _)| id)
519                .collect();
520            for id in doomed {
521                close_scope(rt, id);
522            }
523        });
524        flush();
525    });
526}
527
528/// Polls the ready tasks without ever waiting, so whatever they can
529/// write now (resets, GOAWAY) is written. Bounded, in case tasks keep
530/// waking each other.
531fn flush() {
532    let _ = RUNTIME.try_with(|rt| {
533        if rt.running.get() {
534            return;
535        }
536        for _ in 0..256 {
537            match rt.ready.pop() {
538                Some(ROOT) => {} // the root future is gone with its block_on
539                Some(id) => poll_task(rt, id),
540                None => break,
541            }
542        }
543    });
544}
545
546/// Tasks that belong to one owner (a scan) and end with it: dropping the
547/// scope drops every task it spawned, and aborting the subtransaction it
548/// was opened in drops them sooner (contract §2, "Task scopes"). Either
549/// way the dropped requests' resets are written at once.
550pub struct Scope {
551    id: ScopeId,
552}
553
554impl Scope {
555    pub fn new() -> Self {
556        assert_backend_thread();
557        // SAFETY: called on the backend inside a transaction (an executor
558        // callback); it only reads the current transaction state.
559        let opened = unsafe { pg_sys::GetCurrentSubTransactionId() };
560        RUNTIME.with(|rt| {
561            let id = rt.next_scope.get();
562            rt.next_scope.set(id + 1);
563            rt.scopes.borrow_mut().insert(id, opened);
564            Scope { id }
565        })
566    }
567
568    /// Spawns `future` in this scope. After an abort has closed the scope
569    /// it is dropped unpolled: nothing may start for a failed statement.
570    pub fn spawn(&self, future: impl Future<Output = ()> + 'static) {
571        if RUNTIME.with(|rt| rt.scopes.borrow().contains_key(&self.id)) {
572            spawn(Box::pin(future), Some(self.id));
573        }
574    }
575
576    /// Drops every task spawned so far; the scope stays open.
577    pub fn cancel(&self) {
578        RUNTIME.with(|rt| {
579            let opened = rt.scopes.borrow().get(&self.id).copied();
580            close_scope(rt, self.id);
581            if let Some(opened) = opened {
582                rt.scopes.borrow_mut().insert(self.id, opened);
583            }
584        });
585        flush();
586    }
587}
588
589impl Drop for Scope {
590    fn drop(&mut self) {
591        let _ = RUNTIME.try_with(|rt| close_scope(rt, self.id));
592        flush();
593    }
594}
595
596/// Forgets scope `id` and drops its tasks. The futures are dropped after
597/// the task table is released, because dropping one deregisters sockets
598/// and timers.
599fn close_scope(rt: &Runtime, id: ScopeId) {
600    rt.scopes.borrow_mut().remove(&id);
601    let Ok(mut tasks) = rt.tasks.try_borrow_mut() else { return };
602    let dropped: Vec<Task> = tasks
603        .iter_mut()
604        .filter(|slot| matches!(slot, Slot::Waiting(task) if task.scope == Some(id)))
605        .map(|slot| match std::mem::replace(slot, Slot::Vacant) {
606            Slot::Waiting(task) => task,
607            _ => unreachable!("filtered to waiting tasks"),
608        })
609        .collect();
610    drop(tasks);
611    drop(dropped);
612}
613
614/// Every future's output, in order, polling them all concurrently.
615pub async fn join_all<F: Future>(futures: Vec<F>) -> Vec<F::Output> {
616    let mut futures: Vec<Pin<Box<F>>> = futures.into_iter().map(Box::pin).collect();
617    let mut outputs: Vec<Option<F::Output>> = futures.iter().map(|_| None).collect();
618    std::future::poll_fn(|cx| {
619        let mut pending = false;
620        for (future, output) in futures.iter_mut().zip(outputs.iter_mut()) {
621            if output.is_none() {
622                match future.as_mut().poll(cx) {
623                    Poll::Ready(out) => *output = Some(out),
624                    Poll::Pending => pending = true,
625                }
626            }
627        }
628        if pending { Poll::Pending } else { Poll::Ready(()) }
629    })
630    .await;
631    outputs.into_iter().map(|o| o.expect("every future completed")).collect()
632}
633
634/// Returns `Pending` once, so the executor completes a reactor pass (at
635/// timeout 0) and polls whatever it woke before the caller resumes.
636pub async fn yield_now() {
637    let mut yielded = false;
638    std::future::poll_fn(|cx| {
639        if yielded {
640            return Poll::Ready(());
641        }
642        yielded = true;
643        cx.waker().wake_by_ref();
644        Poll::Pending
645    })
646    .await
647}