Skip to main content

systemprompt_cli/commands/admin/setup/
common.rs

1use 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}