use crate::error::{SageError, SageResult};
use crate::mock::{try_get_mock, MockResponse};
#[cfg(feature = "database")]
use sqlx::{any::AnyRow, AnyPool, Column, Row};
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct DbRow {
pub columns: Vec<String>,
pub values: Vec<String>,
}
#[derive(Debug, Clone)]
pub struct DatabaseClient {
#[cfg(feature = "database")]
pool: AnyPool,
#[cfg(not(feature = "database"))]
_marker: std::marker::PhantomData<()>,
}
impl DatabaseClient {
#[cfg(feature = "database")]
pub async fn connect(url: &str) -> SageResult<Self> {
sqlx::any::install_default_drivers();
let pool = AnyPool::connect(url)
.await
.map_err(|e| SageError::Tool(format!("Database connection failed: {e}")))?;
Ok(Self { pool })
}
#[cfg(not(feature = "database"))]
pub async fn connect(_url: &str) -> SageResult<Self> {
Err(SageError::Tool(
"Database support not enabled. Compile with the 'database' feature.".to_string(),
))
}
#[cfg(feature = "database")]
pub async fn from_env() -> SageResult<Self> {
let url = std::env::var("SAGE_DATABASE_URL").map_err(|_| {
SageError::Tool("SAGE_DATABASE_URL environment variable not set".to_string())
})?;
Self::connect(&url).await
}
#[cfg(not(feature = "database"))]
pub async fn from_env() -> SageResult<Self> {
Err(SageError::Tool(
"Database support not enabled. Compile with the 'database' feature.".to_string(),
))
}
#[cfg(feature = "database")]
pub async fn query(&self, sql: String) -> SageResult<Vec<DbRow>> {
if let Some(mock_response) = try_get_mock("Database", "query") {
return Self::apply_mock_vec(mock_response);
}
let rows: Vec<AnyRow> = sqlx::query(&sql)
.fetch_all(&self.pool)
.await
.map_err(|e| SageError::Tool(format!("Query failed: {e}")))?;
let result: Vec<DbRow> = rows
.iter()
.map(|row| {
let columns: Vec<String> =
row.columns().iter().map(|c| c.name().to_string()).collect();
let values: Vec<String> = (0..row.columns().len())
.map(|i| {
if let Ok(v) = row.try_get::<String, _>(i) {
v
} else if let Ok(v) = row.try_get::<i64, _>(i) {
v.to_string()
} else if let Ok(v) = row.try_get::<i32, _>(i) {
v.to_string()
} else if let Ok(v) = row.try_get::<f64, _>(i) {
v.to_string()
} else if let Ok(v) = row.try_get::<bool, _>(i) {
v.to_string()
} else {
row.try_get::<Option<String>, _>(i)
.ok()
.flatten()
.unwrap_or_else(|| "null".to_string())
}
})
.collect();
DbRow { columns, values }
})
.collect();
Ok(result)
}
#[cfg(not(feature = "database"))]
pub async fn query(&self, _sql: String) -> SageResult<Vec<DbRow>> {
if let Some(mock_response) = try_get_mock("Database", "query") {
return Self::apply_mock_vec(mock_response);
}
Err(SageError::Tool(
"Database support not enabled. Compile with the 'database' feature.".to_string(),
))
}
#[cfg(feature = "database")]
pub async fn execute(&self, sql: String) -> SageResult<i64> {
if let Some(mock_response) = try_get_mock("Database", "execute") {
return Self::apply_mock_i64(mock_response);
}
let result = sqlx::query(&sql)
.execute(&self.pool)
.await
.map_err(|e| SageError::Tool(format!("Execute failed: {e}")))?;
Ok(result.rows_affected() as i64)
}
#[cfg(not(feature = "database"))]
pub async fn execute(&self, _sql: String) -> SageResult<i64> {
if let Some(mock_response) = try_get_mock("Database", "execute") {
return Self::apply_mock_i64(mock_response);
}
Err(SageError::Tool(
"Database support not enabled. Compile with the 'database' feature.".to_string(),
))
}
fn apply_mock_vec(mock_response: MockResponse) -> SageResult<Vec<DbRow>> {
match mock_response {
MockResponse::Value(v) => serde_json::from_value(v)
.map_err(|e| SageError::Tool(format!("mock deserialize: {e}"))),
MockResponse::Fail(msg) => Err(SageError::Tool(msg)),
}
}
fn apply_mock_i64(mock_response: MockResponse) -> SageResult<i64> {
match mock_response {
MockResponse::Value(v) => serde_json::from_value(v)
.map_err(|e| SageError::Tool(format!("mock deserialize: {e}"))),
MockResponse::Fail(msg) => Err(SageError::Tool(msg)),
}
}
}
#[cfg(all(test, feature = "database"))]
mod tests {
use super::*;
#[tokio::test]
async fn database_connect_sqlite() {
let client = DatabaseClient::connect("sqlite:file::memory:?mode=memory&cache=shared")
.await
.unwrap();
drop(client);
}
#[tokio::test]
async fn database_execute_and_query() {
let temp_dir = tempfile::tempdir().unwrap();
let db_path = temp_dir.path().join("test.db");
std::fs::write(&db_path, "").unwrap();
let url = format!("sqlite:{}?mode=rwc", db_path.display());
let client = DatabaseClient::connect(&url).await.unwrap();
client
.execute("CREATE TABLE test (id INTEGER PRIMARY KEY, name TEXT)".to_string())
.await
.unwrap();
let affected = client
.execute("INSERT INTO test (id, name) VALUES (1, 'Alice'), (2, 'Bob')".to_string())
.await
.unwrap();
assert_eq!(affected, 2);
let rows = client
.query("SELECT id, name FROM test ORDER BY id".to_string())
.await
.unwrap();
assert_eq!(rows.len(), 2);
assert_eq!(rows[0].columns, vec!["id", "name"]);
assert_eq!(rows[0].values, vec!["1", "Alice"]);
assert_eq!(rows[1].values, vec!["2", "Bob"]);
}
#[tokio::test]
async fn database_query_select_one() {
let client = DatabaseClient::connect("sqlite:file::memory:?mode=memory&cache=shared")
.await
.unwrap();
let rows = client.query("SELECT 1 as value".to_string()).await.unwrap();
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].columns, vec!["value"]);
assert_eq!(rows[0].values, vec!["1"]);
}
}