Skip to main content

systemprompt_database/admin/
introspection.rs

1//! Schema introspection service.
2//!
3//! Part of the documented sqlx allowlist — every query here is built
4//! dynamically because the table name is supplied at runtime as a
5//! [`SafeIdentifier`].
6//!
7//! Copyright (c) systemprompt.io — Business Source License 1.1.
8//! See <https://systemprompt.io> for licensing details.
9
10use std::sync::Arc;
11
12use sqlx::Row;
13use sqlx::postgres::PgPool;
14
15use crate::admin::identifier::SafeIdentifier;
16use crate::error::{DatabaseResult, RepositoryError, is_undefined_table};
17use crate::models::{ColumnInfo, DatabaseInfo, IndexInfo, TableInfo};
18
19#[derive(Debug)]
20pub struct DatabaseAdminService {
21    pool: Arc<PgPool>,
22}
23
24impl DatabaseAdminService {
25    pub const fn new(pool: Arc<PgPool>) -> Self {
26        Self { pool }
27    }
28
29    pub async fn list_tables(&self) -> DatabaseResult<Vec<TableInfo>> {
30        let rows = sqlx::query(
31            r"
32            SELECT
33                t.table_name as name,
34                COALESCE(s.n_live_tup, 0) as row_count,
35                COALESCE(pg_total_relation_size(to_regclass(quote_ident(t.table_name))), 0) as size_bytes
36            FROM information_schema.tables t
37            LEFT JOIN pg_stat_user_tables s ON t.table_name = s.relname
38            WHERE t.table_schema = 'public'
39            ORDER BY t.table_name
40            ",
41        )
42        .fetch_all(&*self.pool)
43        .await?;
44
45        let tables = rows
46            .iter()
47            .map(|row| {
48                let name: String = row.get("name");
49                let row_count: i64 = row.get("row_count");
50                let size_bytes: i64 = row.get("size_bytes");
51                TableInfo {
52                    name,
53                    row_count,
54                    size_bytes,
55                    columns: vec![],
56                }
57            })
58            .collect();
59
60        Ok(tables)
61    }
62
63    pub async fn describe_table(
64        &self,
65        table_name: &SafeIdentifier,
66    ) -> DatabaseResult<(Vec<ColumnInfo>, i64)> {
67        let rows = sqlx::query(
68            "SELECT column_name, data_type, is_nullable, column_default FROM \
69             information_schema.columns WHERE table_schema = 'public' AND table_name = $1 \
70             ORDER BY ordinal_position",
71        )
72        .bind(table_name.as_str())
73        .fetch_all(&*self.pool)
74        .await?;
75
76        if rows.is_empty() {
77            return Err(RepositoryError::not_found("table", table_name));
78        }
79
80        let pk_rows = sqlx::query(
81            r"
82            SELECT a.attname as column_name
83            FROM pg_index i
84            JOIN pg_attribute a ON a.attrelid = i.indrelid AND a.attnum = ANY(i.indkey)
85            WHERE i.indrelid = $1::regclass AND i.indisprimary
86            ",
87        )
88        .bind(table_name.as_str())
89        .fetch_all(&*self.pool)
90        .await?;
91
92        let pk_columns: Vec<String> = pk_rows
93            .iter()
94            .map(|row| row.get::<String, _>("column_name"))
95            .collect();
96
97        let columns = rows
98            .iter()
99            .map(|row| {
100                let name: String = row.get("column_name");
101                let data_type: String = row.get("data_type");
102                let nullable_str: String = row.get("is_nullable");
103                let nullable = nullable_str.to_uppercase() == "YES";
104                let default: Option<String> = row.get("column_default");
105                let primary_key = pk_columns.contains(&name);
106
107                ColumnInfo {
108                    name,
109                    data_type,
110                    nullable,
111                    primary_key,
112                    default,
113                }
114            })
115            .collect();
116
117        let row_count = self.count_rows(table_name).await?;
118
119        Ok((columns, row_count))
120    }
121
122    pub async fn list_table_indexes(
123        &self,
124        table_name: &SafeIdentifier,
125    ) -> DatabaseResult<Vec<IndexInfo>> {
126        let rows = sqlx::query(
127            r"
128            SELECT
129                i.relname as index_name,
130                ix.indisunique as is_unique,
131                array_agg(a.attname ORDER BY array_position(ix.indkey, a.attnum)) as columns
132            FROM pg_class t
133            JOIN pg_index ix ON t.oid = ix.indrelid
134            JOIN pg_class i ON i.oid = ix.indexrelid
135            JOIN pg_attribute a ON a.attrelid = t.oid AND a.attnum = ANY(ix.indkey)
136            WHERE t.relname = $1 AND t.relkind = 'r'
137            GROUP BY i.relname, ix.indisunique
138            ORDER BY i.relname
139            ",
140        )
141        .bind(table_name.as_str())
142        .fetch_all(&*self.pool)
143        .await?;
144
145        let indexes = rows
146            .iter()
147            .map(|row| {
148                let name: String = row.get("index_name");
149                let unique: bool = row.get("is_unique");
150                let columns: Vec<String> = row.get("columns");
151                IndexInfo {
152                    name,
153                    columns,
154                    unique,
155                }
156            })
157            .collect();
158
159        Ok(indexes)
160    }
161
162    pub async fn list_tables_counted(&self) -> DatabaseResult<Vec<TableInfo>> {
163        let mut counted = Vec::new();
164        for mut table in self.list_tables().await? {
165            let ident = SafeIdentifier::parse(&table.name)
166                .map_err(|e| RepositoryError::decode(format!("table name {}", table.name), e))?;
167            let count_query = format!("SELECT COUNT(*) as count FROM {}", ident.quoted());
168            match sqlx::query_scalar::<_, i64>(sqlx::AssertSqlSafe(count_query))
169                .fetch_one(&*self.pool)
170                .await
171            {
172                Ok(row_count) => table.row_count = row_count,
173                Err(e) if is_undefined_table(&e) => continue,
174                Err(e) => return Err(e.into()),
175            }
176            counted.push(table);
177        }
178        Ok(counted)
179    }
180
181    pub async fn count_rows(&self, table_name: &SafeIdentifier) -> DatabaseResult<i64> {
182        let quoted_table = table_name.quoted();
183        let count_query = format!("SELECT COUNT(*) as count FROM {quoted_table}");
184        let row_count: i64 = sqlx::query_scalar(sqlx::AssertSqlSafe(count_query))
185            .fetch_one(&*self.pool)
186            .await?;
187
188        Ok(row_count)
189    }
190
191    pub async fn get_database_info(&self) -> DatabaseResult<DatabaseInfo> {
192        let version: String = sqlx::query_scalar("SELECT version()")
193            .fetch_one(&*self.pool)
194            .await?;
195
196        let size: i64 = sqlx::query_scalar("SELECT pg_database_size(current_database())")
197            .fetch_one(&*self.pool)
198            .await?;
199
200        let size = u64::try_from(size)
201            .map_err(|e| RepositoryError::decode("pg_database_size(current_database())", e))?;
202
203        let tables = self.list_tables().await?;
204
205        Ok(DatabaseInfo {
206            path: "PostgreSQL".to_owned(),
207            size,
208            version,
209            tables,
210        })
211    }
212}