use serde_json::{Value, json};
use sha2::{Digest, Sha256};
use tokio::spawn;
use tokio::sync::Mutex;
use tokio::task::JoinHandle;
use tokio_postgres::types::Type;
use tokio_postgres::{Client, NoTls, Row};
use tracing::warn;
use super::{PollingProbe, ProbeError, ProbeFuture, ProbeResult};
#[derive(Debug, Clone)]
pub struct SqlProbeConfig {
pub connection_string: String,
pub query: String,
}
struct SqlConnection {
client: Client,
handle: JoinHandle<()>,
}
pub struct SqlProbe {
config: SqlProbeConfig,
conn: Mutex<Option<SqlConnection>>,
}
impl SqlProbe {
pub fn new(config: SqlProbeConfig) -> Self {
Self {
config,
conn: Mutex::new(None),
}
}
fn column_to_json(row: &Row, index: usize, col_type: &Type) -> Value {
match *col_type {
Type::BOOL => row
.try_get::<_, Option<bool>>(index)
.ok()
.flatten()
.map_or(Value::Null, |v| json!(v)),
Type::INT2 => row
.try_get::<_, Option<i16>>(index)
.ok()
.flatten()
.map_or(Value::Null, |v| json!(v)),
Type::INT4 => row
.try_get::<_, Option<i32>>(index)
.ok()
.flatten()
.map_or(Value::Null, |v| json!(v)),
Type::INT8 => row
.try_get::<_, Option<i64>>(index)
.ok()
.flatten()
.map_or(Value::Null, |v| json!(v)),
Type::FLOAT4 => row
.try_get::<_, Option<f32>>(index)
.ok()
.flatten()
.map_or(Value::Null, |v| json!(v)),
Type::FLOAT8 => row
.try_get::<_, Option<f64>>(index)
.ok()
.flatten()
.map_or(Value::Null, |v| json!(v)),
_ => row
.try_get::<_, Option<String>>(index)
.ok()
.flatten()
.map_or(Value::Null, |v| json!(v)),
}
}
fn row_to_json(row: &Row) -> Value {
let mut map = serde_json::Map::new();
for (i, col) in row.columns().iter().enumerate() {
let value = Self::column_to_json(row, i, col.type_());
map.insert(col.name().to_string(), value);
}
Value::Object(map)
}
}
impl Drop for SqlProbe {
fn drop(&mut self) {
if let Some(conn) = self.conn.get_mut().take() {
conn.handle.abort();
}
}
}
impl PollingProbe for SqlProbe {
fn name(&self) -> &str {
"sql"
}
fn poll(&self) -> ProbeFuture<'_> {
Box::pin(async {
let mut guard = self.conn.lock().await;
let needs_reconnect = match &*guard {
Some(c) => c.handle.is_finished(),
None => true,
};
if needs_reconnect {
if let Some(old) = guard.take() {
old.handle.abort();
}
let (client, connection) =
tokio_postgres::connect(&self.config.connection_string, NoTls)
.await
.map_err(|e| ProbeError::Failed(format!("SQL connection failed: {e}")))?;
let handle = spawn(async move {
if let Err(e) = connection.await {
warn!(error = %e, "SQL connection error");
}
});
*guard = Some(SqlConnection { client, handle });
}
let conn = guard.as_ref().expect("connection just established");
let result = conn.client.query(&self.config.query, &[]).await;
if result.is_err()
&& let Some(old) = guard.take()
{
old.handle.abort();
}
let rows = result.map_err(|e| ProbeError::Failed(format!("SQL query failed: {e}")))?;
if rows.is_empty() {
return Ok(None);
}
let json_rows: Vec<Value> = rows.iter().map(Self::row_to_json).collect();
let data = json!({ "rows": json_rows, "count": rows.len() });
let hash_input = serde_json::to_vec(&data).unwrap_or_default();
let content_hash = hex::encode(Sha256::digest(&hash_input));
Ok(Some(ProbeResult::with_hash(data, content_hash)))
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn sql_probe_name() {
let probe = SqlProbe::new(SqlProbeConfig {
connection_string: "host=localhost".to_string(),
query: "SELECT 1".to_string(),
});
assert_eq!(probe.name(), "sql");
}
#[test]
fn sql_probe_config_clone() {
let config = SqlProbeConfig {
connection_string: "host=db dbname=app".to_string(),
query: "SELECT * FROM events".to_string(),
};
let cloned = config.clone();
assert_eq!(cloned.connection_string, config.connection_string);
assert_eq!(cloned.query, config.query);
}
}