1use std::collections::{HashMap, HashSet};
7use std::path::{Path, PathBuf};
8use std::sync::{Arc, RwLock};
9
10use deadpool_postgres::{Object, Pool};
11use nexql_conn::{
12 ConnectionParams, PoolOptions, ProfileConfig, ResolvedConnection, checkout_guarded, create_pool,
13 resolve_profile,
14};
15use nexql_index::IndexStore;
16use nexql_policy::{AccessMode, PolicyCaps, PolicyFilter};
17use tokio::sync::RwLock as AsyncRwLock;
18
19use crate::error::ToolError;
20
21pub fn default_index_root() -> PathBuf {
23 if let Ok(p) = std::env::var("NEXQL_MCP_INDEX_DIR") {
24 return PathBuf::from(p);
25 }
26 let home = std::env::var("HOME").unwrap_or_else(|_| ".".into());
27 PathBuf::from(home).join(".local/share/nexql-mcp")
28}
29
30pub fn resolve_index_root(
33 workspace_root: Option<&Path>,
34 project_config: Option<&nexql_conn::ProjectConfigFile>,
35) -> PathBuf {
36 if let Ok(p) = std::env::var("NEXQL_MCP_INDEX_DIR") {
37 return PathBuf::from(p);
38 }
39 if let (Some(root), Some(cfg)) = (workspace_root, project_config)
40 && let Some(ref index_dir) = cfg.index_dir
41 {
42 return root.join(".nexql").join(index_dir);
43 }
44 default_index_root()
45}
46
47#[derive(Debug, Clone)]
48pub struct ConnectionPolicy {
49 pub access_mode: AccessMode,
50 pub caps: PolicyCaps,
51 pub filter: PolicyFilter,
52 pub environment: Option<String>,
53}
54
55#[derive(Debug, Clone)]
56pub struct ConnectionInfo {
57 pub id: String,
58 pub name: String,
59 pub host: Option<String>,
60 pub port: Option<u16>,
61 pub database: Option<String>,
62 pub params: ConnectionParams,
63 pub policy: ConnectionPolicy,
64}
65
66#[derive(Debug, Clone)]
67struct ActivePolicy {
68 access_mode: AccessMode,
69 caps: PolicyCaps,
70 filter: PolicyFilter,
71 pool_opts: PoolOptions,
72}
73
74pub struct ToolSession {
75 connections: RwLock<Vec<ConnectionInfo>>,
76 stale_index_keys: RwLock<HashSet<String>>,
78 policy: RwLock<ActivePolicy>,
79 pub index_store: Option<IndexStore>,
81 inner: AsyncRwLock<SessionInner>,
82}
83
84fn filter_from_profile(profile: &ProfileConfig) -> PolicyFilter {
85 PolicyFilter {
86 allow_schemas: profile.schemas.clone(),
87 deny_schemas: profile.deny_schemas.clone(),
88 deny_tables: profile.deny_tables.clone(),
89 pii_columns: profile.pii_columns.clone(),
90 }
91}
92
93pub fn policy_from_profile(
94 profile: Option<&ProfileConfig>,
95 default_mode: AccessMode,
96 default_caps: PolicyCaps,
97) -> ConnectionPolicy {
98 let access_mode = profile
99 .and_then(|p| p.access_mode.as_deref())
100 .and_then(|m| m.parse::<AccessMode>().ok())
101 .unwrap_or(default_mode);
102 let mut caps = default_caps;
103 if let Some(n) = profile.and_then(|p| p.max_rows) {
104 caps = caps.with_max_rows(n);
105 }
106 let filter = profile.map(filter_from_profile).unwrap_or_default();
107 ConnectionPolicy {
108 access_mode,
109 caps,
110 filter,
111 environment: None,
112 }
113}
114
115struct SessionInner {
116 active_id: String,
117 database: String,
118 pools: HashMap<String, Pool>,
119}
120
121fn pool_key(connection_id: &str, database: &str) -> String {
122 format!("{connection_id}\0{database}")
123}
124
125fn params_for_database(base: &ConnectionParams, database: &str) -> ConnectionParams {
126 let mut params = base.clone();
127 params.dbname = Some(database.to_string());
128 if params.host.is_some() {
130 params.url = None;
131 }
132 params
133}
134
135fn active_policy_from(policy: &ConnectionPolicy) -> ActivePolicy {
136 ActivePolicy {
137 access_mode: policy.access_mode,
138 caps: policy.caps.clone(),
139 filter: policy.filter.clone(),
140 pool_opts: PoolOptions {
141 read_only: !policy.access_mode.allows_writes(),
142 ..Default::default()
143 },
144 }
145}
146
147impl ToolSession {
148 pub fn connections(&self) -> Vec<ConnectionInfo> {
149 self.connections
150 .read()
151 .map(|c| c.clone())
152 .unwrap_or_default()
153 }
154
155 pub fn register_profile(
157 &self,
158 name: &str,
159 profile: &ProfileConfig,
160 default_mode: AccessMode,
161 default_caps: PolicyCaps,
162 ) -> Result<(), ToolError> {
163 let params = resolve_profile(profile).map_err(ToolError::Conn)?;
164 let policy = policy_from_profile(Some(profile), default_mode, default_caps);
165 let info = ConnectionInfo {
166 id: name.to_string(),
167 name: name.to_string(),
168 host: params.host.clone(),
169 port: params.port,
170 database: params.dbname.clone(),
171 params,
172 policy,
173 };
174 let mut conns = self
175 .connections
176 .write()
177 .map_err(|_| ToolError::Execution("connection registry lock poisoned".into()))?;
178 if let Some(existing) = conns.iter_mut().find(|c| c.id == name) {
179 *existing = info;
180 } else {
181 conns.push(info);
182 }
183 Ok(())
184 }
185
186 pub fn mark_index_stale(&self, connection_id: &str, database: &str) {
187 let key = pool_key(connection_id, database);
188 if let Ok(mut keys) = self.stale_index_keys.write() {
189 keys.insert(key);
190 }
191 }
192
193 pub fn clear_index_stale(&self, connection_id: &str, database: &str) {
194 let key = pool_key(connection_id, database);
195 if let Ok(mut keys) = self.stale_index_keys.write() {
196 keys.remove(&key);
197 }
198 }
199
200 pub fn is_index_stale(&self, connection_id: &str, database: &str) -> bool {
201 let key = pool_key(connection_id, database);
202 self.stale_index_keys
203 .read()
204 .map(|keys| keys.contains(&key))
205 .unwrap_or(false)
206 }
207
208 pub fn access_mode(&self) -> AccessMode {
209 self.policy
210 .read()
211 .map(|p| p.access_mode)
212 .unwrap_or(AccessMode::Read)
213 }
214
215 pub fn caps(&self) -> PolicyCaps {
216 self.policy
217 .read()
218 .map(|p| p.caps.clone())
219 .unwrap_or_default()
220 }
221
222 pub fn filter(&self) -> PolicyFilter {
223 self.policy
224 .read()
225 .map(|p| p.filter.clone())
226 .unwrap_or_default()
227 }
228
229 pub fn pool_opts(&self) -> PoolOptions {
230 self.policy
231 .read()
232 .map(|p| p.pool_opts.clone())
233 .unwrap_or_default()
234 }
235
236 fn apply_policy(&self, policy: &ConnectionPolicy) {
237 if let Ok(mut active) = self.policy.write() {
238 *active = active_policy_from(policy);
239 }
240 }
241
242 pub async fn from_resolved(
243 resolved: ResolvedConnection,
244 access_mode: AccessMode,
245 caps: PolicyCaps,
246 ) -> Result<Arc<Self>, ToolError> {
247 let id = resolved
248 .profile_name
249 .clone()
250 .unwrap_or_else(|| "default".into());
251 let conn_policy = policy_from_profile(resolved.profile.as_ref(), access_mode, caps);
252 let info = ConnectionInfo {
253 id: id.clone(),
254 name: id.clone(),
255 host: resolved.params.host.clone(),
256 port: resolved.params.port,
257 database: resolved.params.dbname.clone(),
258 params: resolved.params,
259 policy: conn_policy,
260 };
261 Self::from_connections(vec![info], Some(id)).await
262 }
263
264 pub async fn from_connections(
265 connections: Vec<ConnectionInfo>,
266 active_id: Option<String>,
267 ) -> Result<Arc<Self>, ToolError> {
268 if connections.is_empty() {
269 return Err(ToolError::Execution("no connections configured".into()));
270 }
271 let active = active_id.unwrap_or_else(|| connections[0].id.clone());
272 let active_conn = connections
273 .iter()
274 .find(|c| c.id == active)
275 .unwrap_or(&connections[0]);
276 let active_policy = active_policy_from(&active_conn.policy);
277 let pool_opts = active_policy.pool_opts.clone();
278 let mut pools = HashMap::new();
279 for c in &connections {
280 let database = c.database.as_deref().unwrap_or("postgres");
281 let pool = create_pool(&c.params, &pool_opts).await?;
282 pools.insert(pool_key(&c.id, database), pool);
283 }
284 let database = active_conn
285 .database
286 .clone()
287 .unwrap_or_else(|| "postgres".into());
288 Ok(Arc::new(Self {
289 connections: RwLock::new(connections),
290 stale_index_keys: RwLock::new(HashSet::new()),
291 policy: RwLock::new(active_policy),
292 index_store: Some(IndexStore::new(default_index_root())),
293 inner: AsyncRwLock::new(SessionInner {
294 active_id: active,
295 database,
296 pools,
297 }),
298 }))
299 }
300
301 pub async fn from_connections_with_filter(
302 connections: Vec<ConnectionInfo>,
303 access_mode: AccessMode,
304 caps: PolicyCaps,
305 active_id: Option<String>,
306 filter: PolicyFilter,
307 ) -> Result<Arc<Self>, ToolError> {
308 let connections = connections
309 .into_iter()
310 .map(|mut c| {
311 c.policy.access_mode = access_mode;
312 c.policy.caps = caps.clone();
313 c.policy.filter = filter.clone();
314 c
315 })
316 .collect();
317 Self::from_connections(connections, active_id).await
318 }
319
320 pub async fn active_context(&self) -> (String, String) {
321 let g = self.inner.read().await;
322 (g.active_id.clone(), g.database.clone())
323 }
324
325 pub async fn switch(
326 &self,
327 connection_id: &str,
328 database: Option<String>,
329 ) -> Result<(), ToolError> {
330 let conn = self
331 .connections()
332 .into_iter()
333 .find(|c| c.id == connection_id)
334 .ok_or_else(|| {
335 ToolError::Execution(format!(
336 "Connection not found for ID: {connection_id} — call list_connections"
337 ))
338 })?;
339 let target_db = database
340 .or_else(|| conn.database.clone())
341 .unwrap_or_else(|| "postgres".into());
342 self.apply_policy(&conn.policy);
343 let pool_opts = self.pool_opts();
344 let key = pool_key(connection_id, &target_db);
345 let mut g = self.inner.write().await;
346 if let std::collections::hash_map::Entry::Vacant(e) = g.pools.entry(key) {
347 let params = params_for_database(&conn.params, &target_db);
348 let pool = create_pool(¶ms, &pool_opts).await?;
349 e.insert(pool);
350 }
351 g.active_id = connection_id.to_string();
352 g.database = target_db;
353 Ok(())
354 }
355
356 pub async fn checkout(&self) -> Result<Object, ToolError> {
357 let g = self.inner.read().await;
358 let key = pool_key(&g.active_id, &g.database);
359 let pool = g
360 .pools
361 .get(&key)
362 .ok_or_else(|| ToolError::Execution("no active pool".into()))?;
363 let pool_opts = self.pool_opts();
364 Ok(checkout_guarded(pool, &pool_opts).await?)
365 }
366
367 #[cfg(test)]
369 pub fn for_tests(
370 connections: Vec<ConnectionInfo>,
371 filter: PolicyFilter,
372 index_store: Option<IndexStore>,
373 ) -> Arc<Self> {
374 assert!(!connections.is_empty());
375 let active = connections[0].id.clone();
376 let database = connections[0]
377 .database
378 .clone()
379 .unwrap_or_else(|| "postgres".into());
380 let policy = ConnectionPolicy {
381 access_mode: AccessMode::Read,
382 caps: PolicyCaps::default(),
383 filter,
384 environment: None,
385 };
386 Arc::new(Self {
387 connections: RwLock::new(connections),
388 stale_index_keys: RwLock::new(HashSet::new()),
389 policy: RwLock::new(active_policy_from(&policy)),
390 index_store,
391 inner: AsyncRwLock::new(SessionInner {
392 active_id: active,
393 database,
394 pools: HashMap::new(),
395 }),
396 })
397 }
398}
399
400#[cfg(test)]
401mod tests {
402 use super::*;
403
404 #[test]
405 fn filter_maps_profile_fields() {
406 let profile = ProfileConfig {
407 schemas: vec!["public".into()],
408 deny_schemas: vec!["pgboss".into()],
409 deny_tables: vec!["auth.*".into()],
410 pii_columns: vec!["public.users.ssn".into()],
411 ..Default::default()
412 };
413 let f = filter_from_profile(&profile);
414 assert_eq!(f.allow_schemas, vec!["public"]);
415 assert_eq!(f.deny_schemas, vec!["pgboss"]);
416 assert_eq!(f.deny_tables, vec!["auth.*"]);
417 assert_eq!(f.pii_columns, vec!["public.users.ssn"]);
418 assert!(f.allows_schema("public"));
419 assert!(!f.allows_schema("pgboss"));
420 assert!(!f.allows_table("auth", "sessions"));
421 }
422
423 #[test]
424 fn register_profile_upserts_connection() {
425 let session = ToolSession::for_tests(
426 vec![ConnectionInfo {
427 id: "existing".into(),
428 name: "existing".into(),
429 host: Some("localhost".into()),
430 port: Some(5432),
431 database: Some("postgres".into()),
432 params: ConnectionParams::default(),
433 policy: policy_from_profile(None, AccessMode::Read, PolicyCaps::default()),
434 }],
435 PolicyFilter::default(),
436 None,
437 );
438 let profile = ProfileConfig {
439 host: Some("db.example.com".into()),
440 port: Some(5432),
441 dbname: Some("app".into()),
442 user: Some("app".into()),
443 password: Some("secret".into()),
444 ..Default::default()
445 };
446 session
447 .register_profile("newdb", &profile, AccessMode::Read, PolicyCaps::default())
448 .unwrap();
449 let names: Vec<_> = session.connections().iter().map(|c| c.id.clone()).collect();
450 assert!(names.contains(&"existing".to_string()));
451 assert!(names.contains(&"newdb".to_string()));
452 }
453
454 #[test]
455 fn index_stale_marker_round_trip() {
456 let session = ToolSession::for_tests(
457 vec![ConnectionInfo {
458 id: "local".into(),
459 name: "local".into(),
460 host: None,
461 port: None,
462 database: Some("postgres".into()),
463 params: ConnectionParams::default(),
464 policy: policy_from_profile(None, AccessMode::Read, PolicyCaps::default()),
465 }],
466 PolicyFilter::default(),
467 None,
468 );
469 assert!(!session.is_index_stale("local", "postgres"));
470 session.mark_index_stale("local", "postgres");
471 assert!(session.is_index_stale("local", "postgres"));
472 session.clear_index_stale("local", "postgres");
473 assert!(!session.is_index_stale("local", "postgres"));
474 }
475}