use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use std::fmt::{Debug, Display};
#[allow(unused_imports)]
use std::sync::Arc;
use thiserror::Error;
#[derive(Debug, Error)]
pub enum DatabaseError {
#[error("connection error: {0}")]
Connection(String),
#[error("query error: {0}")]
Query(String),
#[error("transaction error: {0}")]
Transaction(String),
#[error("pool error: {0}")]
Pool(String),
#[error("configuration error: {0}")]
Configuration(String),
#[error("migration error: {0}")]
Migration(String),
#[error("serialization error: {0}")]
Serialization(String),
}
pub type DatabaseResult<T> = Result<T, DatabaseError>;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum DatabaseType {
Postgres,
MySql,
Sqlite,
}
impl Display for DatabaseType {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
DatabaseType::Postgres => write!(f, "postgres"),
DatabaseType::MySql => write!(f, "mysql"),
DatabaseType::Sqlite => write!(f, "sqlite"),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DatabaseConfig {
pub db_type: DatabaseType,
#[serde(default)]
pub host: Option<String>,
#[serde(default)]
pub port: Option<u16>,
pub database: String,
#[serde(default)]
pub username: Option<String>,
#[serde(default)]
pub password: Option<String>,
#[serde(default)]
pub ssl_mode: Option<String>,
#[serde(default)]
pub pool: PoolConfig,
#[serde(default)]
pub extra_params: std::collections::HashMap<String, String>,
}
impl Default for DatabaseConfig {
fn default() -> Self {
Self {
db_type: DatabaseType::Sqlite,
host: None,
port: None,
database: ":memory:".to_string(),
username: None,
password: None,
ssl_mode: None,
pool: PoolConfig::default(),
extra_params: std::collections::HashMap::new(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PoolConfig {
#[serde(default = "default_max_connections")]
pub max_connections: u32,
#[serde(default = "default_min_connections")]
pub min_connections: u32,
#[serde(default = "default_idle_timeout")]
pub idle_timeout_seconds: u64,
#[serde(default = "default_max_lifetime")]
pub max_lifetime_seconds: u64,
#[serde(default = "default_acquire_timeout")]
pub acquire_timeout_seconds: u64,
}
fn default_max_connections() -> u32 {
10
}
fn default_min_connections() -> u32 {
2
}
fn default_idle_timeout() -> u64 {
300
} fn default_max_lifetime() -> u64 {
1800
} fn default_acquire_timeout() -> u64 {
30
}
impl Default for PoolConfig {
fn default() -> Self {
PoolConfig {
max_connections: default_max_connections(),
min_connections: default_min_connections(),
idle_timeout_seconds: default_idle_timeout(),
max_lifetime_seconds: default_max_lifetime(),
acquire_timeout_seconds: default_acquire_timeout(),
}
}
}
pub trait DatabaseRow: Send + Sync {
fn get_string(&self, column: &str) -> DatabaseResult<String>;
fn get_i64(&self, column: &str) -> DatabaseResult<i64>;
fn get_f64(&self, column: &str) -> DatabaseResult<f64>;
fn get_bool(&self, column: &str) -> DatabaseResult<bool>;
fn get_bytes(&self, column: &str) -> DatabaseResult<Vec<u8>>;
fn try_get_string(&self, column: &str) -> DatabaseResult<Option<String>>;
fn try_get_i64(&self, column: &str) -> DatabaseResult<Option<i64>>;
fn try_get_f64(&self, column: &str) -> DatabaseResult<Option<f64>>;
fn try_get_bool(&self, column: &str) -> DatabaseResult<Option<bool>>;
fn try_get_bytes(&self, column: &str) -> DatabaseResult<Option<Vec<u8>>>;
}
#[async_trait]
pub trait DatabaseConnection: Send + Sync {
async fn execute(&self, query: &str, params: &[DatabaseValue]) -> DatabaseResult<u64>;
async fn query(
&self,
query: &str,
params: &[DatabaseValue],
) -> DatabaseResult<Vec<Box<dyn DatabaseRow>>>;
async fn query_one(
&self,
query: &str,
params: &[DatabaseValue],
) -> DatabaseResult<Option<Box<dyn DatabaseRow>>>;
async fn begin_transaction(&self) -> DatabaseResult<Box<dyn DatabaseTransaction>>;
fn get_database_type(&self) -> DatabaseType;
async fn ping(&self) -> DatabaseResult<()>;
async fn close(&self) -> DatabaseResult<()>;
}
#[async_trait]
pub trait DatabaseTransaction: Send + Sync {
async fn execute(&mut self, query: &str, params: &[DatabaseValue]) -> DatabaseResult<u64>;
async fn query(
&mut self,
query: &str,
params: &[DatabaseValue],
) -> DatabaseResult<Vec<Box<dyn DatabaseRow>>>;
async fn query_one(
&mut self,
query: &str,
params: &[DatabaseValue],
) -> DatabaseResult<Option<Box<dyn DatabaseRow>>>;
async fn commit(self: Box<Self>) -> DatabaseResult<()>;
async fn rollback(self: Box<Self>) -> DatabaseResult<()>;
}
#[derive(Debug, Clone)]
pub enum DatabaseValue {
Null,
Boolean(bool),
Integer(i64),
Float(f64),
Text(String),
Blob(Vec<u8>),
Array(Vec<DatabaseValue>),
}
#[cfg(all(feature = "integration_tests", feature = "database"))]
mod mock {
use super::*;
pub struct MockTransaction {
pub connection: Arc<dyn DatabaseConnection>,
}
#[async_trait]
impl DatabaseTransaction for MockTransaction {
async fn execute(&mut self, query: &str, params: &[DatabaseValue]) -> DatabaseResult<u64> {
self.connection.execute(query, params).await
}
async fn query(
&mut self,
query: &str,
params: &[DatabaseValue],
) -> DatabaseResult<Vec<Box<dyn DatabaseRow>>> {
self.connection.query(query, params).await
}
async fn query_one(
&mut self,
query: &str,
params: &[DatabaseValue],
) -> DatabaseResult<Option<Box<dyn DatabaseRow>>> {
self.connection.query_one(query, params).await
}
async fn commit(self: Box<Self>) -> DatabaseResult<()> {
Ok(())
}
async fn rollback(self: Box<Self>) -> DatabaseResult<()> {
Ok(())
}
}
pub struct MockConnection {
inner: Arc<dyn DatabaseConnection>,
}
#[async_trait]
impl DatabaseConnection for MockConnection {
async fn execute(&self, query: &str, params: &[DatabaseValue]) -> DatabaseResult<u64> {
self.inner.execute(query, params).await
}
async fn query(
&self,
query: &str,
params: &[DatabaseValue],
) -> DatabaseResult<Vec<Box<dyn DatabaseRow>>> {
self.inner.query(query, params).await
}
async fn query_one(
&self,
query: &str,
params: &[DatabaseValue],
) -> DatabaseResult<Option<Box<dyn DatabaseRow>>> {
self.inner.query_one(query, params).await
}
async fn begin_transaction(&self) -> DatabaseResult<Box<dyn DatabaseTransaction>> {
Ok(Box::new(MockTransaction {
connection: Arc::clone(&self.inner),
}))
}
fn get_database_type(&self) -> DatabaseType {
self.inner.get_database_type()
}
async fn ping(&self) -> DatabaseResult<()> {
self.inner.ping().await
}
async fn close(&self) -> DatabaseResult<()> {
self.inner.close().await
}
}
pub fn create_mock_db_connection(
conn: Box<dyn DatabaseConnection>,
) -> Box<dyn DatabaseConnection> {
Box::new(MockConnection { inner: Arc::from(conn) })
}
}
pub async fn create_database_connection(
config: &DatabaseConfig,
) -> DatabaseResult<Box<dyn DatabaseConnection>> {
match config.db_type {
DatabaseType::Postgres => {
#[cfg(feature = "postgres")]
{
let conn = postgres::PostgresConnection::connect(config).await?;
let boxed_conn = Box::new(conn) as Box<dyn DatabaseConnection>;
#[cfg(all(feature = "integration_tests", feature = "database"))]
return Ok(mock::create_mock_db_connection(boxed_conn));
#[cfg(not(all(feature = "integration_tests", feature = "database")))]
return Ok(boxed_conn);
}
#[cfg(not(feature = "postgres"))]
{
Err(DatabaseError::Configuration(
"PostgreSQL support is not enabled. Enable the 'postgres' feature.".to_string(),
))
}
}
DatabaseType::MySql => {
#[cfg(feature = "mysql")]
{
let conn = mysql::MySqlConnection::connect(config).await?;
let boxed_conn = Box::new(conn) as Box<dyn DatabaseConnection>;
#[cfg(all(feature = "integration_tests", feature = "database"))]
return Ok(mock::create_mock_db_connection(boxed_conn));
#[cfg(not(all(feature = "integration_tests", feature = "database")))]
return Ok(boxed_conn);
}
#[cfg(not(feature = "mysql"))]
{
Err(DatabaseError::Configuration(
"MySQL support is not enabled. Enable the 'mysql' feature.".to_string(),
))
}
}
DatabaseType::Sqlite => {
#[cfg(feature = "sqlite")]
{
let conn = sqlite::SqliteConnection::connect(config).await?;
let boxed_conn = Box::new(conn) as Box<dyn DatabaseConnection>;
#[cfg(all(feature = "integration_tests", feature = "database"))]
return Ok(mock::create_mock_db_connection(boxed_conn));
#[cfg(not(all(feature = "integration_tests", feature = "database")))]
return Ok(boxed_conn);
}
#[cfg(not(feature = "sqlite"))]
{
Err(DatabaseError::Configuration(
"SQLite support is not enabled. Enable the 'sqlite' feature.".to_string(),
))
}
}
}
}
fn is_running_as_module() -> bool {
std::env::var("PYWATT_MODULE_ID").is_ok()
}
#[cfg(feature = "postgres")]
mod postgres;
#[cfg(feature = "postgres")]
pub use postgres::PostgresConnection;
#[cfg(feature = "mysql")]
mod mysql;
#[cfg(feature = "mysql")]
pub use mysql::MySqlConnection;
#[cfg(feature = "sqlite")]
mod sqlite;
#[cfg(feature = "sqlite")]
pub use sqlite::SqliteConnection;
#[cfg(feature = "ipc")]
pub mod proxy_connection;
#[cfg(feature = "ipc")]
pub use proxy_connection::ProxyDatabaseConnection;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_database_type_display() {
assert_eq!(DatabaseType::Postgres.to_string(), "postgres");
assert_eq!(DatabaseType::MySql.to_string(), "mysql");
assert_eq!(DatabaseType::Sqlite.to_string(), "sqlite");
}
#[test]
fn test_pool_config_defaults() {
let config = PoolConfig::default();
assert_eq!(config.max_connections, 10);
assert_eq!(config.min_connections, 2);
assert_eq!(config.idle_timeout_seconds, 300);
assert_eq!(config.max_lifetime_seconds, 1800);
assert_eq!(config.acquire_timeout_seconds, 30);
}
}
pub mod extensions {
use super::{DatabaseConfig, DatabaseType};
pub fn postgres_config(
host: impl Into<String>,
port: u16,
database: impl Into<String>,
username: impl Into<String>,
password: impl Into<String>,
) -> DatabaseConfig {
let mut config = DatabaseConfig {
db_type: DatabaseType::Postgres,
host: Some(host.into()),
port: Some(port),
database: database.into(),
username: Some(username.into()),
password: Some(password.into()),
ssl_mode: Some("prefer".to_string()),
..Default::default()
};
config
.extra_params
.insert("application_name".to_string(), "pywatt_sdk".to_string());
config
}
pub fn mysql_config(
host: impl Into<String>,
port: u16,
database: impl Into<String>,
username: impl Into<String>,
password: impl Into<String>,
) -> DatabaseConfig {
let mut config = DatabaseConfig {
db_type: DatabaseType::MySql,
host: Some(host.into()),
port: Some(port),
database: database.into(),
username: Some(username.into()),
password: Some(password.into()),
..Default::default()
};
config
.extra_params
.insert("charset".to_string(), "utf8mb4".to_string());
config
.extra_params
.insert("collation".to_string(), "utf8mb4_unicode_ci".to_string());
config
}
pub fn sqlite_config(database_path: impl Into<String>) -> DatabaseConfig {
DatabaseConfig {
db_type: DatabaseType::Sqlite,
database: database_path.into(),
..Default::default()
}
}
}