1use 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 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 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 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#[derive(Clone)]
276pub struct AppState {
277 pub pool: SqlitePool,
278 pub active_run: ActiveRun,
281 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}