Skip to main content

bhtune_server/
state.rs

1//! [`AppState`]: the shared state every route handler receives via `axum::extract::State`.
2
3use bhtune_db::SqlitePool;
4use chrono::{DateTime, Duration, Utc};
5use std::collections::{HashMap, VecDeque};
6use std::sync::{Arc, RwLock};
7
8use crate::active_run::ActiveRun;
9use bhtune_cli::config::{DemoPolicy, ServerMode};
10use tokio::sync::Mutex;
11use tokio::sync::{OwnedMutexGuard, OwnedSemaphorePermit, Semaphore};
12
13#[derive(Debug, Clone, Copy, PartialEq, Eq)]
14pub struct DemoQuotaExceeded {
15    pub retry_after_secs: u64,
16}
17
18#[derive(Debug, Clone)]
19struct AcceptedStart {
20    id: u64,
21    accepted_at: DateTime<Utc>,
22}
23
24#[derive(Default)]
25struct AcceptedStartWindows {
26    next_id: u64,
27    by_token: HashMap<String, VecDeque<AcceptedStart>>,
28    by_ip: HashMap<String, VecDeque<AcceptedStart>>,
29}
30
31#[derive(Debug)]
32pub struct DemoAcceptedStart {
33    id: u64,
34    token_hash: String,
35    client_ip: String,
36}
37
38#[derive(Clone)]
39pub struct DemoRuntime {
40    run_permits: Arc<Semaphore>,
41    sse_permits: Arc<Semaphore>,
42    ordinary_request_permits: Arc<Semaphore>,
43    start_admission: Arc<Mutex<()>>,
44    accepted_starts: Arc<Mutex<AcceptedStartWindows>>,
45    visitor_run_permits: Arc<Mutex<HashMap<String, Arc<Semaphore>>>>,
46    visitor_sse_permits: Arc<Mutex<HashMap<String, Arc<Semaphore>>>>,
47}
48
49impl DemoRuntime {
50    fn new(policy: DemoPolicy) -> Self {
51        Self {
52            run_permits: Arc::new(Semaphore::new(policy.max_active_runs_global as usize)),
53            sse_permits: Arc::new(Semaphore::new(policy.max_sse_global as usize)),
54            ordinary_request_permits: Arc::new(Semaphore::new(
55                policy.ordinary_request_concurrency as usize,
56            )),
57            start_admission: Arc::new(Mutex::new(())),
58            accepted_starts: Arc::new(Mutex::new(AcceptedStartWindows::default())),
59            visitor_run_permits: Arc::new(Mutex::new(HashMap::new())),
60            visitor_sse_permits: Arc::new(Mutex::new(HashMap::new())),
61        }
62    }
63
64    pub fn try_acquire_global_run(
65        &self,
66    ) -> Result<OwnedSemaphorePermit, tokio::sync::TryAcquireError> {
67        self.run_permits.clone().try_acquire_owned()
68    }
69
70    pub fn try_acquire_global_sse(
71        &self,
72    ) -> Result<OwnedSemaphorePermit, tokio::sync::TryAcquireError> {
73        self.sse_permits.clone().try_acquire_owned()
74    }
75
76    pub fn try_acquire_ordinary_request(
77        &self,
78    ) -> Result<OwnedSemaphorePermit, tokio::sync::TryAcquireError> {
79        self.ordinary_request_permits.clone().try_acquire_owned()
80    }
81
82    /// Serializes the short admission phase that checks global row capacity and persists a
83    /// prepared run. Without this guard, several simultaneous starts could all observe one
84    /// remaining row slot before any of them inserted, exceeding the fixed global Demo cap.
85    pub async fn lock_start_admission(&self) -> OwnedMutexGuard<()> {
86        self.start_admission.clone().lock_owned().await
87    }
88
89    pub async fn try_acquire_visitor_run(
90        &self,
91        token_hash: &str,
92        limit: u32,
93    ) -> Result<OwnedSemaphorePermit, tokio::sync::TryAcquireError> {
94        visitor_semaphore(&self.visitor_run_permits, token_hash, limit)
95            .await
96            .try_acquire_owned()
97    }
98
99    pub async fn try_acquire_visitor_sse(
100        &self,
101        token_hash: &str,
102        limit: u32,
103    ) -> Result<OwnedSemaphorePermit, tokio::sync::TryAcquireError> {
104        visitor_semaphore(&self.visitor_sse_permits, token_hash, limit)
105            .await
106            .try_acquire_owned()
107    }
108
109    /// Atomically reserves one accepted-start slot in both the visitor-token and client-IP
110    /// windows. A caller that fails before scheduling the run must release the returned handle;
111    /// a successfully scheduled run leaves it recorded until the fixed window expires.
112    pub async fn reserve_accepted_start(
113        &self,
114        token_hash: &str,
115        client_ip: &str,
116        now: DateTime<Utc>,
117        policy: DemoPolicy,
118    ) -> Result<DemoAcceptedStart, DemoQuotaExceeded> {
119        let mut windows = self.accepted_starts.lock().await;
120        prune_start_windows(&mut windows, now, policy.accepted_start_window_secs);
121        if let Some(retry_after_secs) =
122            quota_retry_after(&windows, token_hash, client_ip, now, policy)
123        {
124            return Err(DemoQuotaExceeded { retry_after_secs });
125        }
126
127        windows.next_id = windows.next_id.wrapping_add(1);
128        let accepted = AcceptedStart {
129            id: windows.next_id,
130            accepted_at: now,
131        };
132        windows
133            .by_token
134            .entry(token_hash.to_owned())
135            .or_default()
136            .push_back(accepted.clone());
137        windows
138            .by_ip
139            .entry(client_ip.to_owned())
140            .or_default()
141            .push_back(accepted.clone());
142        Ok(DemoAcceptedStart {
143            id: accepted.id,
144            token_hash: token_hash.to_owned(),
145            client_ip: client_ip.to_owned(),
146        })
147    }
148
149    pub async fn release_accepted_start(&self, accepted: DemoAcceptedStart) {
150        let mut windows = self.accepted_starts.lock().await;
151        remove_accepted_start(&mut windows.by_token, &accepted.token_hash, accepted.id);
152        remove_accepted_start(&mut windows.by_ip, &accepted.client_ip, accepted.id);
153    }
154
155    /// Reclaims expired quota entries and inactive per-visitor semaphores during the periodic
156    /// Demo cleanup pass, including tokens that never make another request.
157    pub async fn cleanup(&self, now: DateTime<Utc>, policy: DemoPolicy) {
158        {
159            let mut accepted_starts = self.accepted_starts.lock().await;
160            prune_start_windows(&mut accepted_starts, now, policy.accepted_start_window_secs);
161        }
162        {
163            let mut visitor_runs = self.visitor_run_permits.lock().await;
164            retain_active_visitor_semaphores(&mut visitor_runs, policy.max_active_runs_per_visitor);
165        }
166        {
167            let mut visitor_streams = self.visitor_sse_permits.lock().await;
168            retain_active_visitor_semaphores(&mut visitor_streams, policy.max_sse_per_visitor);
169        }
170    }
171}
172
173async fn visitor_semaphore(
174    permits: &Mutex<HashMap<String, Arc<Semaphore>>>,
175    token_hash: &str,
176    limit: u32,
177) -> Arc<Semaphore> {
178    let mut permits = permits.lock().await;
179    permits.retain(|_, permit| {
180        Arc::strong_count(permit) > 1 || permit.available_permits() < limit as usize
181    });
182    permits
183        .entry(token_hash.to_owned())
184        .or_insert_with(|| Arc::new(Semaphore::new(limit as usize)))
185        .clone()
186}
187
188fn retain_active_visitor_semaphores(permits: &mut HashMap<String, Arc<Semaphore>>, limit: u32) {
189    permits.retain(|_, permit| {
190        Arc::strong_count(permit) > 1 || permit.available_permits() < limit as usize
191    });
192}
193
194fn prune_start_windows(windows: &mut AcceptedStartWindows, now: DateTime<Utc>, window_secs: u64) {
195    let cutoff = now - Duration::seconds(window_secs as i64);
196    windows.by_token.retain(|_, starts| {
197        while starts
198            .front()
199            .is_some_and(|start| start.accepted_at <= cutoff)
200        {
201            starts.pop_front();
202        }
203        !starts.is_empty()
204    });
205    windows.by_ip.retain(|_, starts| {
206        while starts
207            .front()
208            .is_some_and(|start| start.accepted_at <= cutoff)
209        {
210            starts.pop_front();
211        }
212        !starts.is_empty()
213    });
214}
215
216fn quota_retry_after(
217    windows: &AcceptedStartWindows,
218    token_hash: &str,
219    client_ip: &str,
220    now: DateTime<Utc>,
221    policy: DemoPolicy,
222) -> Option<u64> {
223    let token_retry = retry_after_for(
224        windows.by_token.get(token_hash),
225        policy.accepted_starts_per_token,
226        now,
227        policy.accepted_start_window_secs,
228    );
229    let ip_retry = retry_after_for(
230        windows.by_ip.get(client_ip),
231        policy.accepted_starts_per_client_ip,
232        now,
233        policy.accepted_start_window_secs,
234    );
235    match (token_retry, ip_retry) {
236        (Some(left), Some(right)) => Some(left.max(right)),
237        (Some(retry), None) | (None, Some(retry)) => Some(retry),
238        (None, None) => None,
239    }
240}
241
242fn retry_after_for(
243    starts: Option<&VecDeque<AcceptedStart>>,
244    limit: u32,
245    now: DateTime<Utc>,
246    window_secs: u64,
247) -> Option<u64> {
248    let starts = starts?;
249    if starts.len() < limit as usize {
250        return None;
251    }
252    let remaining_ms = (starts.front()?.accepted_at + Duration::seconds(window_secs as i64) - now)
253        .num_milliseconds()
254        .max(1);
255    Some((remaining_ms as u64).div_ceil(1_000))
256}
257
258fn remove_accepted_start(
259    windows: &mut HashMap<String, VecDeque<AcceptedStart>>,
260    key: &str,
261    id: u64,
262) {
263    let Some(starts) = windows.get_mut(key) else {
264        return;
265    };
266    starts.retain(|start| start.id != id);
267    if starts.is_empty() {
268        windows.remove(key);
269    }
270}
271
272/// Cheap to clone (an `Arc`-backed connection pool under the hood, per `sqlx`, and
273/// [`ActiveRun`] is itself `Arc<Mutex<..>>`-backed) -- axum's `State` extractor just requires
274/// `Clone` to hand a copy to each request.
275#[derive(Clone)]
276pub struct AppState {
277    pub pool: SqlitePool,
278    /// The in-flight tune registry and exclusive post-hoc write/revert reservation -- see
279    /// [`ActiveRun`]'s own doc comment for the concurrency and shutdown behavior.
280    pub active_run: ActiveRun,
281    /// The live, revisioned TOML configuration. Route handlers take a fresh snapshot for
282    /// every operation so a configuration-page save is visible without restarting the server.
283    pub config_store: Arc<RwLock<bhtune_cli::config::LoadedConfigStore>>,
284    pub allowed_origin: Option<String>,
285    pub trusted_proxy: Option<String>,
286    pub mode: ServerMode,
287    pub demo_policy: DemoPolicy,
288    pub demo_runtime: DemoRuntime,
289}
290
291impl AppState {
292    pub fn config_snapshot(&self) -> anyhow::Result<bhtune_cli::config::BhtuneConfig> {
293        self.config_store
294            .read()
295            .map(|store| store.config.clone())
296            .map_err(|_| anyhow::anyhow!("configuration store lock is poisoned"))
297    }
298
299    pub fn for_mode(
300        pool: SqlitePool,
301        config_store: Arc<RwLock<bhtune_cli::config::LoadedConfigStore>>,
302        mode: ServerMode,
303        demo_policy: DemoPolicy,
304    ) -> Self {
305        Self::for_mode_with_network_config(pool, config_store, mode, demo_policy, None, None)
306    }
307
308    pub fn for_mode_with_network_config(
309        pool: SqlitePool,
310        config_store: Arc<RwLock<bhtune_cli::config::LoadedConfigStore>>,
311        mode: ServerMode,
312        demo_policy: DemoPolicy,
313        allowed_origin: Option<String>,
314        trusted_proxy: Option<String>,
315    ) -> Self {
316        Self {
317            pool,
318            active_run: ActiveRun::default(),
319            config_store,
320            allowed_origin,
321            trusted_proxy,
322            demo_runtime: DemoRuntime::new(demo_policy),
323            mode,
324            demo_policy,
325        }
326    }
327}
328
329#[cfg(test)]
330mod tests {
331    use super::*;
332
333    #[tokio::test]
334    async fn config_snapshot_returns_an_independent_config_copy() {
335        let state = crate::test_support::in_memory_state().await;
336        let mut snapshot = state.config_snapshot().unwrap();
337        snapshot.allow_uncertain_quality = false;
338
339        assert!(state.config_snapshot().unwrap().allow_uncertain_quality);
340    }
341
342    #[tokio::test]
343    async fn config_snapshot_reports_a_poisoned_store_lock() {
344        let state = crate::test_support::in_memory_state().await;
345        let store = state.config_store.clone();
346        let _ = std::thread::spawn(move || {
347            let _guard = store.write().unwrap();
348            panic!("deliberately poison the test lock");
349        })
350        .join();
351
352        assert!(
353            state
354                .config_snapshot()
355                .unwrap_err()
356                .to_string()
357                .contains("poisoned")
358        );
359    }
360
361    #[tokio::test]
362    async fn demo_runtime_uses_the_global_and_ordinary_policy_limits() {
363        let policy = DemoPolicy {
364            max_active_runs_global: 1,
365            max_sse_global: 1,
366            ordinary_request_concurrency: 1,
367            ..DemoPolicy::default()
368        };
369        let runtime = DemoRuntime::new(policy);
370
371        let _run = runtime.try_acquire_global_run().unwrap();
372        assert!(runtime.try_acquire_global_run().is_err());
373        let _sse = runtime.try_acquire_global_sse().unwrap();
374        assert!(runtime.try_acquire_global_sse().is_err());
375        let _request = runtime.try_acquire_ordinary_request().unwrap();
376        assert!(runtime.try_acquire_ordinary_request().is_err());
377    }
378
379    #[tokio::test]
380    async fn demo_start_admission_is_serialized() {
381        let runtime = DemoRuntime::new(DemoPolicy::default());
382        let admission = runtime.lock_start_admission().await;
383        let waiting_runtime = runtime.clone();
384        let waiter = tokio::spawn(async move {
385            let _admission = waiting_runtime.lock_start_admission().await;
386        });
387        tokio::task::yield_now().await;
388        assert!(!waiter.is_finished());
389
390        drop(admission);
391
392        waiter.await.unwrap();
393    }
394
395    #[tokio::test]
396    async fn accepted_start_quotas_are_independent_for_token_and_client_ip() {
397        let policy = DemoPolicy::default();
398        let now = DateTime::from_timestamp(1_700_000_000, 0).unwrap();
399        let token_runtime = DemoRuntime::new(policy);
400        for index in 0..policy.accepted_starts_per_token {
401            token_runtime
402                .reserve_accepted_start("token", &format!("ip-{index}"), now, policy)
403                .await
404                .unwrap();
405        }
406        let token_error = token_runtime
407            .reserve_accepted_start("token", "unused-ip", now, policy)
408            .await
409            .unwrap_err();
410        assert_eq!(
411            token_error.retry_after_secs,
412            policy.accepted_start_window_secs
413        );
414        assert!(
415            token_runtime
416                .reserve_accepted_start(
417                    "token",
418                    "unused-ip",
419                    now + Duration::seconds(policy.accepted_start_window_secs as i64),
420                    policy,
421                )
422                .await
423                .is_ok()
424        );
425
426        let ip_runtime = DemoRuntime::new(policy);
427        for index in 0..policy.accepted_starts_per_client_ip {
428            ip_runtime
429                .reserve_accepted_start(&format!("token-{index}"), "shared-ip", now, policy)
430                .await
431                .unwrap();
432        }
433        assert!(
434            ip_runtime
435                .reserve_accepted_start("unused-token", "shared-ip", now, policy)
436                .await
437                .is_err()
438        );
439        assert!(
440            ip_runtime
441                .reserve_accepted_start("unused-token", "unused-ip", now, policy)
442                .await
443                .is_ok()
444        );
445    }
446
447    #[tokio::test]
448    async fn accepted_start_quota_uses_the_larger_token_or_ip_retry_after() {
449        let policy = DemoPolicy {
450            accepted_starts_per_token: 1,
451            accepted_starts_per_client_ip: 1,
452            accepted_start_window_secs: 20,
453            ..DemoPolicy::default()
454        };
455        let now = DateTime::from_timestamp(1_700_000_000, 0).unwrap();
456        let runtime = DemoRuntime::new(policy);
457
458        runtime
459            .reserve_accepted_start("token", "other-ip", now - Duration::seconds(5), policy)
460            .await
461            .unwrap();
462        runtime
463            .reserve_accepted_start(
464                "other-token",
465                "shared-ip",
466                now - Duration::seconds(10),
467                policy,
468            )
469            .await
470            .unwrap();
471
472        let error = runtime
473            .reserve_accepted_start("token", "shared-ip", now, policy)
474            .await
475            .unwrap_err();
476
477        assert_eq!(error.retry_after_secs, 15);
478    }
479
480    #[tokio::test]
481    async fn rejected_start_reservations_can_be_released_without_consuming_quota() {
482        let policy = DemoPolicy::default();
483        let now = DateTime::from_timestamp(1_700_000_000, 0).unwrap();
484        let runtime = DemoRuntime::new(policy);
485        let mut starts = Vec::new();
486        for index in 0..policy.accepted_starts_per_token {
487            starts.push(
488                runtime
489                    .reserve_accepted_start("token", &format!("ip-{index}"), now, policy)
490                    .await
491                    .unwrap(),
492            );
493        }
494        assert!(
495            runtime
496                .reserve_accepted_start("token", "another-ip", now, policy)
497                .await
498                .is_err()
499        );
500
501        runtime.release_accepted_start(starts.remove(0)).await;
502
503        assert!(
504            runtime
505                .reserve_accepted_start("token", "another-ip", now, policy)
506                .await
507                .is_ok()
508        );
509    }
510
511    #[tokio::test]
512    async fn releasing_an_accepted_start_removes_empty_token_and_ip_windows() {
513        let policy = DemoPolicy::default();
514        let now = DateTime::from_timestamp(1_700_000_000, 0).unwrap();
515        let runtime = DemoRuntime::new(policy);
516        let accepted = runtime
517            .reserve_accepted_start("token", "ip", now, policy)
518            .await
519            .unwrap();
520
521        runtime.release_accepted_start(accepted).await;
522
523        let windows = runtime.accepted_starts.lock().await;
524        assert!(windows.by_token.is_empty());
525        assert!(windows.by_ip.is_empty());
526    }
527
528    #[test]
529    fn removing_an_unknown_accepted_start_is_a_no_op() {
530        let mut windows = HashMap::new();
531
532        remove_accepted_start(&mut windows, "missing", 1);
533
534        assert!(windows.is_empty());
535    }
536
537    #[tokio::test]
538    async fn cleanup_reclaims_expired_windows_and_idle_visitor_semaphores() {
539        let policy = DemoPolicy {
540            accepted_starts_per_token: 1,
541            accepted_starts_per_client_ip: 1,
542            accepted_start_window_secs: 10,
543            max_active_runs_per_visitor: 1,
544            max_sse_per_visitor: 1,
545            ..DemoPolicy::default()
546        };
547        let runtime = DemoRuntime::new(policy);
548        let now = DateTime::from_timestamp(1_700_000_000, 0).unwrap();
549        runtime
550            .reserve_accepted_start("token", "ip", now, policy)
551            .await
552            .unwrap();
553        drop(runtime.try_acquire_visitor_run("token", 1).await.unwrap());
554        drop(runtime.try_acquire_visitor_sse("token", 1).await.unwrap());
555
556        runtime.cleanup(now + Duration::seconds(10), policy).await;
557
558        assert!(
559            runtime
560                .reserve_accepted_start("token", "ip", now + Duration::seconds(10), policy)
561                .await
562                .is_ok()
563        );
564        assert!(runtime.visitor_run_permits.lock().await.is_empty());
565        assert!(runtime.visitor_sse_permits.lock().await.is_empty());
566    }
567}