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 connections: spin_resource_table::Table<Arc<dyn Connection>>,
22 connection_creators: HashMap<String, Arc<dyn ConnectionCreator>>,
24 otel: OtelFactorState,
25}
26
27impl InstanceState {
28 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 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 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 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}