pub mod base64;
pub mod builder;
pub mod config;
pub mod connections;
pub mod credentials;
pub mod dialect;
pub mod driver;
pub mod migration;
pub mod model;
pub mod mysql;
pub mod pagination;
pub mod pool;
pub mod postgres;
pub mod random;
pub mod row;
pub mod schema;
pub mod sqlserver;
pub mod tls;
pub mod value;
pub use builder::{Direction, QueryBuilder};
pub use config::DatabaseConfig;
pub use dialect::{ColumnType, Dialect, ReturningStyle};
pub use driver::{Driver, DriverConnection, QueryResult};
pub use connections::{Connections, DEFAULT_BUDGET};
pub use migration::{Faker, Migration, Migrator, Seeder};
pub use model::{Model, ModelExt, belongs_to, has_many};
pub use mysql::{MySqlConnection, MySqlDriver};
pub use pagination::{CursorPage, Page};
pub use postgres::connection::{log_bindings, set_log_bindings};
pub use pool::{Pool, PooledConnection};
pub use schema::{Schema, Table};
pub use sqlserver::{SqlServerConnection, SqlServerDriver};
pub use row::{Row, rows_to_json};
pub use value::{FromValue, Value};
pub use rustlavel_core::{Error, Result};
use std::sync::Arc;
fn driver_for(config: DatabaseConfig) -> Result<Arc<dyn Driver>> {
match config.driver.as_str() {
"postgres" => Ok(Arc::new(postgres::PostgresDriver::new(config))),
"mysql" => Ok(Arc::new(mysql::MySqlDriver::new(config))),
"sqlserver" => Ok(Arc::new(sqlserver::SqlServerDriver::new(config))),
other => Err(Error::msg(format!(
"the `{other}` driver is not available in this build. \
Point DATABASE_URL at a database this build supports."
))),
}
}
pub use rustlavel_macros::Model;
pub mod prelude {
pub use crate::connections::Connections;
pub use crate::migration::{Faker, Migrator, Seeder};
pub use crate::model::{ModelExt, belongs_to, has_many};
pub use crate::schema::{Schema, Table};
pub use crate::{CursorPage, Database, Model, Page, QueryBuilder, Row, Value};
pub use rustlavel_core::{Error, Json, Result};
}
#[derive(Clone)]
pub struct Database {
pool: Pool,
dialect: Arc<dyn Dialect>,
}
impl Database {
pub async fn connect(url: &str) -> Result<Database> {
Database::with_config(DatabaseConfig::from_url(url)?).await
}
pub async fn with_config(config: DatabaseConfig) -> Result<Database> {
let database = Database::lazy(config)?;
database.pool.verify().await?;
Ok(database)
}
pub fn lazy(config: DatabaseConfig) -> Result<Database> {
Ok(Database::with_driver(driver_for(config)?))
}
pub fn with_driver(driver: Arc<dyn Driver>) -> Database {
let dialect = driver.dialect();
Database { pool: Pool::new(driver), dialect }
}
pub fn dialect(&self) -> &dyn Dialect {
self.dialect.as_ref()
}
pub fn pool(&self) -> &Pool {
&self.pool
}
pub fn table(&self, name: &str) -> QueryBuilder {
QueryBuilder::new(name)
}
pub async fn select(&self, sql: &str, params: &[Value]) -> Result<Vec<Row>> {
let mut connection = self.pool.acquire().await?;
Ok(connection.query(sql, params).await?.rows)
}
pub async fn select_one(&self, sql: &str, params: &[Value]) -> Result<Option<Row>> {
Ok(self.select(sql, params).await?.into_iter().next())
}
pub async fn execute(&self, sql: &str, params: &[Value]) -> Result<u64> {
let mut connection = self.pool.acquire().await?;
Ok(connection.query(sql, params).await?.affected)
}
pub async fn run(&self, sql: &str) -> Result<u64> {
let mut connection = self.pool.acquire().await?;
Ok(connection.simple_query(sql).await?.affected)
}
pub async fn insert_returning_key(
&self,
sql: &str,
params: &[Value],
column: &str,
) -> Result<Option<Value>> {
let mut connection = self.pool.acquire().await?;
let result = connection.query(sql, params).await?;
if let Some(row) = result.rows.first() {
return Ok(Some(row.value(column).or_else(|_| row.value_at(0))?.clone()));
}
Ok(result.last_insert_id.map(Value::Int))
}
pub async fn scalar<T: FromValue>(&self, sql: &str, params: &[Value]) -> Result<Option<T>> {
match self.select_one(sql, params).await? {
Some(row) => row.get_at::<T>(0).map(Some),
None => Ok(None),
}
}
pub async fn begin(&self) -> Result<Transaction> {
let mut connection = self.pool.acquire().await?;
connection.simple_query(self.dialect.begin_sql()).await?;
Ok(Transaction {
connection: Some(connection),
dialect: Arc::clone(&self.dialect),
finished: false,
})
}
pub async fn close(&self) {
self.pool.close().await;
}
}
pub struct Transaction {
connection: Option<PooledConnection>,
dialect: Arc<dyn Dialect>,
finished: bool,
}
impl Transaction {
fn connection(&mut self) -> Result<&mut PooledConnection> {
self.connection
.as_mut()
.ok_or_else(|| Error::msg("this transaction has already finished"))
}
pub fn dialect(&self) -> &dyn Dialect {
self.dialect.as_ref()
}
pub async fn select(&mut self, sql: &str, params: &[Value]) -> Result<Vec<Row>> {
Ok(self.connection()?.query(sql, params).await?.rows)
}
pub async fn select_one(&mut self, sql: &str, params: &[Value]) -> Result<Option<Row>> {
Ok(self.select(sql, params).await?.into_iter().next())
}
pub async fn execute(&mut self, sql: &str, params: &[Value]) -> Result<u64> {
Ok(self.connection()?.query(sql, params).await?.affected)
}
pub async fn run(&mut self, sql: &str) -> Result<u64> {
Ok(self.connection()?.simple_query(sql).await?.affected)
}
pub async fn scalar<T: FromValue>(&mut self, sql: &str, params: &[Value]) -> Result<Option<T>> {
match self.select_one(sql, params).await? {
Some(row) => row.get_at::<T>(0).map(Some),
None => Ok(None),
}
}
pub async fn savepoint(&mut self, name: &str) -> Result<()> {
validate_identifier(name)?;
let sql = self.dialect.savepoint_sql(name);
self.connection()?.simple_query(&sql).await?;
Ok(())
}
pub async fn rollback_to(&mut self, name: &str) -> Result<()> {
validate_identifier(name)?;
let sql = self.dialect.rollback_to_savepoint_sql(name);
self.connection()?.simple_query(&sql).await?;
Ok(())
}
pub async fn commit(mut self) -> Result<()> {
let sql = self.dialect.commit_sql();
self.connection()?.simple_query(sql).await?;
self.finished = true;
Ok(())
}
pub async fn rollback(mut self) -> Result<()> {
let sql = self.dialect.rollback_sql();
self.connection()?.simple_query(sql).await?;
self.finished = true;
Ok(())
}
}
impl Drop for Transaction {
fn drop(&mut self) {
if self.finished {
return;
}
if let Some(mut connection) = self.connection.take() {
let sql = self.dialect.rollback_sql();
tokio::spawn(async move {
let _ = connection.simple_query(sql).await;
});
}
}
}
pub fn validate_identifier(name: &str) -> Result<()> {
let valid = !name.is_empty()
&& name.len() <= 63
&& name.chars().next().is_some_and(|c| c.is_ascii_alphabetic() || c == '_')
&& name.chars().all(|c| c.is_ascii_alphanumeric() || c == '_');
if valid {
Ok(())
} else {
Err(Error::msg(format!(
"`{name}` is not a valid SQL identifier. Identifiers may contain letters, digits and \
underscores, and must not start with a digit."
)))
}
}
pub fn quote_identifier(name: &str) -> Result<String> {
let quoted: Result<Vec<String>> = name
.split('.')
.map(|part| validate_identifier(part).map(|_| format!("\"{part}\"")))
.collect();
Ok(quoted?.join("."))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn accepts_ordinary_identifiers() {
for name in ["users", "user_profiles", "_private", "t1"] {
assert!(validate_identifier(name).is_ok(), "{name} should be valid");
}
}
#[test]
fn rejects_anything_that_could_alter_a_statement() {
for name in ["users; drop table users", "user\"s", "1abc", "", "a b", "users--"] {
assert!(validate_identifier(name).is_err(), "{name:?} should be rejected");
}
}
#[test]
fn quotes_qualified_names_part_by_part() {
assert_eq!(quote_identifier("public.users").unwrap(), "\"public\".\"users\"");
assert!(quote_identifier("public.users; drop table x").is_err());
}
}