1mod connection;
21mod read;
22mod trace;
23mod values;
24mod write;
25
26use std::sync::atomic::{AtomicUsize, Ordering};
27use std::sync::{Arc, Mutex};
28use std::time::Instant;
29
30use helix_core::effect::{
31 BatchDeleteSpec, BatchUpdateSpec, GetSpec, GuardedBumpSpec, MonotonicUpsertSpec,
32 Row as HelixRow, ScanSpec, ScopedGetSpec, ScopedGuardedBumpSpec, UpsertSpec,
33};
34use helix_core::ports::Storage;
35use helix_core::PortError;
36use rusqlite::Connection;
37
38use crate::metrics::{
39 AsyncMetricSink, LabelKey, MetricEvent, MetricId, MetricLabels, NoopMetricSink,
40};
41
42use connection::{
43 map_join_err, map_lock_err, map_sqlite_err, open_reader, open_writer, sqlite_target_from_url,
44 target_supports_shared_readers, ReadPool, READ_POOL_SIZE,
45};
46use trace::{current_storage_trace_context, trace_sql};
47
48pub use trace::{
49 with_storage_operation_context, with_storage_trace_context, StorageOperationContext,
50 StorageTraceContext,
51};
52
53#[derive(Clone)]
56pub struct HostStorage {
57 writer: Arc<Mutex<Connection>>,
58 readers: Option<Arc<ReadPool>>,
59 metrics: Arc<dyn AsyncMetricSink>,
60 db_target: Arc<String>,
61 inflight: Arc<AtomicUsize>,
62}
63
64impl HostStorage {
65 pub async fn open_sqlite_url(db_url: &str) -> Result<Self, PortError> {
66 let target = sqlite_target_from_url(db_url);
67 let shareable = target_supports_shared_readers(&target);
68
69 let target_for_blocking = target.clone();
70 let (writer, readers) = tokio::task::spawn_blocking(move || {
71 let writer = open_writer(&target_for_blocking)?;
72
73 let readers = if shareable {
74 let mut conns = Vec::with_capacity(READ_POOL_SIZE);
75 for _ in 0..READ_POOL_SIZE {
76 conns.push(open_reader(&target_for_blocking)?);
77 }
78 Some(Arc::new(ReadPool::new(conns)))
79 } else {
80 None
81 };
82
83 Ok::<_, PortError>((writer, readers))
84 })
85 .await
86 .map_err(map_join_err)??;
87
88 let storage = Self {
89 writer: Arc::new(Mutex::new(writer)),
90 readers,
91 metrics: Arc::new(NoopMetricSink),
92 db_target: Arc::new(target),
93 inflight: Arc::new(AtomicUsize::new(0)),
94 };
95 storage.migrate().await?;
96 Ok(storage)
97 }
98
99 pub fn with_metric_sink(mut self, metrics: Arc<dyn AsyncMetricSink>) -> Self {
100 self.metrics = metrics;
101 self
102 }
103
104 pub async fn execute_raw(&self, sql: &'static str) -> Result<(), PortError> {
105 let writer = Arc::clone(&self.writer);
106 let trace = current_storage_trace_context();
107 tokio::task::spawn_blocking(move || {
108 let conn = writer.lock().map_err(map_lock_err)?;
109 let _span = trace_sql(&trace, "EXECUTE", None, sql, &[]);
110 conn.execute_batch(sql).map_err(map_sqlite_err)
111 })
112 .await
113 .map_err(map_join_err)?
114 }
115
116 async fn migrate(&self) -> Result<(), PortError> {
117 self.execute_raw(
118 "PRAGMA journal_mode=WAL;
119 PRAGMA synchronous=NORMAL;
120 PRAGMA busy_timeout=5000;
121 CREATE TABLE IF NOT EXISTS _helix_monotonic
122 (scope_key TEXT PRIMARY KEY NOT NULL, value INTEGER NOT NULL DEFAULT 0);",
123 )
124 .await
125 }
126
127 pub(super) async fn with_reader<T, F>(&self, f: F) -> Result<T, PortError>
129 where
130 T: Send + 'static,
131 F: FnOnce(&Connection) -> Result<T, PortError> + Send + 'static,
132 {
133 match &self.readers {
134 Some(pool) => {
135 let permit = Arc::clone(&pool.permits)
136 .acquire_owned()
137 .await
138 .map_err(|_| PortError::Storage("read pool semaphore closed".to_string()))?;
139 let pool = Arc::clone(pool);
140 tokio::task::spawn_blocking(move || {
141 let conn = {
142 let mut idle = pool.idle.lock().map_err(map_lock_err)?;
143 idle.pop()
144 };
145 let Some(conn) = conn else {
146 return Err(PortError::Storage(
147 "read pool permit/conn mismatch".to_string(),
148 ));
149 };
150 let out = f(&conn);
151 if let Ok(mut idle) = pool.idle.lock() {
152 idle.push(conn);
153 }
154 drop(permit);
155 out
156 })
157 .await
158 .map_err(map_join_err)?
159 }
160 None => {
161 let writer = Arc::clone(&self.writer);
162 tokio::task::spawn_blocking(move || {
163 let conn = writer.lock().map_err(map_lock_err)?;
164 f(&conn)
165 })
166 .await
167 .map_err(map_join_err)?
168 }
169 }
170 }
171}
172
173#[cfg(test)]
174mod tests {
175 use std::path::{Path, PathBuf};
176 use std::sync::atomic::{AtomicU64, Ordering};
177
178 use super::HostStorage;
179
180 fn temp_db_path() -> PathBuf {
182 static SEQUENCE: AtomicU64 = AtomicU64::new(0);
183 let sequence = SEQUENCE.fetch_add(1, Ordering::Relaxed);
184 let path = std::env::temp_dir().join(format!(
185 "helix-host-sqlite-uri-{}-{sequence}.db",
186 std::process::id()
187 ));
188 remove_temp_db(&path);
189 path
190 }
191
192 fn remove_temp_db(path: &Path) {
194 let _ = std::fs::remove_file(path);
195 let _ = std::fs::remove_file(path.with_extension("db-wal"));
196 let _ = std::fs::remove_file(path.with_extension("db-shm"));
197 }
198
199 #[tokio::test]
200 async fn sqlite_mode_rwc_writer_and_reader_share_the_public_target() {
201 let path = temp_db_path();
202 let url = format!("sqlite:{}?mode=rwc", path.display());
203 let storage = HostStorage::open_sqlite_url(&url)
204 .await
205 .expect("HostStorage should open the canonical target");
206 storage
207 .execute_raw(
208 "CREATE TABLE uri_regression (value INTEGER NOT NULL);\
209 INSERT INTO uri_regression VALUES (73);",
210 )
211 .await
212 .expect("writer should persist through the public storage entry point");
213
214 let value = storage
215 .with_reader(|conn| {
216 conn.query_row("SELECT value FROM uri_regression", [], |row| {
217 row.get::<_, i64>(0)
218 })
219 .map_err(super::map_sqlite_err)
220 })
221 .await
222 .expect("reader pool should read the writer's file");
223 assert_eq!(value, 73);
224
225 drop(storage);
226 remove_temp_db(&path);
227 }
228}
229
230#[async_trait::async_trait]
231impl Storage for HostStorage {
232 async fn batch_upsert(&self, spec: UpsertSpec) -> Result<(), PortError> {
233 if !self.metrics.is_enabled() {
234 return write::batch_upsert(self, spec).await;
235 }
236 let rows = spec.rows.len();
237 let started = self.start_storage_operation();
238 let result = write::batch_upsert(self, spec).await;
239 self.record_storage_result("batch_upsert", rows, started, &result);
240 result
241 }
242
243 async fn batch_update(&self, spec: BatchUpdateSpec) -> Result<(), PortError> {
244 if !self.metrics.is_enabled() {
245 return write::batch_update(self, spec).await;
246 }
247 let rows = spec.key_vals.len();
248 let started = self.start_storage_operation();
249 let result = write::batch_update(self, spec).await;
250 self.record_storage_result("batch_update", rows, started, &result);
251 result
252 }
253
254 async fn monotonic_upsert(&self, spec: MonotonicUpsertSpec) -> Result<(), PortError> {
255 if !self.metrics.is_enabled() {
256 return write::monotonic_upsert(self, spec).await;
257 }
258 let started = self.start_storage_operation();
259 let result = write::monotonic_upsert(self, spec).await;
260 self.record_storage_result("monotonic_upsert", 1, started, &result);
261 result
262 }
263
264 async fn guarded_bump(&self, spec: GuardedBumpSpec) -> Result<(), PortError> {
265 if !self.metrics.is_enabled() {
266 return write::guarded_bump(self, spec).await;
267 }
268 let started = self.start_storage_operation();
269 let result = write::guarded_bump(self, spec).await;
270 self.record_storage_result("guarded_bump", 1, started, &result);
271 result
272 }
273
274 async fn scoped_guarded_bump(&self, spec: ScopedGuardedBumpSpec) -> Result<(), PortError> {
276 if !self.metrics.is_enabled() {
277 return write::scoped_guarded_bump(self, spec).await;
278 }
279 let started = self.start_storage_operation();
280 let result = write::scoped_guarded_bump(self, spec).await;
281 self.record_storage_result("scoped_guarded_bump", 1, started, &result);
282 result
283 }
284
285 async fn get(&self, spec: GetSpec) -> Result<Option<HelixRow>, PortError> {
286 if !self.metrics.is_enabled() {
287 return read::get(self, spec).await;
288 }
289 let started = self.start_storage_operation();
290 let result = read::get(self, spec).await;
291 let rows = result
292 .as_ref()
293 .ok()
294 .and_then(Option::as_ref)
295 .map_or(0, |_| 1);
296 self.record_storage_result("get", rows, started, &result);
297 result
298 }
299
300 async fn scoped_get(&self, spec: ScopedGetSpec) -> Result<Option<HelixRow>, PortError> {
302 if !self.metrics.is_enabled() {
303 return read::scoped_get(self, spec).await;
304 }
305 let started = self.start_storage_operation();
306 let result = read::scoped_get(self, spec).await;
307 let rows = result
308 .as_ref()
309 .ok()
310 .and_then(Option::as_ref)
311 .map_or(0, |_| 1);
312 self.record_storage_result("scoped_get", rows, started, &result);
313 result
314 }
315
316 async fn scan(&self, spec: ScanSpec) -> Result<Vec<HelixRow>, PortError> {
317 if !self.metrics.is_enabled() {
318 return read::scan(self, spec).await;
319 }
320 let started = self.start_storage_operation();
321 let result = read::scan(self, spec).await;
322 let rows = result.as_ref().map_or(0, Vec::len);
323 self.record_storage_result("scan", rows, started, &result);
324 result
325 }
326
327 async fn batch_delete(&self, spec: BatchDeleteSpec) -> Result<(), PortError> {
328 if !self.metrics.is_enabled() {
329 return write::batch_delete(self, spec).await;
330 }
331 let rows = spec.key_vals.len();
332 let started = self.start_storage_operation();
333 let result = write::batch_delete(self, spec).await;
334 self.record_storage_result("batch_delete", rows, started, &result);
335 result
336 }
337
338 async fn atomic_write(&self, ops: Vec<helix_core::effect::StorageOp>) -> Result<(), PortError> {
339 if !self.metrics.is_enabled() {
340 return write::atomic_write(self, ops).await;
341 }
342 let operation_count = ops.len();
343 let started = self.start_storage_operation();
344 let result = write::atomic_write(self, ops).await;
345 self.record_storage_result("atomic_write", operation_count, started, &result);
346 result
347 }
348}
349
350impl HostStorage {
351 fn start_storage_operation(&self) -> Instant {
353 let inflight = self.inflight.fetch_add(1, Ordering::Relaxed) + 1;
354 let _ = self.metrics.try_record(MetricEvent::gauge(
355 MetricId::StorageInflight,
356 inflight as f64,
357 MetricLabels::one(LabelKey::Stage, "storage"),
358 ));
359 Instant::now()
360 }
361
362 fn record_storage_result<T>(
364 &self,
365 operation: &'static str,
366 rows: usize,
367 started: Instant,
368 result: &Result<T, PortError>,
369 ) {
370 let remaining = self
371 .inflight
372 .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |current| {
373 Some(current.saturating_sub(1))
374 })
375 .unwrap_or_default()
376 .saturating_sub(1);
377 let _ = self.metrics.try_record(MetricEvent::gauge(
378 MetricId::StorageInflight,
379 remaining as f64,
380 MetricLabels::one(LabelKey::Stage, "storage"),
381 ));
382 let status = if result.is_ok() { "ok" } else { "error" };
383 let labels = MetricLabels::one(LabelKey::Stage, "storage")
384 .with(LabelKey::StorageOp, operation)
385 .with(LabelKey::Status, status);
386 let _ = self.metrics.try_record(MetricEvent::histogram(
387 MetricId::StorageTxDurationSeconds,
388 started.elapsed().as_secs_f64(),
389 labels,
390 ));
391 let _ = self.metrics.try_record(MetricEvent::counter(
392 MetricId::OperationsTotal,
393 1.0,
394 labels.with(LabelKey::Operation, operation),
395 ));
396 let _ = self.metrics.try_record(MetricEvent::histogram(
397 MetricId::StorageBatchRows,
398 rows as f64,
399 labels,
400 ));
401 if result.is_ok() {
402 let _ = self.metrics.try_record(MetricEvent::counter(
403 MetricId::StorageRowsTotal,
404 rows as f64,
405 labels,
406 ));
407 } else {
408 let error_text = result
409 .as_ref()
410 .err()
411 .map(ToString::to_string)
412 .unwrap_or_default()
413 .to_ascii_lowercase();
414 if error_text.contains("busy") || error_text.contains("locked") {
415 let _ = self.metrics.try_record(MetricEvent::counter(
416 MetricId::StorageBusyTotal,
417 1.0,
418 labels,
419 ));
420 }
421 if error_text.contains("timeout") || error_text.contains("timed out") {
422 let _ = self.metrics.try_record(MetricEvent::counter(
423 MetricId::StorageTimeoutTotal,
424 1.0,
425 labels,
426 ));
427 }
428 if operation == "atomic_write" {
429 let _ = self.metrics.try_record(MetricEvent::counter(
430 MetricId::StorageRollbackTotal,
431 1.0,
432 labels,
433 ));
434 }
435 let _ = self.metrics.try_record(MetricEvent::counter(
436 MetricId::ErrorsTotal,
437 1.0,
438 labels.with(LabelKey::ErrorKind, "storage_failed"),
439 ));
440 }
441 }
442}