systemprompt_cli/commands/admin/setup/
common.rs1use anyhow::Result;
2use rand::distr::Alphanumeric;
3use rand::{RngExt, rng};
4use sqlx::postgres::PgPoolOptions;
5use std::net::ToSocketAddrs;
6use std::time::Duration;
7use systemprompt_logging::CliService;
8
9#[derive(Debug, Clone)]
10pub struct PostgresConfig {
11 pub host: String,
12 pub port: u16,
13 pub user: String,
14 pub password: String,
15 pub database: String,
16}
17
18impl PostgresConfig {
19 pub fn database_url(&self) -> String {
20 format!(
21 "postgres://{}:{}@{}:{}/{}",
22 self.user, self.password, self.host, self.port, self.database
23 )
24 }
25}
26
27pub fn generate_password() -> String {
28 let mut rng = rng();
29 (0..16)
30 .map(|_| rng.sample(Alphanumeric))
31 .map(char::from)
32 .collect()
33}
34
35pub fn detect_postgresql(host: &str, port: u16) -> bool {
36 let addr = format!("{}:{}", host, port);
37 let socket_addrs = match addr.to_socket_addrs() {
38 Ok(addrs) => addrs.collect::<Vec<_>>(),
39 Err(e) => {
40 tracing::debug!(host = %host, port = %port, error = %e, "Failed to resolve socket address");
41 return false;
42 },
43 };
44
45 for socket_addr in socket_addrs {
46 if std::net::TcpStream::connect_timeout(&socket_addr, Duration::from_secs(3)).is_ok() {
47 return true;
48 }
49 }
50
51 false
52}
53
54pub async fn test_connection(config: &PostgresConfig) -> bool {
55 let Ok(pool) = PgPoolOptions::new()
56 .max_connections(1)
57 .acquire_timeout(Duration::from_secs(5))
58 .connect(&config.database_url())
59 .await
60 else {
61 return false;
62 };
63
64 let result = sqlx::query_scalar!("SELECT 1 as one")
65 .fetch_one(&pool)
66 .await
67 .is_ok();
68 pool.close().await;
69 result
70}
71
72pub async fn enable_extensions(config: &PostgresConfig) -> Result<()> {
73 let pool = match PgPoolOptions::new()
74 .max_connections(1)
75 .acquire_timeout(Duration::from_secs(5))
76 .connect(&config.database_url())
77 .await
78 {
79 Ok(pool) => pool,
80 Err(e) => {
81 CliService::warning(&format!("Could not enable extensions: {}", e));
82 return Ok(());
83 },
84 };
85
86 let extensions = ["uuid-ossp", "unaccent", "pg_trgm"];
87
88 for ext in extensions {
89 let sql = format!("CREATE EXTENSION IF NOT EXISTS \"{}\"", ext);
90 if let Err(e) = sqlx::query(sqlx::AssertSqlSafe(sql)).execute(&pool).await {
91 CliService::warning(&format!("Could not create extension '{}': {}", ext, e));
92 }
93 }
94
95 pool.close().await;
96 Ok(())
97}