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};
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    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    // `to_url()` prefers `url` over `dbname`; drop stale URL path when host fields exist.
129    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    /// Register or update a profile in the live session (after save/import/setup).
156    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(&params, &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    /// Test helper: session with no live pools (index tools only).
368    #[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}