systemprompt_database/admin/
introspection.rs1use 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}