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, clamp_statement_timeout_ms};
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 if let Some(ms) = profile.and_then(|p| p.statement_timeout_ms) {
107 caps = caps.with_statement_timeout_ms(ms);
108 }
109 let filter = profile.map(filter_from_profile).unwrap_or_default();
110 ConnectionPolicy {
111 access_mode,
112 caps,
113 filter,
114 environment: None,
115 }
116}
117
118struct SessionInner {
119 active_id: String,
120 database: String,
121 pools: HashMap<String, Pool>,
122}
123
124fn pool_key(connection_id: &str, database: &str) -> String {
125 format!("{connection_id}\0{database}")
126}
127
128fn params_for_database(base: &ConnectionParams, database: &str) -> ConnectionParams {
129 let mut params = base.clone();
130 params.dbname = Some(database.to_string());
131 if params.host.is_some() {
133 params.url = None;
134 }
135 params
136}
137
138fn active_policy_from(policy: &ConnectionPolicy) -> ActivePolicy {
139 ActivePolicy {
140 access_mode: policy.access_mode,
141 caps: policy.caps.clone(),
142 filter: policy.filter.clone(),
143 pool_opts: PoolOptions {
144 read_only: !policy.access_mode.allows_writes(),
145 statement_timeout: std::time::Duration::from_millis(
146 policy.caps.statement_timeout_ms as u64,
147 ),
148 ..Default::default()
149 },
150 }
151}
152
153#[derive(Debug, Clone, PartialEq, Eq)]
155pub struct ScopedContext {
156 pub connection_id: String,
157 pub database: String,
158}
159
160pub enum CheckoutTarget<'a> {
162 Active,
163 Scoped(&'a ScopedContext),
164}
165
166impl ToolSession {
167 pub fn connections(&self) -> Vec<ConnectionInfo> {
168 self.connections
169 .read()
170 .map(|c| c.clone())
171 .unwrap_or_default()
172 }
173
174 pub fn register_profile(
176 &self,
177 name: &str,
178 profile: &ProfileConfig,
179 default_mode: AccessMode,
180 default_caps: PolicyCaps,
181 ) -> Result<(), ToolError> {
182 let params = resolve_profile(name, profile).map_err(ToolError::Conn)?;
183 let policy = policy_from_profile(Some(profile), default_mode, default_caps);
184 let info = ConnectionInfo {
185 id: name.to_string(),
186 name: name.to_string(),
187 host: params.host.clone(),
188 port: params.port,
189 database: params.dbname.clone(),
190 params,
191 policy,
192 };
193 let mut conns = self
194 .connections
195 .write()
196 .map_err(|_| ToolError::Execution("connection registry lock poisoned".into()))?;
197 if let Some(existing) = conns.iter_mut().find(|c| c.id == name) {
198 *existing = info;
199 } else {
200 conns.push(info);
201 }
202 Ok(())
203 }
204
205 pub fn mark_index_stale(&self, connection_id: &str, database: &str) {
206 let key = pool_key(connection_id, database);
207 if let Ok(mut keys) = self.stale_index_keys.write() {
208 keys.insert(key);
209 }
210 }
211
212 pub fn clear_index_stale(&self, connection_id: &str, database: &str) {
213 let key = pool_key(connection_id, database);
214 if let Ok(mut keys) = self.stale_index_keys.write() {
215 keys.remove(&key);
216 }
217 }
218
219 pub fn is_index_stale(&self, connection_id: &str, database: &str) -> bool {
220 let key = pool_key(connection_id, database);
221 self.stale_index_keys
222 .read()
223 .map(|keys| keys.contains(&key))
224 .unwrap_or(false)
225 }
226
227 pub fn access_mode(&self) -> AccessMode {
228 self.policy
229 .read()
230 .map(|p| p.access_mode)
231 .unwrap_or(AccessMode::Read)
232 }
233
234 pub fn caps(&self) -> PolicyCaps {
235 self.policy
236 .read()
237 .map(|p| p.caps.clone())
238 .unwrap_or_default()
239 }
240
241 pub fn filter(&self) -> PolicyFilter {
242 self.policy
243 .read()
244 .map(|p| p.filter.clone())
245 .unwrap_or_default()
246 }
247
248 pub fn pool_opts(&self) -> PoolOptions {
249 self.policy
250 .read()
251 .map(|p| p.pool_opts.clone())
252 .unwrap_or_default()
253 }
254
255 fn apply_policy(&self, policy: &ConnectionPolicy) {
256 if let Ok(mut active) = self.policy.write() {
257 *active = active_policy_from(policy);
258 }
259 }
260
261 pub async fn from_resolved(
262 resolved: ResolvedConnection,
263 access_mode: AccessMode,
264 caps: PolicyCaps,
265 ) -> Result<Arc<Self>, ToolError> {
266 let id = resolved
267 .profile_name
268 .clone()
269 .unwrap_or_else(|| "default".into());
270 let conn_policy = policy_from_profile(resolved.profile.as_ref(), access_mode, caps);
271 let info = ConnectionInfo {
272 id: id.clone(),
273 name: id.clone(),
274 host: resolved.params.host.clone(),
275 port: resolved.params.port,
276 database: resolved.params.dbname.clone(),
277 params: resolved.params,
278 policy: conn_policy,
279 };
280 Self::from_connections(vec![info], Some(id)).await
281 }
282
283 pub async fn from_connections(
284 connections: Vec<ConnectionInfo>,
285 active_id: Option<String>,
286 ) -> Result<Arc<Self>, ToolError> {
287 if connections.is_empty() {
288 return Err(ToolError::Execution("no connections configured".into()));
289 }
290 let active = active_id.unwrap_or_else(|| connections[0].id.clone());
291 let active_conn = connections
292 .iter()
293 .find(|c| c.id == active)
294 .unwrap_or(&connections[0]);
295 let active_policy = active_policy_from(&active_conn.policy);
296 let pool_opts = active_policy.pool_opts.clone();
297 let mut pools = HashMap::new();
298 for c in &connections {
299 let database = c.database.as_deref().unwrap_or("postgres");
300 let pool = create_pool(&c.params, &pool_opts).await?;
301 pools.insert(pool_key(&c.id, database), pool);
302 }
303 let database = active_conn
304 .database
305 .clone()
306 .unwrap_or_else(|| "postgres".into());
307 Ok(Arc::new(Self {
308 connections: RwLock::new(connections),
309 stale_index_keys: RwLock::new(HashSet::new()),
310 policy: RwLock::new(active_policy),
311 index_store: Some(IndexStore::new(default_index_root())),
312 inner: AsyncRwLock::new(SessionInner {
313 active_id: active,
314 database,
315 pools,
316 }),
317 }))
318 }
319
320 pub async fn from_connections_with_filter(
321 connections: Vec<ConnectionInfo>,
322 access_mode: AccessMode,
323 caps: PolicyCaps,
324 active_id: Option<String>,
325 filter: PolicyFilter,
326 ) -> Result<Arc<Self>, ToolError> {
327 let connections = connections
328 .into_iter()
329 .map(|mut c| {
330 c.policy.access_mode = access_mode;
331 c.policy.caps = caps.clone();
332 c.policy.filter = filter.clone();
333 c
334 })
335 .collect();
336 Self::from_connections(connections, active_id).await
337 }
338
339 pub async fn active_context(&self) -> (String, String) {
340 let g = self.inner.read().await;
341 (g.active_id.clone(), g.database.clone())
342 }
343
344 pub async fn switch(
345 &self,
346 connection_id: &str,
347 database: Option<String>,
348 ) -> Result<(), ToolError> {
349 let conn = self
350 .connections()
351 .into_iter()
352 .find(|c| c.id == connection_id)
353 .ok_or_else(|| {
354 ToolError::Execution(format!(
355 "Connection not found for ID: {connection_id} — call list_connections"
356 ))
357 })?;
358 let target_db = database
359 .or_else(|| conn.database.clone())
360 .unwrap_or_else(|| "postgres".into());
361 self.apply_policy(&conn.policy);
362 let pool_opts = self.pool_opts();
363 let key = pool_key(connection_id, &target_db);
364 let mut g = self.inner.write().await;
365 if let std::collections::hash_map::Entry::Vacant(e) = g.pools.entry(key) {
366 let params = params_for_database(&conn.params, &target_db);
367 let pool = create_pool(¶ms, &pool_opts).await?;
368 e.insert(pool);
369 }
370 g.active_id = connection_id.to_string();
371 g.database = target_db;
372 Ok(())
373 }
374
375 pub async fn checkout(&self) -> Result<Object, ToolError> {
376 let g = self.inner.read().await;
377 let key = pool_key(&g.active_id, &g.database);
378 let pool = g
379 .pools
380 .get(&key)
381 .ok_or_else(|| ToolError::Execution("no active pool".into()))?;
382 let pool_opts = self.pool_opts();
383 Ok(checkout_guarded(pool, &pool_opts).await?)
384 }
385
386 pub fn connection_policy(&self, connection_id: &str) -> Option<ConnectionPolicy> {
387 self.connections()
388 .into_iter()
389 .find(|c| c.id == connection_id)
390 .map(|c| c.policy.clone())
391 }
392
393 pub fn filter_for(&self, connection_id: &str) -> PolicyFilter {
394 self.connection_policy(connection_id)
395 .map(|p| p.filter)
396 .unwrap_or_else(|| self.filter())
397 }
398
399 pub fn caps_for(&self, connection_id: &str) -> PolicyCaps {
400 self.connection_policy(connection_id)
401 .map(|p| p.caps)
402 .unwrap_or_else(|| self.caps())
403 }
404
405 pub fn pool_opts_for(&self, connection_id: &str) -> PoolOptions {
406 if let Some(policy) = self.connection_policy(connection_id) {
407 PoolOptions {
408 read_only: !policy.access_mode.allows_writes(),
409 statement_timeout: std::time::Duration::from_millis(
410 policy.caps.statement_timeout_ms as u64,
411 ),
412 ..Default::default()
413 }
414 } else {
415 self.pool_opts()
416 }
417 }
418
419 pub async fn resolve_scoped_context(
421 &self,
422 connection_id: Option<&str>,
423 database: Option<&str>,
424 ) -> Result<ScopedContext, ToolError> {
425 if let Some(id) = connection_id {
426 let conn = self
427 .connections()
428 .into_iter()
429 .find(|c| c.id == id)
430 .ok_or_else(|| {
431 ToolError::Execution(format!(
432 "Connection not found for ID: {id} — call list_connections"
433 ))
434 })?;
435 let target_db = database
436 .map(str::to_owned)
437 .or_else(|| conn.database.clone())
438 .unwrap_or_else(|| "postgres".into());
439 Ok(ScopedContext {
440 connection_id: id.to_string(),
441 database: target_db,
442 })
443 } else {
444 if database.is_some() {
445 return Err(ToolError::InvalidArgs(
446 "database requires connectionId when overriding the active session".into(),
447 ));
448 }
449 let (id, db) = self.active_context().await;
450 Ok(ScopedContext {
451 connection_id: id,
452 database: db,
453 })
454 }
455 }
456
457 async fn ensure_pool_for(&self, ctx: &ScopedContext) -> Result<(), ToolError> {
458 let conn = self
459 .connections()
460 .into_iter()
461 .find(|c| c.id == ctx.connection_id)
462 .ok_or_else(|| {
463 ToolError::Execution(format!(
464 "Connection not found for ID: {} — call list_connections",
465 ctx.connection_id
466 ))
467 })?;
468 let pool_opts = self.pool_opts_for(&ctx.connection_id);
469 let key = pool_key(&ctx.connection_id, &ctx.database);
470 let mut g = self.inner.write().await;
471 if g.pools.contains_key(&key) {
472 return Ok(());
473 }
474 let params = params_for_database(&conn.params, &ctx.database);
475 let pool = create_pool(¶ms, &pool_opts).await?;
476 g.pools.insert(key, pool);
477 Ok(())
478 }
479
480 pub async fn checkout_for(
482 &self,
483 target: CheckoutTarget<'_>,
484 ) -> Result<(Object, ScopedContext), ToolError> {
485 match target {
486 CheckoutTarget::Active => {
487 let (id, db) = self.active_context().await;
488 let ctx = ScopedContext {
489 connection_id: id,
490 database: db,
491 };
492 Ok((self.checkout().await?, ctx))
493 }
494 CheckoutTarget::Scoped(ctx) => {
495 self.ensure_pool_for(ctx).await?;
496 let key = pool_key(&ctx.connection_id, &ctx.database);
497 let pool = {
498 let g = self.inner.read().await;
499 g.pools
500 .get(&key)
501 .cloned()
502 .ok_or_else(|| ToolError::Execution("no pool for scoped context".into()))?
503 };
504 let pool_opts = self.pool_opts_for(&ctx.connection_id);
505 let client = checkout_guarded(&pool, &pool_opts).await?;
506 Ok((client, ctx.clone()))
507 }
508 }
509 }
510
511 pub async fn set_statement_timeout(
512 client: &Object,
513 timeout_ms: u32,
514 ) -> Result<(), ToolError> {
515 let ms = clamp_statement_timeout_ms(timeout_ms);
516 client
517 .batch_execute(&format!("SET statement_timeout = '{ms}ms'"))
518 .await
519 .map_err(|e| ToolError::Execution(e.to_string()))?;
520 Ok(())
521 }
522
523 #[cfg(test)]
525 pub fn for_tests(
526 connections: Vec<ConnectionInfo>,
527 filter: PolicyFilter,
528 index_store: Option<IndexStore>,
529 ) -> Arc<Self> {
530 assert!(!connections.is_empty());
531 let active = connections[0].id.clone();
532 let database = connections[0]
533 .database
534 .clone()
535 .unwrap_or_else(|| "postgres".into());
536 let policy = ConnectionPolicy {
537 access_mode: AccessMode::Read,
538 caps: PolicyCaps::default(),
539 filter,
540 environment: None,
541 };
542 Arc::new(Self {
543 connections: RwLock::new(connections),
544 stale_index_keys: RwLock::new(HashSet::new()),
545 policy: RwLock::new(active_policy_from(&policy)),
546 index_store,
547 inner: AsyncRwLock::new(SessionInner {
548 active_id: active,
549 database,
550 pools: HashMap::new(),
551 }),
552 })
553 }
554}
555
556#[cfg(test)]
557mod tests {
558 use super::*;
559 use nexql_conn::ConnectionParams;
560
561 #[test]
562 fn filter_maps_profile_fields() {
563 let profile = ProfileConfig {
564 schemas: vec!["public".into()],
565 deny_schemas: vec!["pgboss".into()],
566 deny_tables: vec!["auth.*".into()],
567 pii_columns: vec!["public.users.ssn".into()],
568 ..Default::default()
569 };
570 let f = filter_from_profile(&profile);
571 assert_eq!(f.allow_schemas, vec!["public"]);
572 assert_eq!(f.deny_schemas, vec!["pgboss"]);
573 assert_eq!(f.deny_tables, vec!["auth.*"]);
574 assert_eq!(f.pii_columns, vec!["public.users.ssn"]);
575 assert!(f.allows_schema("public"));
576 assert!(!f.allows_schema("pgboss"));
577 assert!(!f.allows_table("auth", "sessions"));
578 }
579
580 #[test]
581 fn register_profile_upserts_connection() {
582 let session = ToolSession::for_tests(
583 vec![ConnectionInfo {
584 id: "existing".into(),
585 name: "existing".into(),
586 host: Some("localhost".into()),
587 port: Some(5432),
588 database: Some("postgres".into()),
589 params: ConnectionParams::default(),
590 policy: policy_from_profile(None, AccessMode::Read, PolicyCaps::default()),
591 }],
592 PolicyFilter::default(),
593 None,
594 );
595 let profile = ProfileConfig {
596 host: Some("db.example.com".into()),
597 port: Some(5432),
598 dbname: Some("app".into()),
599 user: Some("app".into()),
600 password: Some("secret".into()),
601 ..Default::default()
602 };
603 session
604 .register_profile("newdb", &profile, AccessMode::Read, PolicyCaps::default())
605 .unwrap();
606 let names: Vec<_> = session.connections().iter().map(|c| c.id.clone()).collect();
607 assert!(names.contains(&"existing".to_string()));
608 assert!(names.contains(&"newdb".to_string()));
609 }
610
611 #[tokio::test]
612 async fn checkout_for_scoped_does_not_change_active_context() {
613 let session = ToolSession::for_tests(
614 vec![
615 ConnectionInfo {
616 id: "a".into(),
617 name: "a".into(),
618 host: Some("localhost".into()),
619 port: Some(5432),
620 database: Some("postgres".into()),
621 params: ConnectionParams::default(),
622 policy: policy_from_profile(None, AccessMode::Read, PolicyCaps::default()),
623 },
624 ConnectionInfo {
625 id: "b".into(),
626 name: "b".into(),
627 host: Some("localhost".into()),
628 port: Some(5432),
629 database: Some("postgres".into()),
630 params: ConnectionParams::default(),
631 policy: policy_from_profile(None, AccessMode::Read, PolicyCaps::default()),
632 },
633 ],
634 PolicyFilter::default(),
635 None,
636 );
637 let scoped = ScopedContext {
638 connection_id: "b".into(),
639 database: "postgres".into(),
640 };
641 let _ = session
642 .checkout_for(CheckoutTarget::Scoped(&scoped))
643 .await;
644 let (active_id, _) = session.active_context().await;
645 assert_eq!(active_id, "a");
646 }
647
648 #[test]
649 fn index_stale_marker_round_trip() {
650 let session = ToolSession::for_tests(
651 vec![ConnectionInfo {
652 id: "local".into(),
653 name: "local".into(),
654 host: None,
655 port: None,
656 database: Some("postgres".into()),
657 params: ConnectionParams::default(),
658 policy: policy_from_profile(None, AccessMode::Read, PolicyCaps::default()),
659 }],
660 PolicyFilter::default(),
661 None,
662 );
663 assert!(!session.is_index_stale("local", "postgres"));
664 session.mark_index_stale("local", "postgres");
665 assert!(session.is_index_stale("local", "postgres"));
666 session.clear_index_stale("local", "postgres");
667 assert!(!session.is_index_stale("local", "postgres"));
668 }
669}