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;
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};
14use nexql_index::IndexStore;
15use nexql_policy::{AccessMode, PolicyCaps, PolicyFilter};
16use tokio::sync::RwLock as AsyncRwLock;
17
18use crate::error::ToolError;
19
20/// Index root: `NEXQL_MCP_INDEX_DIR`, else `~/.local/share/nexql-mcp` (same as CLI).
21pub fn default_index_root() -> PathBuf {
22    if let Ok(p) = std::env::var("NEXQL_MCP_INDEX_DIR") {
23        return PathBuf::from(p);
24    }
25    let home = std::env::var("HOME").unwrap_or_else(|_| ".".into());
26    PathBuf::from(home).join(".local/share/nexql-mcp")
27}
28
29/// Determine index root with workspace override support.
30/// Priority: explicit env var → project config index_dir → global default.
31pub fn resolve_index_root(
32    workspace_root: Option<&Path>,
33    project_config: Option<&nexql_conn::ProjectConfigFile>,
34) -> PathBuf {
35    if let Ok(p) = std::env::var("NEXQL_MCP_INDEX_DIR") {
36        return PathBuf::from(p);
37    }
38    if let (Some(root), Some(cfg)) = (workspace_root, project_config) {
39        if let Some(ref index_dir) = cfg.index_dir {
40            return root.join(".nexql").join(index_dir);
41        }
42    }
43    default_index_root()
44}
45
46#[derive(Debug, Clone)]
47pub struct ConnectionPolicy {
48    pub access_mode: AccessMode,
49    pub caps: PolicyCaps,
50    pub filter: PolicyFilter,
51    pub environment: Option<String>,
52}
53
54#[derive(Debug, Clone)]
55pub struct ConnectionInfo {
56    pub id: String,
57    pub name: String,
58    pub host: Option<String>,
59    pub port: Option<u16>,
60    pub database: Option<String>,
61    pub params: ConnectionParams,
62    pub policy: ConnectionPolicy,
63}
64
65#[derive(Debug, Clone)]
66struct ActivePolicy {
67    access_mode: AccessMode,
68    caps: PolicyCaps,
69    filter: PolicyFilter,
70    pool_opts: PoolOptions,
71}
72
73pub struct ToolSession {
74    pub connections: Vec<ConnectionInfo>,
75    policy: RwLock<ActivePolicy>,
76    /// Schema index root; `None` disables Phase 3 tools with an actionable error.
77    pub index_store: Option<IndexStore>,
78    inner: AsyncRwLock<SessionInner>,
79}
80
81fn filter_from_profile(profile: &ProfileConfig) -> PolicyFilter {
82    PolicyFilter {
83        allow_schemas: profile.schemas.clone(),
84        deny_schemas: profile.deny_schemas.clone(),
85        deny_tables: profile.deny_tables.clone(),
86        pii_columns: profile.pii_columns.clone(),
87    }
88}
89
90pub fn policy_from_profile(
91    profile: Option<&ProfileConfig>,
92    default_mode: AccessMode,
93    default_caps: PolicyCaps,
94) -> ConnectionPolicy {
95    let access_mode = profile
96        .and_then(|p| p.access_mode.as_deref())
97        .and_then(|m| m.parse::<AccessMode>().ok())
98        .unwrap_or(default_mode);
99    let mut caps = default_caps;
100    if let Some(n) = profile.and_then(|p| p.max_rows) {
101        caps = caps.with_max_rows(n);
102    }
103    let filter = profile.map(filter_from_profile).unwrap_or_default();
104    ConnectionPolicy {
105        access_mode,
106        caps,
107        filter,
108        environment: None,
109    }
110}
111
112struct SessionInner {
113    active_id: String,
114    database: String,
115    pools: HashMap<String, Pool>,
116}
117
118fn pool_key(connection_id: &str, database: &str) -> String {
119    format!("{connection_id}\0{database}")
120}
121
122fn params_for_database(base: &ConnectionParams, database: &str) -> ConnectionParams {
123    let mut params = base.clone();
124    params.dbname = Some(database.to_string());
125    // `to_url()` prefers `url` over `dbname`; drop stale URL path when host fields exist.
126    if params.host.is_some() {
127        params.url = None;
128    }
129    params
130}
131
132fn active_policy_from(policy: &ConnectionPolicy) -> ActivePolicy {
133    ActivePolicy {
134        access_mode: policy.access_mode,
135        caps: policy.caps.clone(),
136        filter: policy.filter.clone(),
137        pool_opts: PoolOptions {
138            read_only: !policy.access_mode.allows_writes(),
139            ..Default::default()
140        },
141    }
142}
143
144impl ToolSession {
145    pub fn access_mode(&self) -> AccessMode {
146        self.policy
147            .read()
148            .map(|p| p.access_mode)
149            .unwrap_or(AccessMode::Read)
150    }
151
152    pub fn caps(&self) -> PolicyCaps {
153        self.policy
154            .read()
155            .map(|p| p.caps.clone())
156            .unwrap_or_default()
157    }
158
159    pub fn filter(&self) -> PolicyFilter {
160        self.policy
161            .read()
162            .map(|p| p.filter.clone())
163            .unwrap_or_default()
164    }
165
166    pub fn pool_opts(&self) -> PoolOptions {
167        self.policy
168            .read()
169            .map(|p| p.pool_opts.clone())
170            .unwrap_or_default()
171    }
172
173    fn apply_policy(&self, policy: &ConnectionPolicy) {
174        if let Ok(mut active) = self.policy.write() {
175            *active = active_policy_from(policy);
176        }
177    }
178
179    pub async fn from_resolved(
180        resolved: ResolvedConnection,
181        access_mode: AccessMode,
182        caps: PolicyCaps,
183    ) -> Result<Arc<Self>, ToolError> {
184        let id = resolved
185            .profile_name
186            .clone()
187            .unwrap_or_else(|| "default".into());
188        let conn_policy = policy_from_profile(resolved.profile.as_ref(), access_mode, caps);
189        let info = ConnectionInfo {
190            id: id.clone(),
191            name: id.clone(),
192            host: resolved.params.host.clone(),
193            port: resolved.params.port,
194            database: resolved.params.dbname.clone(),
195            params: resolved.params,
196            policy: conn_policy,
197        };
198        Self::from_connections(vec![info], Some(id)).await
199    }
200
201    pub async fn from_connections(
202        connections: Vec<ConnectionInfo>,
203        active_id: Option<String>,
204    ) -> Result<Arc<Self>, ToolError> {
205        if connections.is_empty() {
206            return Err(ToolError::Execution("no connections configured".into()));
207        }
208        let active = active_id.unwrap_or_else(|| connections[0].id.clone());
209        let active_conn = connections
210            .iter()
211            .find(|c| c.id == active)
212            .unwrap_or(&connections[0]);
213        let active_policy = active_policy_from(&active_conn.policy);
214        let pool_opts = active_policy.pool_opts.clone();
215        let mut pools = HashMap::new();
216        for c in &connections {
217            let database = c.database.as_deref().unwrap_or("postgres");
218            let pool = create_pool(&c.params, &pool_opts).await?;
219            pools.insert(pool_key(&c.id, database), pool);
220        }
221        let database = active_conn
222            .database
223            .clone()
224            .unwrap_or_else(|| "postgres".into());
225        Ok(Arc::new(Self {
226            connections,
227            policy: RwLock::new(active_policy),
228            index_store: Some(IndexStore::new(default_index_root())),
229            inner: AsyncRwLock::new(SessionInner {
230                active_id: active,
231                database,
232                pools,
233            }),
234        }))
235    }
236
237    pub async fn from_connections_with_filter(
238        connections: Vec<ConnectionInfo>,
239        access_mode: AccessMode,
240        caps: PolicyCaps,
241        active_id: Option<String>,
242        filter: PolicyFilter,
243    ) -> Result<Arc<Self>, ToolError> {
244        let connections = connections
245            .into_iter()
246            .map(|mut c| {
247                c.policy.access_mode = access_mode;
248                c.policy.caps = caps.clone();
249                c.policy.filter = filter.clone();
250                c
251            })
252            .collect();
253        Self::from_connections(connections, active_id).await
254    }
255
256    pub async fn active_context(&self) -> (String, String) {
257        let g = self.inner.read().await;
258        (g.active_id.clone(), g.database.clone())
259    }
260
261    pub async fn switch(
262        &self,
263        connection_id: &str,
264        database: Option<String>,
265    ) -> Result<(), ToolError> {
266        let conn = self
267            .connections
268            .iter()
269            .find(|c| c.id == connection_id)
270            .ok_or_else(|| {
271                ToolError::Execution(format!(
272                    "Connection not found for ID: {connection_id} — call list_connections"
273                ))
274            })?;
275        let target_db = database
276            .or_else(|| conn.database.clone())
277            .unwrap_or_else(|| "postgres".into());
278        self.apply_policy(&conn.policy);
279        let pool_opts = self.pool_opts();
280        let key = pool_key(connection_id, &target_db);
281        let mut g = self.inner.write().await;
282        if let std::collections::hash_map::Entry::Vacant(e) = g.pools.entry(key) {
283            let params = params_for_database(&conn.params, &target_db);
284            let pool = create_pool(&params, &pool_opts).await?;
285            e.insert(pool);
286        }
287        g.active_id = connection_id.to_string();
288        g.database = target_db;
289        Ok(())
290    }
291
292    pub async fn checkout(&self) -> Result<Object, ToolError> {
293        let g = self.inner.read().await;
294        let key = pool_key(&g.active_id, &g.database);
295        let pool = g
296            .pools
297            .get(&key)
298            .ok_or_else(|| ToolError::Execution("no active pool".into()))?;
299        let pool_opts = self.pool_opts();
300        Ok(checkout_guarded(pool, &pool_opts).await?)
301    }
302
303    /// Test helper: session with no live pools (index tools only).
304    #[cfg(test)]
305    pub fn for_tests(
306        connections: Vec<ConnectionInfo>,
307        filter: PolicyFilter,
308        index_store: Option<IndexStore>,
309    ) -> Arc<Self> {
310        assert!(!connections.is_empty());
311        let active = connections[0].id.clone();
312        let database = connections[0]
313            .database
314            .clone()
315            .unwrap_or_else(|| "postgres".into());
316        let policy = ConnectionPolicy {
317            access_mode: AccessMode::Read,
318            caps: PolicyCaps::default(),
319            filter,
320            environment: None,
321        };
322        Arc::new(Self {
323            connections,
324            policy: RwLock::new(active_policy_from(&policy)),
325            index_store,
326            inner: AsyncRwLock::new(SessionInner {
327                active_id: active,
328                database,
329                pools: HashMap::new(),
330            }),
331        })
332    }
333}
334
335#[cfg(test)]
336mod tests {
337    use super::*;
338
339    #[test]
340    fn filter_maps_profile_fields() {
341        let profile = ProfileConfig {
342            schemas: vec!["public".into()],
343            deny_schemas: vec!["pgboss".into()],
344            deny_tables: vec!["auth.*".into()],
345            pii_columns: vec!["public.users.ssn".into()],
346            ..Default::default()
347        };
348        let f = filter_from_profile(&profile);
349        assert_eq!(f.allow_schemas, vec!["public"]);
350        assert_eq!(f.deny_schemas, vec!["pgboss"]);
351        assert_eq!(f.deny_tables, vec!["auth.*"]);
352        assert_eq!(f.pii_columns, vec!["public.users.ssn"]);
353        assert!(f.allows_schema("public"));
354        assert!(!f.allows_schema("pgboss"));
355        assert!(!f.allows_table("auth", "sessions"));
356    }
357}