postjevsql.git / crates / postjevsql-pg / src / ratelimit.rs
1//! The account's rate limit, shared by every backend of the cluster
2//! (contract *Execution*, "The account's rate limit is shared
3//! cluster-wide"). `postjevsql_core::gcra` decides; this module gives it
4//! its three words in shared memory and a clock every backend reads
5//! alike.
6//!
7//! The words live in a named DSM segment (`GetNamedDSMSegment`, PG17+),
8//! so the extension need not be in `shared_preload_libraries`: the first
9//! backend to ask creates and zeroes it, and the rest attach. Zeroed
10//! words are a limiter at rest: both TATs in the past, and the adaptive
11//! scale at its ceiling.
12
13use std::cell::Cell;
14use std::ffi::c_void;
15use std::sync::atomic::AtomicU64;
16use std::time::Duration;
17
18use pgrx::pg_sys;
19use postjevsql_core::TokenRatio;
20use postjevsql_core::gcra::{Adaptive, Admission, Ceilings, Limit, Limiter};
21
22use crate::executor;
23
24/// The segment: tokens TAT, requests TAT, adaptive scale.
25#[repr(C)]
26struct Words([AtomicU64; 3]);
27
28thread_local! {
29    static WORDS: Cell<Option<&'static Words>> = const { Cell::new(None) };
30}
31
32/// Attaches this backend to the shared words, creating them if no
33/// backend has yet. Called from `begin`, never from a future.
34fn words() -> &'static Words {
35    if let Some(words) = WORDS.get() {
36        return words;
37    }
38    let mut found = false;
39    // SAFETY: called on the backend thread, outside any future, where an
40    // ERROR from dsm_registry.c longjmps into the caller's pg_guard. The
41    // segment is pinned and its mapping kept for the backend's life, so
42    // the pointer stays valid as `'static`. `init` runs under the
43    // registry's lock, before any other backend can attach.
44    let ptr = unsafe {
45        pg_sys::GetNamedDSMSegment(c"postjevsql_ratelimit".as_ptr(), size_of::<Words>(), Some(init), &mut found)
46    };
47    // SAFETY: the segment is `size_of::<Words>()` bytes, suitably aligned
48    // (DSM segments are MAXALIGNed), zeroed by `init`, and only ever
49    // accessed through atomics.
50    let words = unsafe { &*(ptr as *const Words) };
51    WORDS.set(Some(words));
52    words
53}
54
55#[pgrx::pg_guard]
56unsafe extern "C-unwind" fn init(ptr: *mut c_void) {
57    // SAFETY: the registry hands a fresh mapping of the requested size.
58    unsafe { std::ptr::write_bytes(ptr as *mut u8, 0, size_of::<Words>()) };
59}
60
61/// Nanoseconds on CLOCK_MONOTONIC, which every process on the host
62/// shares (an `Instant` cannot be compared across processes).
63fn now() -> u64 {
64    let mut ts = libc::timespec { tv_sec: 0, tv_nsec: 0 };
65    // SAFETY: a valid out-pointer; CLOCK_MONOTONIC cannot fail on Linux.
66    unsafe { libc::clock_gettime(libc::CLOCK_MONOTONIC, &mut ts) };
67    ts.tv_sec as u64 * 1_000_000_000 + ts.tv_nsec as u64
68}
69
70/// The cluster's limiter under this backend's ceilings. Built in `begin`,
71/// so the segment is attached before any future runs.
72#[derive(Clone, Copy)]
73pub struct RateLimit {
74    words: &'static Words,
75    ceilings: Ceilings,
76    /// Characters per token, learned from the cache when the scan began.
77    ratio: TokenRatio,
78}
79
80impl RateLimit {
81    /// `tokens_per_second` and `requests_per_minute` of 0 are unlimited.
82    /// `ratio` prices a request's body until its answer reports the real
83    /// count.
84    pub fn attach(tokens_per_second: f64, requests_per_minute: f64, ratio: TokenRatio) -> Self {
85        RateLimit { words: words(), ceilings: Ceilings { tokens_per_second, requests_per_minute }, ratio }
86    }
87
88    /// Tokens a request is charged before its answer reports the real
89    /// count: the body's length at the learned characters per token.
90    fn estimate(&self, body: &[u8]) -> f64 {
91        body.len() as f64 / self.ratio.chars_per_token()
92    }
93
94    fn limiter(&self) -> Limiter<'static> {
95        let [tokens, requests, adaptive] = &self.words.0;
96        Limiter { tokens: Limit(tokens), requests: Limit(requests), adaptive: Adaptive(adaptive) }
97    }
98
99    /// Waits until both limits admit a request with this `body`, and
100    /// charges it. The wait is a timer on the executor, so a cancel or
101    /// `statement_timeout` ends it like any other wait.
102    pub async fn admit(&self, body: &[u8]) {
103        let tokens = self.estimate(body);
104        loop {
105            match self.limiter().admit(now(), tokens, self.ceilings) {
106                Admission::Admitted => return,
107                // At least a millisecond, so a rounding sliver cannot spin.
108                Admission::Wait(wait) => executor::sleep(wait.max(Duration::from_millis(1))).await,
109            }
110        }
111    }
112
113    /// The server answered `status` to a request with this `body`.
114    pub fn settle(&self, body: &[u8], status: u16, response: &[u8]) {
115        let limiter = self.limiter();
116        match status {
117            200 => {
118                limiter.adaptive.succeeded();
119                if let Some(reported) = reported_input_tokens(response) {
120                    limiter.correct(self.estimate(body), reported, self.ceilings);
121                }
122            }
123            429 => limiter.throttled(self.ceilings),
124            _ => {}
125        }
126    }
127
128    /// The request never reached the server: nothing was processed, so
129    /// its tokens are returned. (The request slot stays spent; a redial
130    /// is rare and the slot is one of 1,200 a minute.)
131    pub fn unsent(&self, body: &[u8]) {
132        self.limiter().correct(self.estimate(body), 0.0, self.ceilings);
133    }
134}
135
136/// `usage.input_tokens` from an answer's body, if it has one.
137fn reported_input_tokens(body: &[u8]) -> Option<f64> {
138    let value: serde_json::Value = serde_json::from_slice(body).ok()?;
139    value.get("usage")?.get("input_tokens")?.as_f64()
140}