diff --git a/backend/Cargo.lock b/backend/Cargo.lock index abc735f0c7..271ec93546 100644 --- a/backend/Cargo.lock +++ b/backend/Cargo.lock @@ -16773,6 +16773,7 @@ dependencies = [ "pem", "pep440_rs", "postgres-native-tls 0.5.3", + "postgres-protocol", "process-wrap", "prometheus", "prost", diff --git a/backend/windmill-worker/Cargo.toml b/backend/windmill-worker/Cargo.toml index 5e95a1e20a..57b292f8f2 100644 --- a/backend/windmill-worker/Cargo.toml +++ b/backend/windmill-worker/Cargo.toml @@ -101,6 +101,7 @@ mappable-rc.workspace = true git-version.workspace = true once_cell.workspace = true tokio-postgres.workspace = true +postgres-protocol.workspace = true bit-vec.workspace = true url.workspace = true async-stream.workspace = true diff --git a/backend/windmill-worker/src/pg_executor.rs b/backend/windmill-worker/src/pg_executor.rs index 69aa3ee0fb..22a59f82e6 100644 --- a/backend/windmill-worker/src/pg_executor.rs +++ b/backend/windmill-worker/src/pg_executor.rs @@ -1,5 +1,4 @@ use std::collections::{HashMap, HashSet}; -use std::net::IpAddr; use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering}; use std::sync::Arc; use std::time::Duration; @@ -1889,8 +1888,8 @@ pub fn pg_cell_to_json_value_with_state( Type::UUID => get_basic(row, column, column_i, |a: uuid::Uuid| { Ok(JSONValue::String(a.to_string())) })?, - Type::INET => get_basic(row, column, column_i, |a: IpAddr| { - Ok(JSONValue::String(a.to_string())) + Type::INET | Type::CIDR => get_basic(row, column, column_i, |a: InetStr| { + Ok(JSONValue::String(a.0)) })?, Type::INTERVAL => get_basic(row, column, column_i, |a: IntervalStr| { Ok(JSONValue::String(a.0)) @@ -1997,6 +1996,9 @@ pub fn pg_cell_to_json_value_with_state( Type::BYTEA_ARRAY => get_array(row, column, column_i, |a: Vec| { Ok(JSONValue::String(format!("\\x{}", hex::encode(a)))) })?, + Type::INET_ARRAY | Type::CIDR_ARRAY => get_array(row, column, column_i, |a: InetStr| { + Ok(JSONValue::String(a.0)) + })?, Type::VOID => JSONValue::Null, // Default fallback for unhandled column types: read as text. We use // `AnyTextValue` instead of plain `String` so that `Kind::Enum`, @@ -2156,6 +2158,29 @@ fn get_array<'a, T: FromSql<'a>>( }) } +/// `inet` / `cidr` in Postgres's own text form. `IpAddr`'s `FromSql` drops the +/// prefix length, so a network such as `10.0.0.0/8` would read back as a host. +struct InetStr(String); +impl FromSql<'_> for InetStr { + fn from_sql( + ty: &Type, + raw: &[u8], + ) -> Result> { + let inet = postgres_protocol::types::inet_from_sql(raw)?; + let host_prefix = if inet.addr().is_ipv4() { 32 } else { 128 }; + Ok(InetStr( + if *ty == Type::CIDR || inet.netmask() != host_prefix { + format!("{}/{}", inet.addr(), inet.netmask()) + } else { + inet.addr().to_string() + }, + )) + } + fn accepts(ty: &Type) -> bool { + matches!(*ty, Type::INET | Type::CIDR) + } +} + // you can remove this section if not using TS_VECTOR (or other types requiring an intermediary `FromSQL` struct) struct StringCollector(String); impl FromSql<'_> for StringCollector { @@ -2271,6 +2296,19 @@ mod tests { ); } + #[test] + fn inet_keeps_the_prefix_postgres_prints() { + let inet = |netmask: u8, addr: [u8; 4]| { + let mut raw = vec![2, netmask, 0, 4]; + raw.extend(addr); + raw + }; + let read = |ty: &Type, raw: &[u8]| InetStr::from_sql(ty, raw).unwrap().0; + assert_eq!(read(&Type::INET, &inet(32, [10, 1, 2, 3])), "10.1.2.3"); + assert_eq!(read(&Type::INET, &inet(8, [10, 0, 0, 0])), "10.0.0.0/8"); + assert_eq!(read(&Type::CIDR, &inet(32, [10, 1, 2, 3])), "10.1.2.3/32"); + } + #[test] fn test_parse_naive_date() { // chrono's NaiveDate::to_string() format