use chrono::{SecondsFormat, Utc};
use laterite_core::strata::*;
use serde::de::DeserializeOwned;
use serde::Serialize;
use thiserror::Error;
#[derive(Iden)]
pub(crate) enum Settings {
Table,
Code,
Value,
UpdatedAt,
}
#[derive(Debug, Error)]
pub enum SettingsError {
#[error("database error")]
Database(#[from] sqlx::Error),
#[error("settings serialization error")]
Serde(#[from] serde_json::Error),
}
pub trait SettingsModel: Serialize + DeserializeOwned + Default {
const CODE: &'static str;
}
pub async fn load<T: SettingsModel>(db: &Db) -> Result<T, SettingsError> {
match get(db, T::CODE).await? {
Some(value) => Ok(serde_json::from_value(value)?),
None => Ok(T::default()),
}
}
pub async fn save<T: SettingsModel>(db: &Db, model: &T) -> Result<(), SettingsError> {
set(db, T::CODE, &serde_json::to_value(model)?).await
}
pub async fn get(db: &Db, code: &str) -> Result<Option<serde_json::Value>, SettingsError> {
let (sql, values) = build(
db.backend,
Query::select()
.column(Settings::Value)
.from(Settings::Table)
.and_where(Expr::col(Settings::Code).eq(code))
.to_owned(),
);
let row = bind_values(sqlx::query(&sql), values)
.fetch_optional(&db.pool)
.await?;
match row {
Some(r) => {
let text = r.get_text("value")?;
Ok(Some(serde_json::from_str(&text)?))
}
None => Ok(None),
}
}
pub async fn set(db: &Db, code: &str, value: &serde_json::Value) -> Result<(), SettingsError> {
let json = serde_json::to_string(value)?;
let now = Utc::now().to_rfc3339_opts(SecondsFormat::Micros, true);
let (sql, values) = build(
db.backend,
Query::insert()
.into_table(Settings::Table)
.columns([Settings::Code, Settings::Value, Settings::UpdatedAt])
.values_panic([code.into(), json.into(), now.into()])
.on_conflict(
OnConflict::column(Settings::Code)
.update_columns([Settings::Value, Settings::UpdatedAt])
.to_owned(),
)
.to_owned(),
);
bind_values(sqlx::query(&sql), values)
.execute(&db.pool)
.await?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use serde::Deserialize;
async fn test_db() -> (Db, laterite_core::testing::TestGuard) {
laterite_core::testing::connect_test(&[super::super::migrations::migrations()]).await
}
#[derive(Debug, Default, PartialEq, Serialize, Deserialize)]
struct LogSettings {
#[serde(default)]
log_events: bool,
#[serde(default)]
log_requests: bool,
}
impl SettingsModel for LogSettings {
const CODE: &'static str = "test.log";
}
#[tokio::test]
async fn typed_round_trip_and_default() {
let (db, _guard) = test_db().await;
assert_eq!(
load::<LogSettings>(&db).await.unwrap(),
LogSettings::default()
);
let saved = LogSettings {
log_events: true,
log_requests: false,
};
save(&db, &saved).await.unwrap();
assert_eq!(load::<LogSettings>(&db).await.unwrap(), saved);
let raw = get(&db, LogSettings::CODE).await.unwrap().unwrap();
assert_eq!(raw["log_events"], serde_json::json!(true));
}
#[tokio::test]
async fn missing_field_uses_serde_default() {
let (db, _guard) = test_db().await;
set(
&db,
LogSettings::CODE,
&serde_json::json!({ "log_events": true }),
)
.await
.unwrap();
assert_eq!(
load::<LogSettings>(&db).await.unwrap(),
LogSettings {
log_events: true,
log_requests: false,
}
);
}
}