1use std::collections::HashMap;
4use std::path::PathBuf;
5use std::sync::Arc;
6
7use deadpool_postgres::{Object, Pool};
8use nexql_conn::{
9 ConnectionParams, PoolOptions, ProfileConfig, ResolvedConnection, checkout_guarded, create_pool,
10};
11use nexql_index::IndexStore;
12use nexql_policy::{AccessMode, PolicyCaps, PolicyFilter};
13use tokio::sync::RwLock;
14
15use crate::error::ToolError;
16
17pub fn default_index_root() -> PathBuf {
19 if let Ok(p) = std::env::var("NEXQL_MCP_INDEX_DIR") {
20 return PathBuf::from(p);
21 }
22 let home = std::env::var("HOME").unwrap_or_else(|_| ".".into());
23 PathBuf::from(home).join(".local/share/nexql-mcp")
24}
25
26#[derive(Debug, Clone)]
27pub struct ConnectionInfo {
28 pub id: String,
29 pub name: String,
30 pub host: Option<String>,
31 pub port: Option<u16>,
32 pub database: Option<String>,
33 pub params: ConnectionParams,
34}
35
36pub struct ToolSession {
37 pub connections: Vec<ConnectionInfo>,
38 pub access_mode: AccessMode,
39 pub caps: PolicyCaps,
40 pub filter: PolicyFilter,
41 pub pool_opts: PoolOptions,
42 pub index_store: Option<IndexStore>,
44 inner: RwLock<SessionInner>,
45}
46
47fn filter_from_profile(profile: &ProfileConfig) -> PolicyFilter {
48 PolicyFilter {
49 allow_schemas: profile.schemas.clone(),
50 deny_schemas: profile.deny_schemas.clone(),
51 deny_tables: profile.deny_tables.clone(),
52 pii_columns: profile.pii_columns.clone(),
53 }
54}
55
56struct SessionInner {
57 active_id: String,
58 database: String,
59 pools: HashMap<String, Pool>,
60}
61
62fn pool_key(connection_id: &str, database: &str) -> String {
63 format!("{connection_id}\0{database}")
64}
65
66fn params_for_database(base: &ConnectionParams, database: &str) -> ConnectionParams {
67 let mut params = base.clone();
68 params.dbname = Some(database.to_string());
69 if params.host.is_some() {
71 params.url = None;
72 }
73 params
74}
75
76impl ToolSession {
77 pub async fn from_resolved(
78 resolved: ResolvedConnection,
79 access_mode: AccessMode,
80 caps: PolicyCaps,
81 ) -> Result<Arc<Self>, ToolError> {
82 let id = resolved
83 .profile_name
84 .clone()
85 .unwrap_or_else(|| "default".into());
86 let filter = resolved
87 .profile
88 .as_ref()
89 .map(filter_from_profile)
90 .unwrap_or_default();
91 let info = ConnectionInfo {
92 id: id.clone(),
93 name: id.clone(),
94 host: resolved.params.host.clone(),
95 port: resolved.params.port,
96 database: resolved.params.dbname.clone(),
97 params: resolved.params,
98 };
99 Self::from_connections_with_filter(vec![info], access_mode, caps, Some(id), filter).await
100 }
101
102 pub async fn from_connections(
103 connections: Vec<ConnectionInfo>,
104 access_mode: AccessMode,
105 caps: PolicyCaps,
106 active_id: Option<String>,
107 ) -> Result<Arc<Self>, ToolError> {
108 Self::from_connections_with_filter(
109 connections,
110 access_mode,
111 caps,
112 active_id,
113 PolicyFilter::default(),
114 )
115 .await
116 }
117
118 pub async fn from_connections_with_filter(
119 connections: Vec<ConnectionInfo>,
120 access_mode: AccessMode,
121 caps: PolicyCaps,
122 active_id: Option<String>,
123 filter: PolicyFilter,
124 ) -> Result<Arc<Self>, ToolError> {
125 if connections.is_empty() {
126 return Err(ToolError::Execution("no connections configured".into()));
127 }
128 let pool_opts = PoolOptions {
129 read_only: !access_mode.allows_writes(),
130 ..Default::default()
131 };
132 let active = active_id.unwrap_or_else(|| connections[0].id.clone());
133 let mut pools = HashMap::new();
134 for c in &connections {
135 let database = c.database.as_deref().unwrap_or("postgres");
136 let pool = create_pool(&c.params, &pool_opts).await?;
137 pools.insert(pool_key(&c.id, database), pool);
138 }
139 let database = connections
140 .iter()
141 .find(|c| c.id == active)
142 .and_then(|c| c.database.clone())
143 .unwrap_or_else(|| "postgres".into());
144 Ok(Arc::new(Self {
145 connections,
146 access_mode,
147 caps,
148 filter,
149 pool_opts,
150 index_store: Some(IndexStore::new(default_index_root())),
151 inner: RwLock::new(SessionInner {
152 active_id: active,
153 database,
154 pools,
155 }),
156 }))
157 }
158
159 pub async fn active_context(&self) -> (String, String) {
160 let g = self.inner.read().await;
161 (g.active_id.clone(), g.database.clone())
162 }
163
164 pub async fn switch(
165 &self,
166 connection_id: &str,
167 database: Option<String>,
168 ) -> Result<(), ToolError> {
169 let conn = self
170 .connections
171 .iter()
172 .find(|c| c.id == connection_id)
173 .ok_or_else(|| {
174 ToolError::Execution(format!(
175 "Connection not found for ID: {connection_id} — call list_connections"
176 ))
177 })?;
178 let target_db = database
179 .or_else(|| conn.database.clone())
180 .unwrap_or_else(|| "postgres".into());
181 let key = pool_key(connection_id, &target_db);
182 let mut g = self.inner.write().await;
183 if let std::collections::hash_map::Entry::Vacant(e) = g.pools.entry(key) {
184 let params = params_for_database(&conn.params, &target_db);
185 let pool = create_pool(¶ms, &self.pool_opts).await?;
186 e.insert(pool);
187 }
188 g.active_id = connection_id.to_string();
189 g.database = target_db;
190 Ok(())
191 }
192
193 pub async fn checkout(&self) -> Result<Object, ToolError> {
194 let g = self.inner.read().await;
195 let key = pool_key(&g.active_id, &g.database);
196 let pool = g
197 .pools
198 .get(&key)
199 .ok_or_else(|| ToolError::Execution("no active pool".into()))?;
200 Ok(checkout_guarded(pool, &self.pool_opts).await?)
201 }
202
203 #[cfg(test)]
205 pub fn for_tests(
206 connections: Vec<ConnectionInfo>,
207 filter: PolicyFilter,
208 index_store: Option<IndexStore>,
209 ) -> Arc<Self> {
210 assert!(!connections.is_empty());
211 let active = connections[0].id.clone();
212 let database = connections[0]
213 .database
214 .clone()
215 .unwrap_or_else(|| "postgres".into());
216 Arc::new(Self {
217 connections,
218 access_mode: AccessMode::Read,
219 caps: PolicyCaps::default(),
220 filter,
221 pool_opts: PoolOptions::default(),
222 index_store,
223 inner: RwLock::new(SessionInner {
224 active_id: active,
225 database,
226 pools: HashMap::new(),
227 }),
228 })
229 }
230}
231
232#[cfg(test)]
233mod tests {
234 use super::*;
235
236 #[test]
237 fn filter_maps_profile_fields() {
238 let profile = ProfileConfig {
239 schemas: vec!["public".into()],
240 deny_schemas: vec!["pgboss".into()],
241 deny_tables: vec!["auth.*".into()],
242 pii_columns: vec!["public.users.ssn".into()],
243 ..Default::default()
244 };
245 let f = filter_from_profile(&profile);
246 assert_eq!(f.allow_schemas, vec!["public"]);
247 assert_eq!(f.deny_schemas, vec!["pgboss"]);
248 assert_eq!(f.deny_tables, vec!["auth.*"]);
249 assert_eq!(f.pii_columns, vec!["public.users.ssn"]);
250 assert!(f.allows_schema("public"));
251 assert!(!f.allows_schema("pgboss"));
252 assert!(!f.allows_table("auth", "sessions"));
253 }
254}