use anyhow::{self, Context};
use mysql::prelude::Queryable;
use rhai::plugin::{
mem, Dynamic, FnAccess, FnNamespace, ImmutableString, Module, NativeCallContext,
PluginFunction, RhaiResult, TypeId,
};
#[derive(Debug, serde::Deserialize)]
struct MySQLDatabaseParameters {
pub url: String,
#[serde(default = "default_timeout", with = "humantime_serde")]
pub timeout: std::time::Duration,
#[serde(default = "default_connections")]
pub connections: rhai::INT,
}
const fn default_connections() -> rhai::INT {
4
}
const fn default_timeout() -> std::time::Duration {
std::time::Duration::from_secs(30)
}
#[derive(Clone, Debug)]
#[allow(clippy::module_name_repetitions)]
pub struct ConnectionManager {
params: mysql::Opts,
}
impl ConnectionManager {
pub fn new(params: mysql::OptsBuilder) -> Self {
Self {
params: mysql::Opts::from(params),
}
}
}
impl r2d2::ManageConnection for ConnectionManager {
type Connection = mysql::Conn;
type Error = mysql::Error;
fn connect(&self) -> Result<mysql::Conn, mysql::Error> {
mysql::Conn::new(self.params.clone())
}
fn is_valid(&self, conn: &mut mysql::Conn) -> Result<(), mysql::Error> {
mysql::prelude::Queryable::query(conn, "SELECT version()").map(|_: Vec<String>| ())
}
fn has_broken(&self, conn: &mut mysql::Conn) -> bool {
self.is_valid(conn).is_err()
}
}
#[derive(Debug, Clone)]
pub struct MySQLConnector {
pub url: String,
pub pool: r2d2::Pool<ConnectionManager>,
}
impl MySQLConnector {
pub fn query(&self, query: &str) -> anyhow::Result<Vec<rhai::Map>> {
let result = self
.pool
.get()?
.query::<mysql::Row, _>(query)
.context("failed to execute query on sql database")?;
let mut rows = Vec::with_capacity(result.len());
for row in &result {
let mut values = rhai::Map::new();
for (index, column) in row.columns().iter().enumerate() {
values.insert(
column.name_str().into(),
row.as_ref(index)
.ok_or_else(|| {
anyhow::anyhow!("failed to convert sql row value to string")
})?
.as_sql(false)
.into(),
);
}
rows.push(values);
}
Ok(rows)
}
}
#[rhai::plugin::export_module]
pub mod mysql_api {
pub type MySQL = rhai::Shared<MySQLConnector>;
#[rhai_fn(global, return_raw)]
pub fn connect(parameters: rhai::Map) -> Result<MySQL, Box<rhai::EvalAltResult>> {
let parameters = rhai::serde::from_dynamic::<MySQLDatabaseParameters>(¶meters.into())?;
let opts = mysql::Opts::from_url(¶meters.url)
.map_err::<Box<rhai::EvalAltResult>, _>(|err| err.to_string().into())?;
let builder = mysql::OptsBuilder::from_opts(opts);
let manager = ConnectionManager::new(builder);
Ok(rhai::Shared::new(MySQLConnector {
url: parameters.url,
pool: r2d2::Pool::builder()
.max_size(
u32::try_from(parameters.connections)
.map_err::<Box<rhai::EvalAltResult>, _>(|err| err.to_string().into())?,
)
.connection_timeout(parameters.timeout)
.build(manager)
.map_err::<Box<rhai::EvalAltResult>, _>(|err| err.to_string().into())?,
}))
}
#[rhai_fn(global, name = "query", return_raw, pure)]
pub fn query_str(
database: &mut MySQL,
query: &str,
) -> Result<rhai::Array, Box<rhai::EvalAltResult>> {
super::query(database, query)
}
#[rhai_fn(global, name = "query", return_raw, pure)]
#[allow(clippy::needless_pass_by_value)]
pub fn query_obj(
database: &mut MySQL,
query: vsmtp_rule_engine::api::SharedObject,
) -> Result<rhai::Array, Box<rhai::EvalAltResult>> {
super::query(database, &query.to_string())
}
}
fn query(
database: &mysql_api::MySQL,
query: &str,
) -> Result<rhai::Array, Box<rhai::EvalAltResult>> {
database.query(query).map_or_else(
|_| Ok(rhai::Array::default()),
|record| Ok(record.into_iter().map(rhai::Dynamic::from).collect()),
)
}