Skip to main content

spin_factor_sqlite/
host.rs

1use std::collections::{HashMap, HashSet};
2use std::sync::Arc;
3
4use opentelemetry_semantic_conventions::attribute as otel_attribute;
5use spin_core::wasmtime::component::{Accessor, FutureReader, StreamReader};
6use spin_factor_otel::OtelFactorState;
7use spin_factors::wasmtime::component::Resource;
8use spin_factors::{SelfInstanceBuilder, anyhow};
9use spin_world::MAX_HOST_BUFFERED_BYTES;
10use spin_world::spin::sqlite3_1_0::sqlite as v3;
11use spin_world::v1::sqlite as v1;
12use spin_world::v2::sqlite as v2;
13use tracing::field::Empty;
14use tracing::{Level, instrument};
15
16use crate::{Connection, ConnectionCreator, QueryAsyncResult};
17
18pub struct InstanceState {
19    allowed_databases: Arc<HashSet<String>>,
20    /// A resource table of connections.
21    connections: spin_resource_table::Table<Arc<dyn Connection>>,
22    /// A map from database label to connection creators.
23    connection_creators: HashMap<String, Arc<dyn ConnectionCreator>>,
24    otel: OtelFactorState,
25}
26
27impl InstanceState {
28    /// Create a new `InstanceState`
29    ///
30    /// Takes the list of allowed databases, and a function for getting a connection creator given a database label.
31    pub fn new(
32        allowed_databases: Arc<HashSet<String>>,
33        connection_creators: HashMap<String, Arc<dyn ConnectionCreator>>,
34        otel: OtelFactorState,
35    ) -> Self {
36        Self {
37            allowed_databases,
38            connections: spin_resource_table::Table::new(256),
39            connection_creators,
40            otel,
41        }
42    }
43
44    /// Get a connection for a given database label.
45    fn get_connection<T: 'static>(
46        &self,
47        connection: Resource<T>,
48    ) -> Result<Arc<dyn Connection>, v3::Error> {
49        self.connections
50            .get(connection.rep())
51            .cloned()
52            .ok_or(v3::Error::InvalidConnection)
53    }
54
55    async fn open_impl<T: 'static>(&mut self, database: String) -> Result<Resource<T>, v3::Error> {
56        if !self.allowed_databases.contains(&database) {
57            return Err(v3::Error::AccessDenied);
58        }
59        let conn = self
60            .connection_creators
61            .get(&database)
62            .ok_or(v3::Error::NoSuchDatabase)?
63            .create_connection(&database)
64            .await?;
65        tracing::Span::current().record(
66            "sqlite.backend",
67            conn.summary().as_deref().unwrap_or("unknown"),
68        );
69        self.connections
70            .push(conn)
71            .map_err(|()| v3::Error::Io("too many connections opened".to_string()))
72            .map(Resource::new_own)
73    }
74
75    async fn execute_impl<T: 'static>(
76        &mut self,
77        connection: Resource<T>,
78        query: String,
79        parameters: Vec<v3::Value>,
80    ) -> Result<v3::QueryResult, v3::Error> {
81        let conn = self.get_connection(connection)?;
82        tracing::Span::current().record(
83            "sqlite.backend",
84            conn.summary().as_deref().unwrap_or("unknown"),
85        );
86        conn.query(&query, parameters, MAX_HOST_BUFFERED_BYTES)
87            .await
88    }
89
90    /// Get the set of allowed databases.
91    pub fn allowed_databases(&self) -> &HashSet<String> {
92        &self.allowed_databases
93    }
94}
95
96impl SelfInstanceBuilder for InstanceState {}
97
98impl v3::Host for InstanceState {
99    fn convert_error(&mut self, error: v3::Error) -> anyhow::Result<v3::Error> {
100        Ok(error)
101    }
102}
103
104impl v3::HostConnection for InstanceState {
105    #[instrument(name = "spin_sqlite.open", skip(self), err(level = Level::INFO),
106        fields(otel.kind = "client", {otel_attribute::DB_SYSTEM_NAME} = "sqlite", sqlite.backend = Empty))]
107    async fn open(&mut self, database: String) -> Result<Resource<v3::Connection>, v3::Error> {
108        self.open_impl(database).await
109    }
110
111    #[instrument(name = "spin_sqlite.execute", skip(self, connection, parameters), err(level = Level::INFO),
112        fields(otel.kind = "client", {otel_attribute::DB_SYSTEM_NAME} = "sqlite", sqlite.backend = Empty))]
113    async fn execute(
114        &mut self,
115        connection: Resource<v3::Connection>,
116        query: String,
117        parameters: Vec<v3::Value>,
118    ) -> Result<v3::QueryResult, v3::Error> {
119        self.execute_impl(connection, query, parameters).await
120    }
121
122    async fn changes(&mut self, connection: Resource<v3::Connection>) -> anyhow::Result<u64> {
123        let conn = match self.get_connection(connection) {
124            Ok(c) => c,
125            Err(err) => return Err(err.into()),
126        };
127        tracing::Span::current().record(
128            "sqlite.backend",
129            conn.summary().as_deref().unwrap_or("unknown"),
130        );
131        conn.changes().await.map_err(|e| e.into())
132    }
133
134    async fn last_insert_rowid(
135        &mut self,
136        connection: Resource<v3::Connection>,
137    ) -> anyhow::Result<i64> {
138        let conn = match self.get_connection(connection) {
139            Ok(c) => c,
140            Err(err) => return Err(err.into()),
141        };
142        tracing::Span::current().record(
143            "sqlite.backend",
144            conn.summary().as_deref().unwrap_or("unknown"),
145        );
146        conn.last_insert_rowid().await.map_err(|e| e.into())
147    }
148
149    async fn drop(&mut self, connection: Resource<v3::Connection>) -> anyhow::Result<()> {
150        let _ = self.connections.remove(connection.rep());
151        Ok(())
152    }
153}
154
155impl<T> v3::HostConnectionWithStore<T> for crate::SqliteFactorData {
156    async fn open_async(
157        accessor: &Accessor<T, Self>,
158        database: String,
159    ) -> Result<Resource<v3::Connection>, v3::Error> {
160        // TODO: this duplicates `open_impl` logic but split up to move
161        // in and out of the Accessor. How to dedupe?
162        let conn_creator = accessor.with(|mut access| {
163            let host = access.get();
164            if !host.allowed_databases.contains(&database) {
165                return Err(v3::Error::AccessDenied);
166            }
167            host.connection_creators
168                .get(&database)
169                .ok_or(v3::Error::NoSuchDatabase)
170                .cloned()
171        })?;
172
173        let conn = conn_creator.create_connection(&database).await?;
174
175        tracing::Span::current().record(
176            "sqlite.backend",
177            conn.summary().as_deref().unwrap_or("unknown"),
178        );
179
180        accessor.with(|mut access| {
181            let host = access.get();
182            host.connections
183                .push(conn)
184                .map_err(|()| v3::Error::Io("too many connections opened".to_string()))
185                .map(Resource::new_own)
186        })
187    }
188
189    async fn execute_async(
190        accessor: &Accessor<T, Self>,
191        connection: Resource<v3::Connection>,
192        query: String,
193        parameters: Vec<v3::Value>,
194    ) -> Result<
195        (
196            Vec<String>,
197            StreamReader<v3::RowResult>,
198            FutureReader<Result<(), v3::Error>>,
199        ),
200        v3::Error,
201    > {
202        let conn = accessor.with(|mut access| {
203            let host = access.get();
204            host.get_connection(connection)
205        })?;
206
207        tracing::Span::current().record(
208            "sqlite.backend",
209            conn.summary().as_deref().unwrap_or("unknown"),
210        );
211
212        let QueryAsyncResult {
213            columns,
214            rows,
215            error,
216        } = conn
217            .query_async(&query, parameters, MAX_HOST_BUFFERED_BYTES)
218            .await?;
219        let row_producer = spin_wasi_async::stream::producer(rows);
220
221        let (sr, efr) = accessor
222            .with(|mut access| {
223                let sr = StreamReader::new(&mut access, row_producer)?;
224                let efr = FutureReader::new(&mut access, error)?;
225                anyhow::Ok((sr, efr))
226            })
227            .map_err(|e| v3::Error::Io(e.to_string()))?;
228
229        Ok((columns, sr, efr))
230    }
231
232    async fn changes_async(
233        accessor: &Accessor<T, Self>,
234        connection: Resource<v3::Connection>,
235    ) -> anyhow::Result<u64> {
236        let conn = accessor.with(|mut access| {
237            let host = access.get();
238            host.get_connection(connection)
239        });
240
241        let conn = match conn {
242            Ok(c) => c,
243            Err(err) => return Err(err.into()),
244        };
245        tracing::Span::current().record(
246            "sqlite.backend",
247            conn.summary().as_deref().unwrap_or("unknown"),
248        );
249        conn.changes().await.map_err(|e| e.into())
250    }
251
252    async fn last_insert_rowid_async(
253        accessor: &Accessor<T, Self>,
254        connection: Resource<v3::Connection>,
255    ) -> anyhow::Result<i64> {
256        let conn = accessor.with(|mut access| {
257            let host = access.get();
258            host.get_connection(connection)
259        });
260
261        let conn = match conn {
262            Ok(c) => c,
263            Err(err) => return Err(err.into()),
264        };
265        tracing::Span::current().record(
266            "sqlite.backend",
267            conn.summary().as_deref().unwrap_or("unknown"),
268        );
269        conn.last_insert_rowid().await.map_err(|e| e.into())
270    }
271}
272
273impl v2::Host for InstanceState {
274    fn convert_error(&mut self, error: v2::Error) -> anyhow::Result<v2::Error> {
275        Ok(error)
276    }
277}
278
279impl v2::HostConnection for InstanceState {
280    #[instrument(name = "spin_sqlite.open", skip(self), err(level = Level::INFO),
281        fields(otel.kind = "client", {otel_attribute::DB_SYSTEM_NAME} = "sqlite", sqlite.backend = Empty))]
282    async fn open(&mut self, database: String) -> Result<Resource<v2::Connection>, v2::Error> {
283        self.otel.reparent_tracing_span();
284        self.open_impl(database).await.map_err(to_v2_error)
285    }
286
287    #[instrument(name = "spin_sqlite.execute", skip(self, connection, parameters), err(level = Level::INFO),
288        fields(otel.kind = "client", {otel_attribute::DB_SYSTEM_NAME} = "sqlite", sqlite.backend = Empty))]
289    async fn execute(
290        &mut self,
291        connection: Resource<v2::Connection>,
292        query: String,
293        parameters: Vec<v2::Value>,
294    ) -> Result<v2::QueryResult, v2::Error> {
295        self.otel.reparent_tracing_span();
296        self.execute_impl(
297            connection,
298            query,
299            parameters.into_iter().map(from_v2_value).collect(),
300        )
301        .await
302        .map(to_v2_query_result)
303        .map_err(to_v2_error)
304    }
305
306    async fn drop(&mut self, connection: Resource<v2::Connection>) -> anyhow::Result<()> {
307        let _ = self.connections.remove(connection.rep());
308        Ok(())
309    }
310}
311
312impl v1::Host for InstanceState {
313    async fn open(&mut self, database: String) -> Result<u32, v1::Error> {
314        let result = <Self as v3::HostConnection>::open(self, database).await;
315        result.map_err(to_legacy_error).map(|s| s.rep())
316    }
317
318    async fn execute(
319        &mut self,
320        connection: u32,
321        query: String,
322        parameters: Vec<spin_world::v1::sqlite::Value>,
323    ) -> Result<spin_world::v1::sqlite::QueryResult, v1::Error> {
324        let this = Resource::new_borrow(connection);
325        let result = <Self as v3::HostConnection>::execute(
326            self,
327            this,
328            query,
329            parameters.into_iter().map(from_legacy_value).collect(),
330        )
331        .await;
332        result.map_err(to_legacy_error).map(to_legacy_query_result)
333    }
334
335    async fn close(&mut self, connection: u32) -> anyhow::Result<()> {
336        <Self as v2::HostConnection>::drop(self, Resource::new_own(connection)).await
337    }
338
339    fn convert_error(&mut self, error: v1::Error) -> anyhow::Result<v1::Error> {
340        Ok(error)
341    }
342}
343
344fn to_v2_error(error: v3::Error) -> v2::Error {
345    match error {
346        v3::Error::NoSuchDatabase => v2::Error::NoSuchDatabase,
347        v3::Error::AccessDenied => v2::Error::AccessDenied,
348        v3::Error::InvalidConnection => v2::Error::InvalidConnection,
349        v3::Error::DatabaseFull => v2::Error::DatabaseFull,
350        v3::Error::Io(s) => v2::Error::Io(s),
351    }
352}
353
354fn to_legacy_error(error: v3::Error) -> v1::Error {
355    match error {
356        v3::Error::NoSuchDatabase => v1::Error::NoSuchDatabase,
357        v3::Error::AccessDenied => v1::Error::AccessDenied,
358        v3::Error::InvalidConnection => v1::Error::InvalidConnection,
359        v3::Error::DatabaseFull => v1::Error::DatabaseFull,
360        v3::Error::Io(s) => v1::Error::Io(s),
361    }
362}
363
364fn to_v2_query_result(result: v3::QueryResult) -> v2::QueryResult {
365    v2::QueryResult {
366        columns: result.columns,
367        rows: result.rows.into_iter().map(to_v2_row_result).collect(),
368    }
369}
370
371fn to_legacy_query_result(result: v3::QueryResult) -> v1::QueryResult {
372    v1::QueryResult {
373        columns: result.columns,
374        rows: result.rows.into_iter().map(to_legacy_row_result).collect(),
375    }
376}
377
378fn to_v2_row_result(result: v3::RowResult) -> v2::RowResult {
379    v2::RowResult {
380        values: result.values.into_iter().map(to_v2_value).collect(),
381    }
382}
383
384fn to_legacy_row_result(result: v3::RowResult) -> v1::RowResult {
385    v1::RowResult {
386        values: result.values.into_iter().map(to_legacy_value).collect(),
387    }
388}
389
390fn to_v2_value(value: v3::Value) -> v2::Value {
391    match value {
392        v3::Value::Integer(i) => v2::Value::Integer(i),
393        v3::Value::Real(r) => v2::Value::Real(r),
394        v3::Value::Text(t) => v2::Value::Text(t),
395        v3::Value::Blob(b) => v2::Value::Blob(b),
396        v3::Value::Null => v2::Value::Null,
397    }
398}
399
400fn to_legacy_value(value: v3::Value) -> v1::Value {
401    match value {
402        v3::Value::Integer(i) => v1::Value::Integer(i),
403        v3::Value::Real(r) => v1::Value::Real(r),
404        v3::Value::Text(t) => v1::Value::Text(t),
405        v3::Value::Blob(b) => v1::Value::Blob(b),
406        v3::Value::Null => v1::Value::Null,
407    }
408}
409
410fn from_v2_value(value: v2::Value) -> v3::Value {
411    match value {
412        v2::Value::Integer(i) => v3::Value::Integer(i),
413        v2::Value::Real(r) => v3::Value::Real(r),
414        v2::Value::Text(t) => v3::Value::Text(t),
415        v2::Value::Blob(b) => v3::Value::Blob(b),
416        v2::Value::Null => v3::Value::Null,
417    }
418}
419
420fn from_legacy_value(value: v1::Value) -> v3::Value {
421    match value {
422        v1::Value::Integer(i) => v3::Value::Integer(i),
423        v1::Value::Real(r) => v3::Value::Real(r),
424        v1::Value::Text(t) => v3::Value::Text(t),
425        v1::Value::Blob(b) => v3::Value::Blob(b),
426        v1::Value::Null => v3::Value::Null,
427    }
428}