Skip to main content

nexql_tools/
session.rs

1// SPDX-License-Identifier: GPL-3.0-only
2// Copyright (C) 2026 NexQL-OSS Team
3
4//! Session state: resolved profiles + active pool + optional schema index.
5
6use 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
21/// Index root: `NEXQL_MCP_INDEX_DIR`, else `~/.local/share/nexql-mcp` (same as CLI).
22pub 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
30/// Determine index root with workspace override support.
31/// Priority: explicit env var → project config index_dir → global default.
32pub 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    /// `connection_id\0database` keys marked stale after DDL until index rebuild/refresh.
77    stale_index_keys: RwLock<HashSet<String>>,
78    policy: RwLock<ActivePolicy>,
79    /// Schema index root; `None` disables Phase 3 tools with an actionable error.
80    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    // `to_url()` prefers `url` over `dbname`; drop stale URL path when host fields exist.
132    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/// Target connection/database for a single tool call without mutating session context.
154#[derive(Debug, Clone, PartialEq, Eq)]
155pub struct ScopedContext {
156    pub connection_id: String,
157    pub database: String,
158}
159
160/// Checkout target: active session or an explicit scoped context.
161pub 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    /// Register or update a profile in the live session (after save/import/setup).
175    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(&params, &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    /// Resolve effective connection/database for a tool call.
420    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(&params, &pool_opts).await?;
476        g.pools.insert(key, pool);
477        Ok(())
478    }
479
480    /// Checkout a client for the given target without changing active session context.
481    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    /// Test helper: session with no live pools (index tools only).
524    #[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}