use saya_agent::ToolError;
use saya_connectors::DatabaseConnector;
use saya_types::SqlDialect;
use std::collections::HashMap;
use std::fmt;
#[allow(dead_code)]
pub(crate) struct ConnectionEntry {
pub(crate) connector: Box<dyn DatabaseConnector>,
pub(crate) dialect: SqlDialect,
pub(crate) profile_id: Option<String>,
}
impl fmt::Debug for ConnectionEntry {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ConnectionEntry")
.field("dialect", &self.dialect)
.field("profile_id", &self.profile_id)
.finish()
}
}
#[allow(dead_code)]
pub(crate) struct ConnectionRegistry {
primary: String,
names: Vec<String>,
map: HashMap<String, ConnectionEntry>,
}
#[allow(dead_code)]
impl ConnectionRegistry {
pub(crate) fn new(primary: impl Into<String>) -> Self {
Self {
primary: primary.into(),
names: Vec::new(),
map: HashMap::new(),
}
}
pub(crate) fn insert(&mut self, name: impl Into<String>, entry: ConnectionEntry) {
let name = name.into();
if !self.map.contains_key(&name) {
self.names.push(name.clone());
}
self.map.insert(name, entry);
}
pub(crate) fn primary_name(&self) -> &str {
&self.primary
}
pub(crate) fn len(&self) -> usize {
self.map.len()
}
pub(crate) fn is_empty(&self) -> bool {
self.map.is_empty()
}
pub(crate) fn names(&self) -> Vec<&str> {
self.names.iter().map(String::as_str).collect()
}
pub(crate) fn entries(&self) -> Vec<(&str, &ConnectionEntry)> {
self.names
.iter()
.filter_map(|name| self.map.get(name).map(|entry| (name.as_str(), entry)))
.collect()
}
pub(crate) fn resolve(&self, name: Option<&str>) -> Result<&ConnectionEntry, ToolError> {
let target = match name {
None | Some("") => self.primary.as_str(),
Some(n) => n,
};
if let Some(entry) = self.map.get(target) {
Ok(entry)
} else if self.is_empty() {
Err(ToolError::NoConnectionSelected)
} else {
let available = self.names().join(", ");
Err(ToolError::UnknownConnection {
target: target.to_string(),
available,
})
}
}
pub(crate) fn name_for_identity(&self, identity: &str) -> Option<&str> {
self.entries()
.into_iter()
.find(|(_, entry)| entry.profile_id.as_deref() == Some(identity))
.map(|(name, _)| name)
}
pub(crate) fn describe_context(&self) -> Option<String> {
if self.len() <= 1 {
return None;
}
let mut lines = Vec::new();
lines.push("Available database connections:".to_string());
for name in &self.names {
if let Some(entry) = self.map.get(name) {
lines.push(format!("- {name} ({})", entry.dialect.as_str()));
}
}
lines.push(
"To inspect a database, pass its `connection` argument to schema and query tools. Inspect each database separately and combine your findings. When the same query should run against every connected database, call `bounded_sql_query_all` once instead of repeating `bounded_sql_query` per connection."
.to_string(),
);
Some(lines.join("\n"))
}
}
#[cfg(test)]
#[path = "registry_tests.rs"]
mod tests;