use crate::error::PostgisError;
use crate::geometry::Geometry;
use crate::postgis::{PostgisExt, RealPgConfig};
use async_trait::async_trait;
use tokio::sync::OnceCell;
use tokio_postgres::Client;
pub struct RealPostgis {
config: RealPgConfig,
client: OnceCell<Client>,
}
impl RealPostgis {
pub fn new(config: RealPgConfig) -> Result<Self, PostgisError> {
Ok(Self {
config,
client: OnceCell::new(),
})
}
async fn client(&self) -> Result<&Client, PostgisError> {
self.client
.get_or_try_init(|| async {
let conn_str = format!(
"host={} port={} dbname={} user={} password={}",
self.config.host,
self.config.port,
self.config.database,
self.config.username,
self.config.password
);
let (client, connection) =
tokio_postgres::connect(&conn_str, tokio_postgres::NoTls)
.await
.map_err(|e| PostgisError::Connection(e.to_string()))?;
tokio::spawn(async move {
if let Err(e) = connection.await {
eprintln!("[sz-orm-postgis] postgres connection error: {}", e);
}
});
Ok::<Client, PostgisError>(client)
})
.await
}
async fn query_f64(
&self,
sql: &str,
params: &[&(dyn tokio_postgres::types::ToSql + Sync)],
) -> Result<f64, PostgisError> {
let client = self.client().await?;
let row = client
.query_one(sql, params)
.await
.map_err(|e| PostgisError::Query(e.to_string()))?;
let v: f64 = row
.try_get(0)
.map_err(|e| PostgisError::Query(format!("type conversion failed: {}", e)))?;
Ok(v)
}
async fn query_bool(
&self,
sql: &str,
params: &[&(dyn tokio_postgres::types::ToSql + Sync)],
) -> Result<bool, PostgisError> {
let client = self.client().await?;
let row = client
.query_one(sql, params)
.await
.map_err(|e| PostgisError::Query(e.to_string()))?;
let v: bool = row
.try_get(0)
.map_err(|e| PostgisError::Query(format!("type conversion failed: {}", e)))?;
Ok(v)
}
async fn query_ewkt(
&self,
sql: &str,
params: &[&(dyn tokio_postgres::types::ToSql + Sync)],
) -> Result<Geometry, PostgisError> {
let client = self.client().await?;
let row = client
.query_one(sql, params)
.await
.map_err(|e| PostgisError::Query(e.to_string()))?;
let ewkt: String = row
.try_get(0)
.map_err(|e| PostgisError::Query(format!("conversion failed: {}", e)))?;
Geometry::from_ewkt(&ewkt)
}
async fn execute_ddl(
&self,
sql: &str,
params: &[&(dyn tokio_postgres::types::ToSql + Sync)],
) -> Result<(), PostgisError> {
let client = self.client().await?;
client
.execute(sql, params)
.await
.map_err(|e| PostgisError::Query(e.to_string()))?;
Ok(())
}
}
fn validate_identifier(name: &str, kind: &str) -> Result<(), PostgisError> {
if name.is_empty() || name.len() > 63 {
return Err(PostgisError::Query(format!(
"invalid {}: empty or too long (max 63 chars): {}",
kind, name
)));
}
let mut chars = name.chars();
let first = chars.next().unwrap();
if !first.is_ascii_alphabetic() && first != '_' {
return Err(PostgisError::Query(format!(
"invalid {}: must start with letter or underscore, got '{}'",
kind, name
)));
}
if !chars.all(|c| c.is_ascii_alphanumeric() || c == '_') {
return Err(PostgisError::Query(format!(
"invalid {}: only alphanumeric and underscore allowed, got '{}'",
kind, name
)));
}
Ok(())
}
fn validate_geom_type(geom_type: &str) -> Result<(), PostgisError> {
let allowed = [
"POINT",
"LINESTRING",
"POLYGON",
"MULTIPOINT",
"MULTILINESTRING",
"MULTIPOLYGON",
"GEOMETRYCOLLECTION",
"GEOMETRY",
];
if !allowed.contains(&geom_type.to_uppercase().as_str()) {
return Err(PostgisError::Query(format!(
"invalid geom_type: {}, allowed: {:?}",
geom_type, allowed
)));
}
Ok(())
}
fn validate_dim(dim: &str) -> Result<(), PostgisError> {
if !["2", "3", "4"].contains(&dim) {
return Err(PostgisError::Query(format!(
"invalid dim: {}, allowed: 2/3/4",
dim
)));
}
Ok(())
}
#[async_trait]
impl PostgisExt for RealPostgis {
async fn st_distance(&self, g1: &Geometry, g2: &Geometry) -> Result<f64, PostgisError> {
let sql = "SELECT ST_Distance($1::geometry, $2::geometry)";
self.query_f64(sql, &[&g1.to_ewkt(), &g2.to_ewkt()]).await
}
async fn st_contains(&self, outer: &Geometry, inner: &Geometry) -> Result<bool, PostgisError> {
let sql = "SELECT ST_Contains($1::geometry, $2::geometry)";
self.query_bool(sql, &[&outer.to_ewkt(), &inner.to_ewkt()])
.await
}
async fn st_within(&self, inner: &Geometry, outer: &Geometry) -> Result<bool, PostgisError> {
let sql = "SELECT ST_Within($1::geometry, $2::geometry)";
self.query_bool(sql, &[&inner.to_ewkt(), &outer.to_ewkt()])
.await
}
async fn st_intersects(&self, g1: &Geometry, g2: &Geometry) -> Result<bool, PostgisError> {
let sql = "SELECT ST_Intersects($1::geometry, $2::geometry)";
self.query_bool(sql, &[&g1.to_ewkt(), &g2.to_ewkt()]).await
}
async fn st_area(&self, geom: &Geometry) -> Result<f64, PostgisError> {
let sql = "SELECT ST_Area($1::geometry)";
self.query_f64(sql, &[&geom.to_ewkt()]).await
}
async fn st_length(&self, geom: &Geometry) -> Result<f64, PostgisError> {
let sql = "SELECT ST_Length($1::geometry)";
self.query_f64(sql, &[&geom.to_ewkt()]).await
}
async fn st_buffer(&self, geom: &Geometry, distance: f64) -> Result<Geometry, PostgisError> {
let sql = "SELECT ST_AsEWKT(ST_Buffer($1::geometry, $2))";
self.query_ewkt(sql, &[&geom.to_ewkt(), &distance]).await
}
async fn st_union(&self, g1: &Geometry, g2: &Geometry) -> Result<Geometry, PostgisError> {
let sql = "SELECT ST_AsEWKT(ST_Union($1::geometry, $2::geometry))";
self.query_ewkt(sql, &[&g1.to_ewkt(), &g2.to_ewkt()]).await
}
async fn add_geometry_column(
&self,
table: &str,
column: &str,
srid: i32,
geom_type: &str,
dim: &str,
) -> Result<(), PostgisError> {
validate_identifier(table, "table")?;
validate_identifier(column, "column")?;
validate_geom_type(geom_type)?;
validate_dim(dim)?;
let sql = format!(
"SELECT AddGeometryColumn('{}', '{}', $1, '{}', '{}')",
table,
column,
geom_type.to_uppercase(),
dim
);
self.execute_ddl(&sql, &[&srid]).await
}
async fn create_spatial_index(&self, table: &str, column: &str) -> Result<(), PostgisError> {
validate_identifier(table, "table")?;
validate_identifier(column, "column")?;
let idx_name = format!("idx_{}_{}", table, column);
let sql = format!(
"CREATE INDEX IF NOT EXISTS {} ON {} USING GIST(\"{}\")",
idx_name, table, column
);
self.execute_ddl(&sql, &[]).await
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_validate_identifier_valid() {
assert!(validate_identifier("users", "table").is_ok());
assert!(validate_identifier("_idx", "index").is_ok());
assert!(validate_identifier("geom_2026", "column").is_ok());
}
#[test]
fn test_validate_identifier_invalid() {
assert!(validate_identifier("users; DROP TABLE", "table").is_err());
assert!(validate_identifier("col'--", "column").is_err());
assert!(validate_identifier("col\"x", "column").is_err());
assert!(validate_identifier("1col", "column").is_err());
assert!(validate_identifier("", "table").is_err());
let long_name = "a".repeat(64);
assert!(validate_identifier(&long_name, "table").is_err());
}
#[test]
fn test_validate_geom_type() {
assert!(validate_geom_type("POINT").is_ok());
assert!(validate_geom_type("point").is_ok()); assert!(validate_geom_type("POLYGON").is_ok());
assert!(validate_geom_type("GEOMETRY").is_ok());
assert!(validate_geom_type("EVIL_TYPE").is_err());
assert!(validate_geom_type("POINT'; DROP TABLE").is_err());
}
#[test]
fn test_validate_dim() {
assert!(validate_dim("2").is_ok());
assert!(validate_dim("3").is_ok());
assert!(validate_dim("4").is_ok());
assert!(validate_dim("5").is_err());
assert!(validate_dim("'; DROP TABLE").is_err());
}
#[test]
fn test_realpostgis_new_does_not_connect() {
let config = RealPgConfig {
host: "nonexistent.invalid".to_string(),
port: 5432,
database: "test".to_string(),
username: "postgres".to_string(),
password: "secret".to_string(),
};
let _real = RealPostgis::new(config).expect("new() should not connect");
}
}