1use crate::wit_bindgen;
25use std::sync::Arc;
26
27#[doc(hidden)]
28pub mod wit {
32 #![allow(missing_docs)]
33 use crate::wit_bindgen;
34
35 wit_bindgen::generate!({
36 runtime_path: "crate::wit_bindgen::rt",
37 world: "spin-sdk-mysql-v3",
38 path: "wit",
39 generate_all,
40 });
41
42 pub use spin::mysql::mysql;
43}
44
45pub struct Connection(wit::mysql::Connection);
113
114impl Connection {
115 pub async fn open(address: impl Into<String>) -> Result<Self, Error> {
120 let inner = wit::mysql::Connection::open(address.into()).await?;
121 Ok(Self(inner))
122 }
123
124 pub async fn query(
129 &self,
130 statement: impl Into<String>,
131 params: impl Into<Vec<ParameterValue>>,
132 ) -> Result<QueryResult, Error> {
133 let (columns, rows, result) = self.0.query(statement.into(), params.into()).await?;
134 Ok(QueryResult {
135 columns: Arc::new(columns),
136 rows,
137 result,
138 })
139 }
140
141 pub async fn execute(
146 &self,
147 statement: impl Into<String>,
148 params: impl Into<Vec<ParameterValue>>,
149 ) -> Result<(), Error> {
150 self.0
151 .execute(statement.into(), params.into())
152 .await
153 .map_err(Error::MysqlError)
154 }
155}
156
157#[doc(inline)]
158pub use wit::mysql::Error as MysqlError;
159
160#[doc(inline)]
161pub use wit::mysql::{Column, DbDataType, DbValue, ParameterValue};
162
163pub struct QueryResult {
165 columns: Arc<Vec<Column>>,
166 rows: wit_bindgen::StreamReader<Vec<DbValue>>,
167 result: wit_bindgen::FutureReader<Result<(), MysqlError>>,
168}
169
170impl QueryResult {
171 pub fn columns(&self) -> &[Column] {
173 &self.columns
174 }
175
176 pub async fn next(&mut self) -> Option<Row> {
182 self.rows.next().await.map(|r| Row {
183 columns: self.columns.clone(),
184 result: r,
185 })
186 }
187
188 pub async fn result(self) -> Result<(), Error> {
190 self.result.await.map_err(Error::MysqlError)
191 }
192
193 pub async fn collect(mut self) -> Result<Vec<Row>, Error> {
198 let mut rows = vec![];
199 while let Some(row) = self.next().await {
200 rows.push(row);
201 }
202 self.result.await.map_err(Error::MysqlError)?;
203 Ok(rows)
204 }
205
206 pub fn rows(&mut self) -> &mut wit_bindgen::StreamReader<Vec<DbValue>> {
217 &mut self.rows
218 }
219
220 #[allow(
222 clippy::type_complexity,
223 reason = "sorry clippy that's just what the inner bits are"
224 )]
225 pub fn into_inner(
226 self,
227 ) -> (
228 Vec<Column>,
229 wit_bindgen::StreamReader<Vec<DbValue>>,
230 wit_bindgen::FutureReader<Result<(), MysqlError>>,
231 ) {
232 ((*self.columns).clone(), self.rows, self.result)
233 }
234}
235
236pub struct Row {
243 columns: Arc<Vec<wit::mysql::Column>>,
244 result: Vec<DbValue>,
245}
246
247impl Row {
248 pub fn get<T: Decode>(&self, column: &str) -> Option<T> {
282 let i = self.columns.iter().position(|c| c.name == column)?;
283 let db_value = self.result.get(i)?;
284 Decode::decode(db_value).ok()
285 }
286}
287
288impl std::ops::Index<usize> for Row {
289 type Output = DbValue;
290
291 fn index(&self, index: usize) -> &Self::Output {
292 &self.result[index]
293 }
294}
295
296#[derive(Debug, thiserror::Error)]
298pub enum Error {
299 #[error("error value decoding: {0}")]
301 Decode(String),
302 #[error(transparent)]
304 MysqlError(#[from] MysqlError),
305}
306
307pub trait Decode: Sized {
309 fn decode(value: &DbValue) -> Result<Self, Error>;
311}
312
313impl<T> Decode for Option<T>
314where
315 T: Decode,
316{
317 fn decode(value: &DbValue) -> Result<Self, Error> {
318 match value {
319 DbValue::DbNull => Ok(None),
320 v => Ok(Some(T::decode(v)?)),
321 }
322 }
323}
324
325impl Decode for bool {
326 fn decode(value: &DbValue) -> Result<Self, Error> {
327 match value {
328 DbValue::Int8(0) => Ok(false),
329 DbValue::Int8(1) => Ok(true),
330 _ => Err(Error::Decode(format_decode_err(
331 "TINYINT(1), BOOLEAN",
332 value,
333 ))),
334 }
335 }
336}
337
338impl Decode for i8 {
339 fn decode(value: &DbValue) -> Result<Self, Error> {
340 match value {
341 DbValue::Int8(n) => Ok(*n),
342 _ => Err(Error::Decode(format_decode_err("TINYINT", value))),
343 }
344 }
345}
346
347impl Decode for i16 {
348 fn decode(value: &DbValue) -> Result<Self, Error> {
349 match value {
350 DbValue::Int16(n) => Ok(*n),
351 _ => Err(Error::Decode(format_decode_err("SMALLINT", value))),
352 }
353 }
354}
355
356impl Decode for i32 {
357 fn decode(value: &DbValue) -> Result<Self, Error> {
358 match value {
359 DbValue::Int32(n) => Ok(*n),
360 _ => Err(Error::Decode(format_decode_err("INT", value))),
361 }
362 }
363}
364
365impl Decode for i64 {
366 fn decode(value: &DbValue) -> Result<Self, Error> {
367 match value {
368 DbValue::Int64(n) => Ok(*n),
369 _ => Err(Error::Decode(format_decode_err("BIGINT", value))),
370 }
371 }
372}
373
374impl Decode for u8 {
375 fn decode(value: &DbValue) -> Result<Self, Error> {
376 match value {
377 DbValue::Uint8(n) => Ok(*n),
378 _ => Err(Error::Decode(format_decode_err("UNSIGNED TINYINT", value))),
379 }
380 }
381}
382
383impl Decode for u16 {
384 fn decode(value: &DbValue) -> Result<Self, Error> {
385 match value {
386 DbValue::Uint16(n) => Ok(*n),
387 _ => Err(Error::Decode(format_decode_err("UNSIGNED SMALLINT", value))),
388 }
389 }
390}
391
392impl Decode for u32 {
393 fn decode(value: &DbValue) -> Result<Self, Error> {
394 match value {
395 DbValue::Uint32(n) => Ok(*n),
396 _ => Err(Error::Decode(format_decode_err(
397 "UNISIGNED MEDIUMINT, UNSIGNED INT",
398 value,
399 ))),
400 }
401 }
402}
403
404impl Decode for u64 {
405 fn decode(value: &DbValue) -> Result<Self, Error> {
406 match value {
407 DbValue::Uint64(n) => Ok(*n),
408 _ => Err(Error::Decode(format_decode_err("UNSIGNED BIGINT", value))),
409 }
410 }
411}
412
413impl Decode for f32 {
414 fn decode(value: &DbValue) -> Result<Self, Error> {
415 match value {
416 DbValue::Floating32(n) => Ok(*n),
417 _ => Err(Error::Decode(format_decode_err("FLOAT", value))),
418 }
419 }
420}
421
422impl Decode for f64 {
423 fn decode(value: &DbValue) -> Result<Self, Error> {
424 match value {
425 DbValue::Floating64(n) => Ok(*n),
426 _ => Err(Error::Decode(format_decode_err("DOUBLE", value))),
427 }
428 }
429}
430
431impl Decode for Vec<u8> {
432 fn decode(value: &DbValue) -> Result<Self, Error> {
433 match value {
434 DbValue::Binary(n) => Ok(n.to_owned()),
435 _ => Err(Error::Decode(format_decode_err("BINARY, VARBINARY", value))),
436 }
437 }
438}
439
440impl Decode for String {
441 fn decode(value: &DbValue) -> Result<Self, Error> {
442 match value {
443 DbValue::Str(s) => Ok(s.to_owned()),
444 _ => Err(Error::Decode(format_decode_err(
445 "CHAR, VARCHAR, TEXT",
446 value,
447 ))),
448 }
449 }
450}
451
452macro_rules! impl_parameter_value_conversions {
453 ($($ty:ty => $id:ident),*) => {
454 $(
455 impl From<$ty> for ParameterValue {
456 fn from(v: $ty) -> ParameterValue {
457 ParameterValue::$id(v)
458 }
459 }
460 )*
461 };
462}
463
464impl_parameter_value_conversions! {
465 i8 => Int8,
466 i16 => Int16,
467 i32 => Int32,
468 i64 => Int64,
469 f32 => Floating32,
470 f64 => Floating64,
471 bool => Boolean,
472 String => Str,
473 Vec<u8> => Binary
474}
475
476fn format_decode_err(types: &str, value: &DbValue) -> String {
477 format!("Expected {} from the DB but got {:?}", types, value)
478}
479
480#[cfg(test)]
481mod tests {
482 use super::*;
483
484 #[test]
485 fn boolean() {
486 assert!(bool::decode(&DbValue::Int8(1)).unwrap());
487 assert!(bool::decode(&DbValue::Int8(3)).is_err());
488 assert!(bool::decode(&DbValue::Int32(0)).is_err());
489 assert!(Option::<bool>::decode(&DbValue::DbNull).unwrap().is_none());
490 }
491
492 #[test]
493 fn int8() {
494 assert_eq!(i8::decode(&DbValue::Int8(0)).unwrap(), 0);
495 assert!(i8::decode(&DbValue::Int32(0)).is_err());
496 assert!(Option::<i8>::decode(&DbValue::DbNull).unwrap().is_none());
497 }
498
499 #[test]
500 fn int16() {
501 assert_eq!(i16::decode(&DbValue::Int16(0)).unwrap(), 0);
502 assert!(i16::decode(&DbValue::Int32(0)).is_err());
503 assert!(Option::<i16>::decode(&DbValue::DbNull).unwrap().is_none());
504 }
505
506 #[test]
507 fn int32() {
508 assert_eq!(i32::decode(&DbValue::Int32(0)).unwrap(), 0);
509 assert!(i32::decode(&DbValue::Boolean(false)).is_err());
510 assert!(Option::<i32>::decode(&DbValue::DbNull).unwrap().is_none());
511 }
512
513 #[test]
514 fn int64() {
515 assert_eq!(i64::decode(&DbValue::Int64(0)).unwrap(), 0);
516 assert!(i64::decode(&DbValue::Boolean(false)).is_err());
517 assert!(Option::<i64>::decode(&DbValue::DbNull).unwrap().is_none());
518 }
519
520 #[test]
521 fn uint8() {
522 assert_eq!(u8::decode(&DbValue::Uint8(0)).unwrap(), 0);
523 assert!(u8::decode(&DbValue::Uint32(0)).is_err());
524 assert!(Option::<u16>::decode(&DbValue::DbNull).unwrap().is_none());
525 }
526
527 #[test]
528 fn uint16() {
529 assert_eq!(u16::decode(&DbValue::Uint16(0)).unwrap(), 0);
530 assert!(u16::decode(&DbValue::Uint32(0)).is_err());
531 assert!(Option::<u16>::decode(&DbValue::DbNull).unwrap().is_none());
532 }
533
534 #[test]
535 fn uint32() {
536 assert_eq!(u32::decode(&DbValue::Uint32(0)).unwrap(), 0);
537 assert!(u32::decode(&DbValue::Boolean(false)).is_err());
538 assert!(Option::<u32>::decode(&DbValue::DbNull).unwrap().is_none());
539 }
540
541 #[test]
542 fn uint64() {
543 assert_eq!(u64::decode(&DbValue::Uint64(0)).unwrap(), 0);
544 assert!(u64::decode(&DbValue::Boolean(false)).is_err());
545 assert!(Option::<u64>::decode(&DbValue::DbNull).unwrap().is_none());
546 }
547
548 #[test]
549 fn floating32() {
550 assert!(f32::decode(&DbValue::Floating32(0.0)).is_ok());
551 assert!(f32::decode(&DbValue::Boolean(false)).is_err());
552 assert!(Option::<f32>::decode(&DbValue::DbNull).unwrap().is_none());
553 }
554
555 #[test]
556 fn floating64() {
557 assert!(f64::decode(&DbValue::Floating64(0.0)).is_ok());
558 assert!(f64::decode(&DbValue::Boolean(false)).is_err());
559 assert!(Option::<f64>::decode(&DbValue::DbNull).unwrap().is_none());
560 }
561
562 #[test]
563 fn str() {
564 assert_eq!(
565 String::decode(&DbValue::Str(String::from("foo"))).unwrap(),
566 String::from("foo")
567 );
568
569 assert!(String::decode(&DbValue::Int32(0)).is_err());
570 assert!(
571 Option::<String>::decode(&DbValue::DbNull)
572 .unwrap()
573 .is_none()
574 );
575 }
576}