use std::sync::Arc;
use std::time::Duration;
use saya_store::SqliteStateStore;
use crate::connection::ConnectionRegistry;
use crate::contracts::RecallReceipt;
mod chart_tool;
mod definitions;
mod dispatch;
mod fan_out;
mod observations;
mod recorder;
mod override_log;
pub(crate) use observations::{
DrainedObservations, ObservationLog, ObservationOutcome, ToolObservation,
};
pub(crate) use override_log::OverrideLog;
pub(crate) struct DatabaseTools {
pub(super) registry: ConnectionRegistry,
pub(super) max_rows: usize,
pub(super) allow_query_data: bool,
pub(super) state_db: Option<SqliteStateStore>,
pub(super) max_concurrent_fan_out_queries: usize,
pub(super) fan_out_query_timeout: Duration,
pub(super) observations: Option<Arc<ObservationLog>>,
pub(super) supplied_objects: Vec<String>,
pub(super) recall_receipt: Option<Arc<RecallReceipt>>,
pub(super) override_log: Option<Arc<OverrideLog>>,
}
impl DatabaseTools {
const MAX_CONCURRENT_FAN_OUT_QUERIES: usize = 4;
const FAN_OUT_QUERY_TIMEOUT: Duration = Duration::from_secs(30);
#[cfg(test)]
pub(crate) fn new(
connector: Option<Box<dyn saya_connectors::DatabaseConnector>>,
max_rows: usize,
allow_query_data: bool,
) -> Self {
use crate::connection::ConnectionEntry;
let mut registry = ConnectionRegistry::new("primary");
if let Some(c) = connector {
let dialect = c.dialect();
registry.insert(
"primary",
ConnectionEntry {
connector: c,
dialect,
profile_id: None,
},
);
}
Self {
registry,
max_rows,
allow_query_data,
state_db: None,
max_concurrent_fan_out_queries: Self::MAX_CONCURRENT_FAN_OUT_QUERIES,
fan_out_query_timeout: Self::FAN_OUT_QUERY_TIMEOUT,
observations: None,
supplied_objects: Vec::new(),
recall_receipt: None,
override_log: None,
}
}
#[cfg(test)]
pub(crate) fn with_registry(
registry: ConnectionRegistry,
max_rows: usize,
allow_query_data: bool,
state_db: Option<SqliteStateStore>,
) -> Self {
Self {
registry,
max_rows,
allow_query_data,
state_db,
max_concurrent_fan_out_queries: Self::MAX_CONCURRENT_FAN_OUT_QUERIES,
fan_out_query_timeout: Self::FAN_OUT_QUERY_TIMEOUT,
observations: None,
supplied_objects: Vec::new(),
recall_receipt: None,
override_log: None,
}
}
pub(crate) fn with_learning(
registry: ConnectionRegistry,
max_rows: usize,
allow_query_data: bool,
state_db: Option<SqliteStateStore>,
observations: Option<Arc<ObservationLog>>,
) -> Self {
Self {
registry,
max_rows,
allow_query_data,
state_db,
max_concurrent_fan_out_queries: Self::MAX_CONCURRENT_FAN_OUT_QUERIES,
fan_out_query_timeout: Self::FAN_OUT_QUERY_TIMEOUT,
observations,
supplied_objects: Vec::new(),
recall_receipt: None,
override_log: None,
}
}
pub(crate) fn registry(&self) -> &ConnectionRegistry {
&self.registry
}
pub(crate) fn state_db(&self) -> Option<&SqliteStateStore> {
self.state_db.as_ref()
}
pub(crate) fn with_supplied_objects(mut self, supplied_objects: Vec<String>) -> Self {
self.supplied_objects = supplied_objects;
self
}
#[cfg(test)]
pub(super) fn with_registry_and_fan_out_limits(
registry: ConnectionRegistry,
max_rows: usize,
allow_query_data: bool,
max_concurrent_fan_out_queries: usize,
fan_out_query_timeout: Duration,
) -> Self {
Self {
registry,
max_rows,
allow_query_data,
state_db: None,
max_concurrent_fan_out_queries: max_concurrent_fan_out_queries.max(1),
fan_out_query_timeout,
observations: None,
supplied_objects: Vec::new(),
recall_receipt: None,
override_log: None,
}
}
#[cfg(test)]
pub(super) fn with_registry_and_observations(
registry: ConnectionRegistry,
max_rows: usize,
allow_query_data: bool,
state_db: Option<SqliteStateStore>,
observations: Arc<ObservationLog>,
) -> Self {
Self {
registry,
max_rows,
allow_query_data,
state_db,
max_concurrent_fan_out_queries: Self::MAX_CONCURRENT_FAN_OUT_QUERIES,
fan_out_query_timeout: Self::FAN_OUT_QUERY_TIMEOUT,
observations: Some(observations),
supplied_objects: Vec::new(),
recall_receipt: None,
override_log: None,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[cfg(unix)]
use saya_agent::ToolExecutor;
#[test]
fn render_chart_requires_approval() {
let tools = DatabaseTools::definitions(true, false, false);
let chart_tool = tools
.iter()
.find(|tool| tool.name == "render_chart")
.expect("render_chart definition exists");
assert!(chart_tool.effect.requires_approval);
assert!(chart_tool.effect.external_side_effect);
assert!(!chart_tool.effect.database_data);
}
#[cfg(unix)]
#[tokio::test]
async fn render_chart_creates_0600_permissions_file() {
use async_trait::async_trait;
use saya_connectors::DatabaseConnector;
use saya_types::{ConnectionError, QueryRequest, QueryResult, SchemaTree, SqlDialect};
struct NonEmptyConnector;
#[async_trait]
impl DatabaseConnector for NonEmptyConnector {
fn dialect(&self) -> SqlDialect {
SqlDialect::DuckDb
}
async fn connect(&self) -> Result<(), ConnectionError> {
Ok(())
}
async fn schema(&self) -> Result<SchemaTree, ConnectionError> {
Ok(SchemaTree::default())
}
async fn execute(&self, req: QueryRequest) -> Result<QueryResult, ConnectionError> {
Ok(QueryResult {
columns: vec!["cat".into(), "val".into()],
rows: vec![serde_json::json!(["A", 10])],
row_count: 1,
truncated: false,
executed_sql: req.sql,
})
}
}
let tools = DatabaseTools::new(Some(Box::new(NonEmptyConnector)), 100, true);
let res = tools
.execute(
"render_chart",
serde_json::json!({"sql": "SELECT 1", "chart_type": "bar"}),
)
.await
.expect("render_chart should succeed");
let path_str = res["path"].as_str().expect("path in response");
let path = std::path::Path::new(path_str);
let meta = std::fs::metadata(path).expect("file should exist");
use std::os::unix::fs::PermissionsExt;
assert_eq!(meta.permissions().mode() & 0o777, 0o600);
let _ = std::fs::remove_file(path);
}
}