use crate::config::{DatabaseConfig, DatabaseType};
use crate::error::{Error, Result};
use sqlx::any::{AnyConnectOptions, AnyPoolOptions};
use sqlx::{Any, Pool};
use std::str::FromStr;
use std::time::Duration;
pub type DatabasePool = Pool<Any>;
#[derive(Debug, Clone)]
pub struct Database {
pool: DatabasePool,
config: DatabaseConfig,
}
impl Database {
pub async fn from_config(config: DatabaseConfig) -> Result<Self> {
if !config.enabled {
return Err(Error::Database("数据库未启用".to_string()));
}
let connect_options = AnyConnectOptions::from_str(&config.url)
.map_err(|e| Error::Database(format!("无效的数据库连接 URL: {}", e)))?;
let pool = AnyPoolOptions::new()
.max_connections(config.max_connections)
.min_connections(config.min_connections)
.acquire_timeout(Duration::from_secs(config.connect_timeout))
.connect_with(connect_options)
.await
.map_err(|e| Error::Database(format!("数据库连接失败: {}", e)))?;
tracing::info!(
"数据库连接成功: {:?}, 最大连接数: {}, 最小连接数: {}",
config.db_type,
config.max_connections,
config.min_connections
);
Ok(Self { pool, config })
}
pub fn pool(&self) -> &DatabasePool {
&self.pool
}
pub fn config(&self) -> &DatabaseConfig {
&self.config
}
pub async fn ping(&self) -> Result<()> {
sqlx::query("SELECT 1")
.execute(&self.pool)
.await
.map_err(|e| Error::Database(format!("数据库连接测试失败: {}", e)))?;
Ok(())
}
pub async fn close(&self) {
self.pool.close().await;
tracing::info!("数据库连接已关闭");
}
pub fn db_type(&self) -> &DatabaseType {
&self.config.db_type
}
pub async fn execute_raw(&self, sql: &str) -> Result<()> {
sqlx::query(sql)
.execute(&self.pool)
.await
.map_err(|e| Error::Database(format!("SQL 执行失败: {}", e)))?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
#[ignore] async fn test_database_connection() {
let config = DatabaseConfig {
enabled: true,
db_type: DatabaseType::Sqlite,
url: "sqlite::memory:".to_string(),
max_connections: 5,
min_connections: 1,
connect_timeout: 30,
auto_migrate: false,
};
let db = Database::from_config(config).await;
assert!(db.is_ok());
if let Ok(db) = db {
assert!(db.ping().await.is_ok());
}
}
}