summaryrefslogtreecommitdiff
path: root/server/src/state.rs
diff options
context:
space:
mode:
Diffstat (limited to 'server/src/state.rs')
-rw-r--r--server/src/state.rs55
1 files changed, 55 insertions, 0 deletions
diff --git a/server/src/state.rs b/server/src/state.rs
new file mode 100644
index 0000000..512bcb3
--- /dev/null
+++ b/server/src/state.rs
@@ -0,0 +1,55 @@
+use std::sync::Arc;
+use tokio::sync::{Mutex as AsyncMutex, Semaphore};
+
+use crate::config::Config;
+
+/// Shared application state. The single SQLite connection is serialized behind
+/// an async mutex; concurrent fetch/backtest work is bounded by semaphores.
+pub struct AppState {
+ pub cfg: Config,
+ pub db: AsyncMutex<rusqlite::Connection>,
+ pub run_sem: Arc<Semaphore>,
+ pub fetch_sem: Arc<Semaphore>,
+}
+
+/// Fixed shared contract alias for handler arguments: the axum State extractor
+/// over `Arc<AppState>`. Handlers take `cx: Cx`, helpers take `&Cx`.
+pub type Cx = axum::extract::State<std::sync::Arc<AppState>>;
+
+/// Legacy alias kept for modules (auth, projects, datasets) whose handlers take
+/// the bare `Arc<AppState>`; both styles are valid extractors for this state.
+pub type S = Arc<AppState>;
+
+impl AppState {
+ /// Lock the connection, run the closure, return its result exactly.
+ /// Never `.await` inside the closure; keep transactions in ONE closure.
+ pub async fn with_db<R>(&self, f: impl FnOnce(&mut rusqlite::Connection) -> R) -> R {
+ let mut db = self.db.lock().await;
+ f(&mut db)
+ }
+}
+
+#[cfg(test)]
+mod tests {
+ use super::*;
+ use crate::config::Config;
+
+ #[tokio::test]
+ async fn with_db_returns_closure_result_and_composes() {
+ let cfg = Config::from_env();
+ let conn = rusqlite::Connection::open_in_memory().unwrap();
+ let st = AppState {
+ cfg: cfg.clone(),
+ db: AsyncMutex::new(conn),
+ run_sem: Arc::new(Semaphore::new(1)),
+ fetch_sem: Arc::new(Semaphore::new(1)),
+ };
+ let n: i64 = st.with_db(|db| {
+ db.execute("CREATE TABLE t(x INTEGER)", []).unwrap();
+ db.execute("INSERT INTO t VALUES (42)", []).unwrap();
+ db.query_row("SELECT SUM(x) FROM t", [], |r| r.get(0)).unwrap()
+ }).await;
+ assert_eq!(n, 42);
+ st.with_db(|_db| ()).await;
+ }
+}