Skip to main content

helix_driver_host/
storage.rs

1//! Shared rusqlite storage implementation for PC Tauri and Flutter FFI.
2//!
3//! ## 读写分离(B1 回归修复)
4//!
5//! 旧实现(commit a0595607 重构后)把 native 三池塌成单个 `Arc<Mutex<Connection>>`,
6//! 导致 get/scan 与任何 op 串行在同一把锁。本实现恢复读写分离:
7//!
8//! - 1 个写连接:batch_upsert / batch_update / monotonic_upsert 串行兑现。
9//! - N=4 个只读连接:get / scan 在 WAL 下与写者及其他读者并发。
10//! - 不支持共享连接的内存数据库回退到写连接读取,语义优先。
11//!
12//! ## 不变量(不可破)
13//!
14//! - batch_upsert 单事务(E2 写放大 1x)
15//! - monotonic MAX guard(E5)与 GuardedBump 单条 UPDATE(HX-C005)
16//! - `_helix_monotonic` 表在 open()/migrate() 建(不在热路径)
17//! - PRAGMA journal_mode=WAL + synchronous=NORMAL + busy_timeout=5000
18//! - SQL trace 的 statement、parameters、table 与错误映射保持一致
19
20mod 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/// 内存数据库若未启用 shared cache,第二个连接会得到独立空库;文件 WAL 库与显式
54/// `cache=shared` 的内存库才启用只读池。
55#[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    /// 在只读连接上跑闭包;无可共享只读池时回退到写连接。
128    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    /// 为公开 HostStorage 入口生成隔离的真实文件库路径。
181    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    /// 删除主库和 WAL/SHM sidecar,避免回归测试污染下一次运行。
193    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    /// 在复合主键作用域内兑现单条守卫式计数更新。
275    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    /// 在复合主键作用域内读取唯一行。
301    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    /// 进入真实 SQLite 操作区间并发布并发 Gauge。
352    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    /// 闭合存储结果、批大小和故障分类,并将 inflight 饱和归还。
363    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}