use crate::value::Value;
use std::collections::HashMap;
#[derive(Debug, Clone, PartialEq)]
pub enum QueryError {
ColumnCountMismatch {
expected: usize,
actual: usize,
},
TypeMismatch {
column: std::borrow::Cow<'static, str>,
expected: &'static str,
},
MissingColumn {
column: &'static str,
},
Custom(String),
}
impl std::fmt::Display for QueryError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
QueryError::ColumnCountMismatch { expected, actual } => {
write!(f, "列数不匹配: 期望 {}, 实际 {}", expected, actual)
}
QueryError::TypeMismatch { column, expected } => {
write!(f, "列 {:?} 类型不匹配, 期望 {}", column, expected)
}
QueryError::MissingColumn { column } => {
write!(f, "缺少列: {}", column)
}
QueryError::Custom(msg) => write!(f, "{}", msg),
}
}
}
impl std::error::Error for QueryError {}
#[derive(Debug, Clone)]
pub struct RowDesc {
pub columns: Vec<String>,
}
impl RowDesc {
pub fn new(columns: Vec<String>) -> Self {
Self { columns }
}
pub fn len(&self) -> usize {
self.columns.len()
}
pub fn is_empty(&self) -> bool {
self.columns.is_empty()
}
pub fn index_of(&self, name: &str) -> Option<usize> {
self.columns.iter().position(|c| c == name)
}
}
pub trait Queryable: Sized {
fn from_values(values: Vec<Value>) -> Result<Self, QueryError>;
fn from_values_with_desc(values: Vec<Value>, desc: &RowDesc) -> Result<Self, QueryError> {
if values.len() != desc.len() {
return Err(QueryError::ColumnCountMismatch {
expected: desc.len(),
actual: values.len(),
});
}
Self::from_values(values)
}
}
pub trait FromRow: Sized {
fn from_row(row: HashMap<String, Value>) -> Result<Self, QueryError>;
}
impl Queryable for Value {
fn from_values(values: Vec<Value>) -> Result<Self, QueryError> {
if values.len() != 1 {
return Err(QueryError::ColumnCountMismatch {
expected: 1,
actual: values.len(),
});
}
Ok(values.into_iter().next().unwrap()) }
}
impl Queryable for (Value, Value) {
fn from_values(values: Vec<Value>) -> Result<Self, QueryError> {
if values.len() != 2 {
return Err(QueryError::ColumnCountMismatch {
expected: 2,
actual: values.len(),
});
}
let mut iter = values.into_iter();
Ok((iter.next().unwrap(), iter.next().unwrap())) }
}
impl Queryable for (Value, Value, Value) {
fn from_values(values: Vec<Value>) -> Result<Self, QueryError> {
if values.len() != 3 {
return Err(QueryError::ColumnCountMismatch {
expected: 3,
actual: values.len(),
});
}
let mut iter = values.into_iter();
Ok((
iter.next().unwrap(), iter.next().unwrap(),
iter.next().unwrap(),
))
}
}
pub fn value_as_i64(v: &Value) -> Option<i64> {
v.as_i64()
}
pub fn value_as_f64(v: &Value) -> Option<f64> {
v.as_f64()
}
pub fn value_as_string(v: &Value) -> Option<String> {
v.as_str().map(|s| s.to_string())
}
pub fn value_as_bool(v: &Value) -> Option<bool> {
v.as_bool()
}
pub fn value_as_nullable_i64(v: &Value) -> Option<i64> {
if v.is_null() {
None
} else {
v.as_i64()
}
}
pub fn value_as_nullable_string(v: &Value) -> Option<String> {
if v.is_null() {
None
} else {
v.as_str().map(|s| s.to_string())
}
}
pub struct Query {
sql: String,
}
impl Query {
pub fn new(sql: impl Into<String>) -> Self {
Self { sql: sql.into() }
}
pub fn sql(&self) -> &str {
&self.sql
}
}
impl std::fmt::Display for Query {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.sql)
}
}
impl Query {
pub async fn fetch_all(
&self,
conn: &mut dyn crate::Connection,
) -> Result<Vec<std::collections::HashMap<String, crate::value::Value>>, crate::DbError> {
conn.query(&self.sql).await
}
pub async fn fetch_one(
&self,
conn: &mut dyn crate::Connection,
) -> Result<std::collections::HashMap<String, crate::value::Value>, crate::DbError> {
let rows = conn.query(&self.sql).await?;
match rows.len() {
0 => Err(crate::DbError::NotFound(
"fetch_one: no rows returned".into(),
)),
1 => Ok(rows.into_iter().next().unwrap()), n => Err(crate::DbError::QueryError(format!(
"fetch_one expected 1 row, got {}",
n
))),
}
}
pub async fn fetch_optional(
&self,
conn: &mut dyn crate::Connection,
) -> Result<Option<std::collections::HashMap<String, crate::value::Value>>, crate::DbError>
{
let rows = conn.query(&self.sql).await?;
Ok(rows.into_iter().next())
}
}
pub struct QueryAs<T> {
sql: String,
_marker: std::marker::PhantomData<T>,
}
impl<T> QueryAs<T> {
pub fn new(sql: impl Into<String>) -> Self {
Self {
sql: sql.into(),
_marker: std::marker::PhantomData,
}
}
pub fn sql(&self) -> &str {
&self.sql
}
}
impl<T> std::fmt::Display for QueryAs<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.sql)
}
}
impl<T: crate::value::FromQueryResult> QueryAs<T> {
pub async fn fetch_all(
&self,
conn: &mut dyn crate::Connection,
) -> Result<Vec<T>, crate::DbError> {
let rows = conn.query(&self.sql).await?;
if rows.is_empty() {
return Ok(Vec::new());
}
let expected: Vec<&str> = T::row_desc();
if !expected.is_empty() {
let actual: Vec<&str> = rows[0].keys().map(|s| s.as_str()).collect();
validate_columns(&actual, &expected)?;
validate_column_types(&rows[0], T::column_types())?;
}
rows.into_iter()
.map(|row| {
T::from_query_result(&row).map_err(|e| crate::DbError::QueryError(e.to_string()))
})
.collect()
}
pub async fn fetch_one(&self, conn: &mut dyn crate::Connection) -> Result<T, crate::DbError> {
let rows = conn.query(&self.sql).await?;
match rows.len() {
0 => Err(crate::DbError::NotFound(
"fetch_one: no rows returned".into(),
)),
1 => {
let row = &rows[0];
let expected: Vec<&str> = T::row_desc();
if !expected.is_empty() {
let actual: Vec<&str> = row.keys().map(|s| s.as_str()).collect();
validate_columns(&actual, &expected)?;
validate_column_types(row, T::column_types())?;
}
T::from_query_result(row).map_err(|e| crate::DbError::QueryError(e.to_string()))
}
n => Err(crate::DbError::QueryError(format!(
"fetch_one expected 1 row, got {}",
n
))),
}
}
pub async fn fetch_optional(
&self,
conn: &mut dyn crate::Connection,
) -> Result<Option<T>, crate::DbError> {
let rows = conn.query(&self.sql).await?;
match rows.len() {
0 => Ok(None),
_ => {
let row = &rows[0];
let expected: Vec<&str> = T::row_desc();
if !expected.is_empty() {
let actual: Vec<&str> = row.keys().map(|s| s.as_str()).collect();
validate_columns(&actual, &expected)?;
validate_column_types(row, T::column_types())?;
}
T::from_query_result(row)
.map_err(|e| crate::DbError::QueryError(e.to_string()))
.map(Some)
}
}
}
}
fn validate_columns(actual: &[&str], expected: &[&str]) -> Result<(), crate::DbError> {
use std::collections::HashSet;
let actual_set: HashSet<&str> = actual.iter().copied().collect();
let expected_set: HashSet<&str> = expected.iter().copied().collect();
let missing: Vec<&str> = expected_set.difference(&actual_set).copied().collect();
let extra: Vec<&str> = actual_set.difference(&expected_set).copied().collect();
if !missing.is_empty() || !extra.is_empty() {
return Err(crate::DbError::QueryError(format!(
"query_as! column mismatch: struct expects {:?}, DB returned {:?}{}",
expected,
actual,
if missing.is_empty() {
String::new()
} else {
format!(" (missing: {:?})", missing)
}
)));
}
Ok(())
}
fn validate_column_types(
row: &std::collections::HashMap<String, crate::value::Value>,
expected: &[(&str, &str)],
) -> Result<(), crate::DbError> {
use crate::value::{ColType, Value};
for (col_name, expected_type) in expected {
let value = match row.get(*col_name) {
Some(v) => v,
None => continue, };
if matches!(value, Value::Null) {
continue; }
let actual_type = value_to_sql_type_name(value);
let actual_col = ColType::from_type_name(actual_type);
let expected_col = ColType::from_type_name(expected_type);
if actual_col != expected_col
&& actual_col != ColType::Unknown
&& expected_col != ColType::Unknown
{
return Err(crate::DbError::QueryError(format!(
"query_as! column TYPE mismatch: column '{}' has actual type '{}' (ColType::{:?}), \
but struct expects '{}' (ColType::{:?})",
col_name, actual_type, actual_col, expected_type, expected_col
)));
}
}
Ok(())
}
fn value_to_sql_type_name(value: &Value) -> &'static str {
match value {
Value::Null => "NULL",
Value::Bool(_) => "BOOLEAN",
Value::I8(_) | Value::U8(_) => "TINYINT",
Value::I16(_) | Value::U16(_) => "SMALLINT",
Value::I32(_) | Value::U32(_) => "INT",
Value::I64(_) | Value::U64(_) => "BIGINT",
Value::F32(_) => "FLOAT",
Value::F64(_) => "DOUBLE",
Value::Decimal(_) => "DECIMAL",
Value::String(_)
| Value::Uuid(_)
| Value::Date(_)
| Value::DateTime(_)
| Value::Time(_)
| Value::Json(_) => "VARCHAR",
#[cfg(feature = "perf-box-str")]
Value::BoxedStr(_) => "VARCHAR",
Value::Bytes(_) => "BLOB",
Value::Array(_) | Value::Object(_) => "JSON",
}
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Debug, Default, PartialEq)]
struct UserRow {
id: i64,
name: String,
}
impl Queryable for UserRow {
fn from_values(values: Vec<Value>) -> Result<Self, QueryError> {
if values.len() != 2 {
return Err(QueryError::ColumnCountMismatch {
expected: 2,
actual: values.len(),
});
}
let id = values[0].as_i64().ok_or(QueryError::TypeMismatch {
column: "0".into(),
expected: "i64",
})?;
let name = values[1]
.as_str()
.ok_or(QueryError::TypeMismatch {
column: "1".into(),
expected: "String",
})?
.to_string();
Ok(UserRow { id, name })
}
}
impl FromRow for UserRow {
fn from_row(row: HashMap<String, Value>) -> Result<Self, QueryError> {
let id = row
.get("id")
.ok_or(QueryError::MissingColumn { column: "id" })?
.as_i64()
.ok_or(QueryError::TypeMismatch {
column: "id".into(),
expected: "i64",
})?;
let name = row
.get("name")
.ok_or(QueryError::MissingColumn { column: "name" })?
.as_str()
.ok_or(QueryError::TypeMismatch {
column: "name".into(),
expected: "String",
})?
.to_string();
Ok(UserRow { id, name })
}
}
#[test]
fn test_query_error_display() {
let e = QueryError::ColumnCountMismatch {
expected: 3,
actual: 2,
};
assert!(format!("{}", e).contains("3"));
assert!(format!("{}", e).contains("2"));
let e = QueryError::TypeMismatch {
column: "age".into(),
expected: "i64",
};
assert!(format!("{}", e).contains("age"));
let e = QueryError::MissingColumn { column: "id" };
assert!(format!("{}", e).contains("id"));
let e = QueryError::Custom("custom".into());
assert_eq!(format!("{}", e), "custom");
}
#[test]
fn test_row_desc_basic() {
let desc = RowDesc::new(vec!["id".into(), "name".into(), "age".into()]);
assert_eq!(desc.len(), 3);
assert!(!desc.is_empty());
assert_eq!(desc.index_of("name"), Some(1));
assert_eq!(desc.index_of("missing"), None);
}
#[test]
fn test_row_desc_empty() {
let desc = RowDesc::new(vec![]);
assert!(desc.is_empty());
assert_eq!(desc.len(), 0);
}
#[test]
fn test_user_row_from_values_success() {
let row =
UserRow::from_values(vec![Value::I64(42), Value::String("Alice".into())]).unwrap();
assert_eq!(row.id, 42);
assert_eq!(row.name, "Alice");
}
#[test]
fn test_user_row_from_values_count_mismatch() {
let result = UserRow::from_values(vec![Value::I64(42)]);
assert!(matches!(
result,
Err(QueryError::ColumnCountMismatch {
expected: 2,
actual: 1
})
));
}
#[test]
fn test_user_row_from_values_type_mismatch() {
let result = UserRow::from_values(vec![
Value::String("not_an_int".into()),
Value::String("Alice".into()),
]);
assert!(matches!(result, Err(QueryError::TypeMismatch { .. })));
}
#[test]
fn test_user_row_from_values_with_desc() {
let desc = RowDesc::new(vec!["id".into(), "name".into()]);
let row =
UserRow::from_values_with_desc(vec![Value::I64(1), Value::String("Bob".into())], &desc)
.unwrap();
assert_eq!(row.id, 1);
assert_eq!(row.name, "Bob");
}
#[test]
fn test_user_row_from_values_with_desc_mismatch() {
let desc = RowDesc::new(vec!["id".into(), "name".into(), "age".into()]);
let result =
UserRow::from_values_with_desc(vec![Value::I64(1), Value::String("Bob".into())], &desc);
assert!(matches!(
result,
Err(QueryError::ColumnCountMismatch { .. })
));
}
#[test]
fn test_user_row_from_row_success() {
let mut map = HashMap::new();
map.insert("id".into(), Value::I64(99));
map.insert("name".into(), Value::String("Charlie".into()));
let row = UserRow::from_row(map).unwrap();
assert_eq!(row.id, 99);
assert_eq!(row.name, "Charlie");
}
#[test]
fn test_user_row_from_row_missing_column() {
let mut map = HashMap::new();
map.insert("id".into(), Value::I64(99));
let result = UserRow::from_row(map);
assert!(matches!(
result,
Err(QueryError::MissingColumn { column: "name" })
));
}
#[test]
fn test_user_row_from_row_extra_columns_ignored() {
let mut map = HashMap::new();
map.insert("id".into(), Value::I64(1));
map.insert("name".into(), Value::String("X".into()));
map.insert("extra".into(), Value::String("ignored".into()));
let row = UserRow::from_row(map).unwrap();
assert_eq!(row.id, 1);
}
#[test]
fn test_value_queryable_single() {
let v = Value::from_values(vec![Value::I64(42)]).unwrap();
assert_eq!(v.as_i64(), Some(42));
}
#[test]
fn test_value_queryable_count_mismatch() {
let result = Value::from_values(vec![Value::I64(1), Value::I64(2)]);
assert!(matches!(
result,
Err(QueryError::ColumnCountMismatch { .. })
));
}
#[test]
fn test_tuple_2_queryable() {
let (a, b) =
<(Value, Value)>::from_values(vec![Value::I64(1), Value::String("hello".into())])
.unwrap();
assert_eq!(a.as_i64(), Some(1));
assert_eq!(b.as_str(), Some("hello"));
}
#[test]
fn test_tuple_3_queryable() {
let (a, b, c) = <(Value, Value, Value)>::from_values(vec![
Value::I64(1),
Value::String("two".into()),
Value::F64(3.5),
])
.unwrap();
assert_eq!(a.as_i64(), Some(1));
assert_eq!(b.as_str(), Some("two"));
assert_eq!(c.as_f64(), Some(3.5));
}
#[test]
fn test_value_helpers() {
assert_eq!(value_as_i64(&Value::I64(42)), Some(42));
assert_eq!(value_as_i64(&Value::String("42".into())), Some(42));
assert_eq!(value_as_f64(&Value::F64(3.5)), Some(3.5));
assert_eq!(
value_as_string(&Value::String("hi".into())),
Some("hi".into())
);
assert_eq!(value_as_bool(&Value::Bool(true)), Some(true));
}
#[test]
fn test_nullable_helpers() {
assert_eq!(value_as_nullable_i64(&Value::Null), None);
assert_eq!(value_as_nullable_i64(&Value::I64(42)), Some(42));
assert_eq!(value_as_nullable_string(&Value::Null), None);
assert_eq!(
value_as_nullable_string(&Value::String("hi".into())),
Some("hi".into())
);
}
#[test]
fn test_full_flow_queryable() {
let values = vec![Value::I64(1), Value::String("Alice".into())];
let row = UserRow::from_values(values).unwrap();
assert_eq!(
row,
UserRow {
id: 1,
name: "Alice".into()
}
);
}
#[test]
fn test_full_flow_from_row_with_extra_data() {
let mut map = HashMap::new();
map.insert("id".into(), Value::I64(7));
map.insert("name".into(), Value::String("Bob".into()));
map.insert("email".into(), Value::String("bob@example.com".into()));
map.insert("created_at".into(), Value::String("2026-01-01".into()));
let row = UserRow::from_row(map).unwrap();
assert_eq!(row.id, 7);
assert_eq!(row.name, "Bob");
}
#[test]
fn test_query_new_and_sql() {
let q = Query::new("SELECT id, name FROM users");
assert_eq!(q.sql(), "SELECT id, name FROM users");
}
#[test]
fn test_query_from_str() {
let q = Query::new("SELECT 1");
assert_eq!(q.sql(), "SELECT 1");
}
#[test]
fn test_query_with_params() {
let q = Query::new("SELECT * FROM users WHERE id = ? AND name = ?");
assert!(q.sql().contains("WHERE id = ?"));
}
#[test]
fn test_validate_columns_match() {
let actual = vec!["id", "name"];
let expected = vec!["id", "name"];
assert!(validate_columns(&actual, &expected).is_ok());
}
#[test]
fn test_validate_columns_missing() {
let actual = vec!["id"];
let expected = vec!["id", "name"];
let result = validate_columns(&actual, &expected);
assert!(result.is_err());
let err = result.unwrap_err().to_string();
assert!(err.contains("missing"));
}
#[test]
fn test_validate_columns_extra() {
let actual = vec!["id", "name", "age"];
let expected = vec!["id", "name"];
let result = validate_columns(&actual, &expected);
assert!(result.is_err());
let err = result.unwrap_err().to_string();
assert!(err.contains("column mismatch"));
}
#[test]
fn test_query_as_new_and_sql() {
let q = QueryAs::<UserRow>::new("SELECT id, name FROM users");
assert_eq!(q.sql(), "SELECT id, name FROM users");
}
#[test]
fn test_query_as_fetch_all_empty() {
let _q = QueryAs::<UserRow>::new("SELECT 1");
}
#[test]
fn test_query_as_fetch_optional_empty_result() {
let _q = QueryAs::<(i64, String)>::new("SELECT id, name FROM t");
}
#[test]
fn test_validate_column_types_compatible() {
let mut row = HashMap::new();
row.insert("id".to_string(), Value::I64(42));
row.insert("name".to_string(), Value::String("Alice".into()));
let expected = vec![("id", "BIGINT"), ("name", "TEXT")];
assert!(validate_column_types(&row, &expected).is_ok());
}
#[test]
fn test_validate_column_types_case_insensitive() {
let mut row = HashMap::new();
row.insert("id".to_string(), Value::I64(1));
let expected = vec![("id", "bigint")];
assert!(validate_column_types(&row, &expected).is_ok());
}
#[test]
fn test_validate_column_types_null_skipped() {
let mut row = HashMap::new();
row.insert("id".to_string(), Value::Null);
let expected = vec![("id", "TEXT")]; assert!(validate_column_types(&row, &expected).is_ok());
}
#[test]
fn test_validate_column_types_missing_column_skipped() {
let row = HashMap::new();
let expected = vec![("missing", "BIGINT")];
assert!(validate_column_types(&row, &expected).is_ok());
}
#[test]
fn test_validate_column_types_type_mismatch() {
let mut row = HashMap::new();
row.insert("id".to_string(), Value::I64(42));
let expected = vec![("id", "TEXT")];
let result = validate_column_types(&row, &expected);
assert!(result.is_err());
let err = result.unwrap_err().to_string();
assert!(err.contains("TYPE mismatch"));
assert!(err.contains("id"));
}
#[test]
fn test_validate_column_types_bool_vs_int_mismatch() {
let mut row = HashMap::new();
row.insert("flag".to_string(), Value::Bool(true));
let expected = vec![("flag", "INT")];
let result = validate_column_types(&row, &expected);
assert!(result.is_err());
}
#[test]
fn test_validate_column_types_unknown_actual_skipped() {
let mut row = HashMap::new();
row.insert("data".to_string(), Value::String("x".into()));
let expected = vec![("data", "UNKNOWN_CUSTOM_TYPE")];
assert!(validate_column_types(&row, &expected).is_ok());
}
#[test]
fn test_value_to_sql_type_name_coverage() {
use crate::value::Value;
assert_eq!(super::value_to_sql_type_name(&Value::Bool(true)), "BOOLEAN");
assert_eq!(super::value_to_sql_type_name(&Value::I64(1)), "BIGINT");
assert_eq!(super::value_to_sql_type_name(&Value::F64(1.0)), "DOUBLE");
assert_eq!(
super::value_to_sql_type_name(&Value::Decimal("1.0".into())),
"DECIMAL"
);
assert_eq!(
super::value_to_sql_type_name(&Value::Bytes(vec![1, 2])),
"BLOB"
);
assert_eq!(super::value_to_sql_type_name(&Value::Null), "NULL");
}
}