Skip to main content

ocre/runtime/
d1.rs

1use std::sync::Arc;
2
3use serde::de::DeserializeOwned;
4use wasm_bindgen::{JsCast, JsValue};
5use worker::{
6    D1Database, D1DatabaseSession, D1PreparedStatement,
7    js_sys::{Array, Reflect},
8    send::{SendFuture, SendWrapper},
9    wasm_bindgen_futures::JsFuture,
10};
11
12use super::ctx::Memo;
13use crate::{
14    Error, Param, Result, Statement,
15    cache::{is_read_query, query_key},
16    sql::Value,
17};
18use crate::{instrument::Timings, log::Logger};
19
20/// Handle to the application's D1 (SQLite) database, from [`Ctx::db`](crate::Ctx::db).
21///
22/// Every method returns a `Send` future, so plain axum handlers can await it.
23/// Always pass values through [`params!`](crate::params) and `?1, ?2`
24/// placeholders, never by formatting them into the SQL string. Rows are
25/// deserialized with serde, by column name; SQLite booleans and JSON columns
26/// need [`bool_from_sql`](crate::bool_from_sql) and
27/// [`json_from_sql`](crate::json_from_sql).
28///
29/// A failed query is [`Error::Internal`] (500) with the D1 message and the SQL;
30/// the details go to the Worker logs (Workers Logs), never to the client.
31///
32/// # Free plan
33///
34/// D1 bills rows read (every row a query scans, not only those returned) and
35/// rows written. Add indexes for `WHERE` and `ORDER BY` columns, and bound
36/// lists with `LIMIT` (see [`Page`](crate::Page)).
37///
38/// # Query cache
39///
40/// Like Rails' query cache, a `SELECT` run twice with the same parameters
41/// during one request (or one job, or one cron run) is answered from memory
42/// the second time: no D1 round trip and no rows read. Any other statement
43/// ([`execute`](Self::execute), [`batch`](Self::batch), `INSERT ...
44/// RETURNING` through [`first`](Self::first)) empties the cache, before and
45/// after it runs, so a request always reads its own writes. The cache holds
46/// up to 100 results and is emptied when full. Writes by other requests are
47/// not seen until the next request; use [`uncached`](Self::uncached) for a
48/// query that must hit D1 (polling in a loop, a row another Worker updates).
49///
50/// # Examples
51///
52/// ```no_run
53/// use axum::extract::{Path, State};
54/// use ocre::{Ctx, OptionExt, Result, params};
55/// use serde::Deserialize;
56///
57/// #[derive(Deserialize)]
58/// struct Post {
59///     id: i64,
60///     title: String,
61/// }
62///
63/// async fn show(State(ctx): State<Ctx>, Path(id): Path<i64>) -> Result<String> {
64///     let db = ctx.db()?;
65///     let post: Post = db.first("SELECT id, title FROM posts WHERE id = ?1", params![id]).await?.or_404()?;
66///     Ok(format!("#{} {}", post.id, post.title))
67/// }
68/// # let _ = show;
69/// ```
70pub struct Db {
71    inner: SendWrapper<Handle>,
72    binding: String,
73    memo: Option<Arc<Memo>>,
74    /// Whether `SELECT`s may be served from `memo` (false after [`Db::uncached`]).
75    cached: bool,
76    probe: Option<Probe>,
77}
78
79/// The database itself, or the request's session on it (read replicas).
80#[derive(Clone)]
81pub(crate) enum Handle {
82    Database(Arc<D1Database>),
83    Session(Arc<D1DatabaseSession>),
84}
85
86impl Handle {
87    fn prepare(&self, sql: &str) -> D1PreparedStatement {
88        match self {
89            Self::Database(db) => db.prepare(sql),
90            Self::Session(session) => session.prepare(sql),
91        }
92    }
93
94    async fn batch(&self, statements: Vec<D1PreparedStatement>) -> worker::Result<Vec<worker::D1Result>> {
95        match self {
96            Self::Database(db) => db.batch(statements).await,
97            Self::Session(session) => session.batch(statements).await,
98        }
99    }
100}
101
102/// Where a handle reports its statements: the request's logger (a `debug`
103/// line per statement) and timings (`Server-Timing`, the development error page).
104#[derive(Clone)]
105struct Probe {
106    log: Logger,
107    timings: Timings,
108}
109
110impl Db {
111    pub(crate) fn new(db: Handle, binding: &str, memo: Option<Arc<Memo>>) -> Self {
112        Self { inner: SendWrapper::new(db), binding: binding.to_owned(), memo, cached: true, probe: None }
113    }
114
115    /// This handle, timing and logging each statement for the request.
116    pub(crate) fn probed(self, log: Logger, timings: Timings) -> Self {
117        Self { probe: Some(Probe { log, timings }), ..self }
118    }
119
120    /// This handle without the per-request query cache: every query goes to D1.
121    ///
122    /// Writes through it still empty the cache of the other handles of the
123    /// request. Rails' `uncached` block.
124    ///
125    /// # Examples
126    ///
127    /// ```no_run
128    /// use axum::extract::State;
129    /// use ocre::{Ctx, Result, params};
130    /// use serde::Deserialize;
131    ///
132    /// #[derive(Deserialize)]
133    /// struct Import {
134    ///     status: String,
135    /// }
136    ///
137    /// async fn status(State(ctx): State<Ctx>) -> Result<Option<String>> {
138    ///     // Another Worker may have changed it since this request's last read.
139    ///     let db = ctx.db()?.uncached();
140    ///     let row: Option<Import> = db.first("SELECT status FROM imports WHERE id = ?1", params![1]).await?;
141    ///     Ok(row.map(|import| import.status))
142    /// }
143    /// # let _ = status;
144    /// ```
145    pub fn uncached(self) -> Self {
146        Self { cached: false, ..self }
147    }
148
149    /// Where a query's rows are remembered: `Some` for a `SELECT` when the
150    /// cache is on. A statement that may write empties the cache instead.
151    fn cache_slot(&self, sql: &str, params: &[Param]) -> Option<(Arc<Memo>, String)> {
152        let memo = self.memo.as_ref()?;
153        if !is_read_query(sql) {
154            memo.clear_queries();
155            memo.wrote(&self.binding);
156            return None;
157        }
158        self.cached.then(|| (Arc::clone(memo), query_key(&self.binding, sql, params)))
159    }
160
161    /// The memo to empty once a statement that may write has run.
162    fn write_memo(&self, sql: &str) -> Option<Arc<Memo>> {
163        self.memo.as_ref().filter(|_| !is_read_query(sql)).map(Arc::clone)
164    }
165
166    /// Runs a query and returns every row, deserialized into `T`.
167    ///
168    /// All rows are buffered in memory; bound the query with `LIMIT`.
169    ///
170    /// # Errors
171    ///
172    /// [`Error::Internal`] (500, logged with the SQL) when the statement does
173    /// not prepare or bind, D1 rejects it, or a row does not deserialize into `T`.
174    ///
175    /// # Free plan
176    ///
177    /// Counts every row the query scans as a D1 row read.
178    ///
179    /// # Examples
180    ///
181    /// ```no_run
182    /// use axum::extract::State;
183    /// use ocre::{Ctx, Page, Result, params};
184    /// use serde::Deserialize;
185    ///
186    /// #[derive(Deserialize)]
187    /// struct Post {
188    ///     title: String,
189    /// }
190    ///
191    /// async fn index(State(ctx): State<Ctx>, page: Page) -> Result<String> {
192    ///     let sql = "SELECT title FROM posts ORDER BY id DESC LIMIT ?1 OFFSET ?2";
193    ///     let posts: Vec<Post> = ctx.db()?.all(sql, params![page.limit, page.offset]).await?;
194    ///     Ok(posts.into_iter().map(|p| p.title).collect::<Vec<_>>().join("\n"))
195    /// }
196    /// # let _ = index;
197    /// ```
198    pub fn all<'q, T: DeserializeOwned>(
199        &self,
200        sql: &'q str,
201        params: Vec<Param>,
202    ) -> impl Future<Output = Result<Vec<T>>> + Send + use<'q, T> {
203        let slot = self.cache_slot(sql, &params);
204        let written = self.write_memo(sql);
205        let stmt = self.prepare(sql, params);
206        SendFuture::new(timed(self.probe.clone(), sql, async move {
207            let rows = rows(stmt?, sql, slot).await?;
208            if let Some(memo) = written {
209                memo.clear_queries();
210            }
211            rows.iter().map(|row| deserialize(&row, sql)).collect()
212        }))
213    }
214
215    /// Runs a query and returns its first row, if any, deserialized into `T`.
216    ///
217    /// Use with `INSERT ... RETURNING *` to get the created row back, and with
218    /// [`OptionExt::or_404`](crate::OptionExt::or_404) to turn a missing record
219    /// into a 404.
220    ///
221    /// # Errors
222    ///
223    /// [`Error::Internal`] (500, logged with the SQL) when the statement does
224    /// not prepare or bind, D1 rejects it (e.g. a `UNIQUE` or foreign-key
225    /// constraint), or the row does not deserialize into `T`.
226    ///
227    /// # Free plan
228    ///
229    /// D1 counts the rows the query scans, not only the one returned: add
230    /// `LIMIT 1` or look rows up by an indexed column. `INSERT ... RETURNING`
231    /// also counts the rows written.
232    ///
233    /// # Examples
234    ///
235    /// ```no_run
236    /// use axum::extract::State;
237    /// use ocre::{ApiResult, Created, Ctx, Error, Json, params};
238    /// use serde::{Deserialize, Serialize};
239    ///
240    /// #[derive(Deserialize, Serialize)]
241    /// struct Post {
242    ///     id: i64,
243    ///     title: String,
244    /// }
245    ///
246    /// async fn create(State(ctx): State<Ctx>, Json(title): Json<String>) -> ApiResult<Created<Post>> {
247    ///     let sql = "INSERT INTO posts (title) VALUES (?1) RETURNING *";
248    ///     let post: Option<Post> = ctx.db()?.first(sql, params![title]).await?;
249    ///     Ok(Created(post.ok_or_else(|| Error::internal("INSERT returned no row"))?))
250    /// }
251    /// # let _ = create;
252    /// ```
253    pub fn first<'q, T: DeserializeOwned>(
254        &self,
255        sql: &'q str,
256        params: Vec<Param>,
257    ) -> impl Future<Output = Result<Option<T>>> + Send + use<'q, T> {
258        let slot = self.cache_slot(sql, &params);
259        let written = self.write_memo(sql);
260        let stmt = self.prepare(sql, params);
261        SendFuture::new(timed(self.probe.clone(), sql, async move {
262            let stmt = stmt?;
263            let row = match slot {
264                Some(slot) => rows(stmt, sql, Some(slot)).await?.iter().next().map(|row| deserialize(&row, sql)),
265                None => {
266                    let row = stmt.first::<T>(None).await.map_err(|err| query_error(sql, err))?;
267                    if let Some(memo) = written {
268                        memo.clear_queries();
269                    }
270                    return Ok(row);
271                }
272            };
273            row.transpose()
274        }))
275    }
276
277    /// Runs a statement that returns no rows (`INSERT`, `UPDATE`, `DELETE`) and gives the number of rows changed.
278    ///
279    /// # Errors
280    ///
281    /// [`Error::Internal`] (500, logged with the SQL) when the statement does
282    /// not prepare or bind, or D1 rejects it (constraint violation, syntax
283    /// error, missing table: run `ocre migrate`).
284    ///
285    /// # Free plan
286    ///
287    /// Counts the rows written, plus the rows scanned to find them (rows read).
288    ///
289    /// # Examples
290    ///
291    /// ```no_run
292    /// use axum::extract::{Path, State};
293    /// use ocre::{Ctx, Error, Result, params};
294    ///
295    /// async fn destroy(State(ctx): State<Ctx>, Path(id): Path<i64>) -> Result<()> {
296    ///     match ctx.db()?.execute("DELETE FROM posts WHERE id = ?1", params![id]).await? {
297    ///         0 => Err(Error::NotFound),
298    ///         _ => Ok(()),
299    ///     }
300    /// }
301    /// # let _ = destroy;
302    /// ```
303    pub fn execute<'q>(
304        &self,
305        sql: &'q str,
306        params: Vec<Param>,
307    ) -> impl Future<Output = Result<usize>> + Send + use<'q> {
308        self.cache_slot(sql, &params);
309        let written = self.write_memo(sql);
310        let stmt = self.prepare(sql, params);
311        SendFuture::new(timed(self.probe.clone(), sql, async move {
312            let result = stmt?.run().await.map_err(|err| query_error(sql, err))?;
313            if let Some(memo) = written {
314                memo.clear_queries();
315            }
316            let meta = result.meta().map_err(|err| query_error(sql, err))?;
317            Ok(meta.and_then(|m| m.changes).unwrap_or(0))
318        }))
319    }
320
321    /// Whether a query returns at least one row.
322    ///
323    /// Write it as `SELECT 1 FROM ... WHERE ... LIMIT 1`, so D1 stops at the
324    /// first match; the generated models use it for uniqueness checks ("has
325    /// already been taken").
326    ///
327    /// # Errors
328    ///
329    /// [`Error::Internal`] (500, logged with the SQL) when the statement does
330    /// not prepare or bind, or D1 rejects it.
331    ///
332    /// # Free plan
333    ///
334    /// Counts the rows scanned as rows read: with an index on the `WHERE`
335    /// column and `LIMIT 1`, one row.
336    ///
337    /// # Examples
338    ///
339    /// ```no_run
340    /// use axum::extract::State;
341    /// use ocre::{Ctx, Result, Validator, params};
342    ///
343    /// async fn check_email(State(ctx): State<Ctx>, email: String) -> Result<()> {
344    ///     let taken = ctx.db()?.exists("SELECT 1 FROM users WHERE email = ?1 LIMIT 1", params![&email]).await?;
345    ///     Validator::new().check("email", taken, "has already been taken").finish()
346    /// }
347    /// # let _ = check_email;
348    /// ```
349    pub fn exists<'q>(&self, sql: &'q str, params: Vec<Param>) -> impl Future<Output = Result<bool>> + Send + use<'q> {
350        let slot = self.cache_slot(sql, &params);
351        let stmt = self.prepare(sql, params);
352        SendFuture::new(timed(self.probe.clone(), sql, async move {
353            match slot {
354                Some(slot) => Ok(rows(stmt?, sql, Some(slot)).await?.length() > 0),
355                None => {
356                    let row = stmt?.first::<serde_json::Value>(None).await.map_err(|err| query_error(sql, err))?;
357                    Ok(row.is_some())
358                }
359            }
360        }))
361    }
362
363    /// Runs every statement in one transaction and returns the rows changed by each one.
364    ///
365    /// All statements succeed or none is applied (D1 `batch`). Results are in
366    /// the order of `statements`. It is also one round trip to D1 instead of
367    /// one per statement.
368    ///
369    /// # Errors
370    ///
371    /// [`Error::Internal`] (500, logged with every statement's SQL joined by
372    /// `; `) when a statement does not prepare or bind, or when D1 rejects
373    /// any of them; nothing is written in that case.
374    ///
375    /// # Free plan
376    ///
377    /// Counts the rows read and written by every statement, as if each ran alone.
378    ///
379    /// # Examples
380    ///
381    /// ```no_run
382    /// use axum::extract::State;
383    /// use ocre::{Ctx, Result, Statement, params};
384    ///
385    /// async fn transfer(State(ctx): State<Ctx>) -> Result<()> {
386    ///     let changed = ctx
387    ///         .db()?
388    ///         .batch(vec![
389    ///             Statement::new("UPDATE accounts SET balance = balance - ?1 WHERE id = ?2", params![10, 1]),
390    ///             Statement::new("UPDATE accounts SET balance = balance + ?1 WHERE id = ?2", params![10, 2]),
391    ///         ])
392    ///         .await?;
393    ///     assert_eq!(changed.len(), 2);
394    ///     Ok(())
395    /// }
396    /// # let _ = transfer;
397    /// ```
398    pub fn batch(&self, statements: Vec<Statement>) -> impl Future<Output = Result<Vec<usize>>> + Send + '_ {
399        let sql = statements.iter().map(|s| s.sql.as_str()).collect::<Vec<_>>().join("; ");
400        let prepared: Result<Vec<D1PreparedStatement>> =
401            statements.into_iter().map(|s| self.prepare(&s.sql, s.params)).collect();
402        let db = &self.inner;
403        let memo = self.memo.clone();
404        if let Some(memo) = &memo {
405            memo.clear_queries();
406            memo.wrote(&self.binding);
407        }
408        let probe = self.probe.clone();
409        SendFuture::new(async move {
410            let batch = async { db.batch(prepared?).await.map_err(|err| query_error(&sql, err)) };
411            let results = timed(probe, &sql, batch).await;
412            if let Some(memo) = memo {
413                memo.clear_queries();
414            }
415            let results = results?;
416            results
417                .iter()
418                .map(|result| {
419                    Ok(result.meta().map_err(|err| query_error(&sql, err))?.and_then(|m| m.changes).unwrap_or(0))
420                })
421                .collect()
422        })
423    }
424
425    fn prepare(&self, sql: &str, params: Vec<Param>) -> Result<D1PreparedStatement> {
426        let values: Vec<JsValue> = params.into_iter().map(to_js).collect();
427        self.inner.prepare(sql).bind(&values).map_err(|err| query_error(sql, err))
428    }
429}
430
431/// Awaits a statement, then records how long it took (Workers' clock
432/// advances during I/O, so this is D1's time) and logs it at `debug`:
433/// `SQL (1.2 ms) SELECT ...`, like Rails' `Post Load (0.3ms) SELECT ...`.
434async fn timed<T>(probe: Option<Probe>, sql: &str, future: impl Future<Output = Result<T>>) -> Result<T> {
435    let Some(probe) = probe else { return future.await };
436    let started = crate::clock::now_millis();
437    let result = future.await;
438    let ms = crate::clock::now_millis() - started;
439    probe.timings.record(sql, ms);
440    if probe.log.enabled(crate::log::Level::Debug) {
441        let log = probe.log.with("duration_ms", ms);
442        let log = if result.is_err() { log.with("failed", true) } else { log };
443        log.debug(format_args!("SQL ({ms} ms) {sql}"));
444    }
445    result
446}
447
448/// The rows of a query, from the request's cache when `slot` holds them.
449/// A statement without a slot (uncached, or one that may write) runs as is.
450async fn rows(stmt: D1PreparedStatement, sql: &str, slot: Option<(Arc<Memo>, String)>) -> Result<Array> {
451    if let Some((memo, key)) = &slot
452        && let Some(rows) = memo.query(key)
453    {
454        return Ok(rows);
455    }
456    let promise = stmt.inner().all().map_err(|err| query_error(sql, js_error(err)))?;
457    let result = JsFuture::from(promise).await.map_err(|err| query_error(sql, js_error(err)))?;
458    let results =
459        Reflect::get(&result, &JsValue::from_str("results")).map_err(|err| query_error(sql, js_error(err)))?;
460    let rows = results.dyn_into::<Array>().unwrap_or_else(|_| Array::new());
461    if let Some((memo, key)) = slot {
462        memo.remember_query(key, rows.clone());
463    }
464    Ok(rows)
465}
466
467/// A rejected D1 promise as an error naming D1's message (`D1_ERROR: no such table: posts...`).
468fn js_error(err: JsValue) -> worker::Error {
469    match err.dyn_ref::<worker::js_sys::Error>() {
470        Some(err) => worker::Error::RustError(String::from(err.message())),
471        None => err.into(),
472    }
473}
474
475fn deserialize<T: DeserializeOwned>(row: &JsValue, sql: &str) -> Result<T> {
476    serde_wasm_bindgen::from_value(row.clone()).map_err(|err| query_error(sql, err.into()))
477}
478
479fn to_js(param: Param) -> JsValue {
480    match param.0 {
481        Value::Null => JsValue::NULL,
482        Value::Number(n) => JsValue::from_f64(n),
483        Value::Text(s) => JsValue::from(s),
484    }
485}
486
487fn query_error(sql: &str, err: worker::Error) -> Error {
488    Error::internal(format!("D1 query failed: {err}. SQL: {sql}"))
489}