use crate::config::PostgresSourceConfig;
use async_trait::async_trait;
use faucet_core::shard::ShardSpec;
use faucet_core::util::quote_ident;
use faucet_core::{FaucetError, Stream, StreamPage};
use futures::TryStreamExt;
use serde_json::Value;
use sqlx::postgres::PgPoolOptions;
use sqlx::{Column, PgPool, Row};
use std::pin::Pin;
use std::sync::Mutex;
pub struct PostgresSource {
config: PostgresSourceConfig,
pool: PgPool,
applied_shard: Mutex<Option<ShardBounds>>,
}
#[derive(Clone, Debug)]
struct ShardBounds {
key: String,
lo: i64,
hi: i64,
lo_unbounded: bool,
hi_unbounded: bool,
include_null: bool,
}
impl ShardBounds {
fn from_spec(spec: &ShardSpec) -> Option<Self> {
let d = &spec.descriptor;
Some(Self {
key: d.get("key")?.as_str()?.to_string(),
lo: d.get("lo")?.as_i64()?,
hi: d.get("hi")?.as_i64()?,
lo_unbounded: d
.get("lo_unbounded")
.and_then(Value::as_bool)
.unwrap_or(false),
hi_unbounded: d
.get("hi_unbounded")
.and_then(Value::as_bool)
.unwrap_or(false),
include_null: d
.get("include_null")
.and_then(Value::as_bool)
.unwrap_or(false),
})
}
fn wrap(&self, inner: &str) -> String {
let key = quote_ident(&self.key);
let mut parts: Vec<String> = Vec::with_capacity(2);
if !self.lo_unbounded {
parts.push(format!("{key} >= {lo}", lo = self.lo));
}
if !self.hi_unbounded {
parts.push(format!("{key} < {hi}", hi = self.hi));
}
let range = parts.join(" AND ");
let predicate = if self.include_null {
if range.is_empty() {
"TRUE".to_string()
} else {
format!("(({range}) OR {key} IS NULL)")
}
} else if range.is_empty() {
"TRUE".to_string()
} else {
range
};
format!("SELECT * FROM ({inner}) AS _faucet_shard WHERE {predicate}")
}
}
impl PostgresSource {
pub async fn new(config: PostgresSourceConfig) -> Result<Self, FaucetError> {
faucet_core::validate_batch_size(config.batch_size)?;
let pool = PgPoolOptions::new()
.max_connections(config.max_connections)
.connect(&config.connection_url)
.await
.map_err(|e| FaucetError::Config(format!("PostgreSQL connection failed: {e}")))?;
Ok(Self {
config,
pool,
applied_shard: Mutex::new(None),
})
}
fn shard_wrap(&self, query: String) -> String {
match &*self.applied_shard.lock().expect("shard mutex poisoned") {
Some(bounds) => bounds.wrap(&query),
None => query,
}
}
}
fn plan_pk_shards(key: &str, min: i64, max: i64, target: usize) -> Vec<ShardSpec> {
let target = target.max(1);
let width = (max as i128 - min as i128 + 1).max(1) as u128;
let n = (target as u128).min(width) as usize; let step = width.div_ceil(n as u128);
let mut shards = Vec::with_capacity(n);
let mut lo = min as i128;
for i in 0..n {
let mut hi = lo + step as i128;
let is_first = i == 0;
let is_last = i == n - 1;
if is_last || hi > max as i128 {
hi = max as i128; }
let descriptor = serde_json::json!({
"key": key,
"lo": lo as i64,
"hi": hi as i64,
"lo_unbounded": is_first,
"hi_unbounded": is_last,
"include_null": is_last,
});
let size = (hi - lo).max(0) as u64 + if is_last { 1 } else { 0 };
shards.push(ShardSpec::new(i.to_string(), descriptor).with_size(size));
if is_last {
break;
}
lo = hi;
}
shards
}
fn pg_value_to_json(row: &sqlx::postgres::PgRow, col_name: &str) -> Value {
if let Ok(v) = row.try_get::<Value, _>(col_name) {
return v;
}
if let Ok(v) = row.try_get::<String, _>(col_name) {
return Value::String(v);
}
if let Ok(v) = row.try_get::<i64, _>(col_name) {
return Value::Number(v.into());
}
if let Ok(v) = row.try_get::<i32, _>(col_name) {
return Value::Number(v.into());
}
if let Ok(v) = row.try_get::<i16, _>(col_name) {
return Value::Number(v.into());
}
if let Ok(v) = row.try_get::<f64, _>(col_name) {
return serde_json::Number::from_f64(v)
.map(Value::Number)
.unwrap_or(Value::Null);
}
if let Ok(v) = row.try_get::<f32, _>(col_name) {
return serde_json::Number::from_f64(v as f64)
.map(Value::Number)
.unwrap_or(Value::Null);
}
if let Ok(v) = row.try_get::<bool, _>(col_name) {
return Value::Bool(v);
}
if let Ok(v) =
row.try_get::<sqlx::types::chrono::DateTime<sqlx::types::chrono::Utc>, _>(col_name)
{
return Value::String(v.to_rfc3339());
}
if let Ok(v) = row.try_get::<sqlx::types::chrono::NaiveDateTime, _>(col_name) {
return Value::String(v.to_string());
}
if let Ok(v) = row.try_get::<sqlx::types::chrono::NaiveDate, _>(col_name) {
return Value::String(v.to_string());
}
if let Ok(v) = row.try_get::<sqlx::types::chrono::NaiveTime, _>(col_name) {
return Value::String(v.to_string());
}
if let Ok(v) = row.try_get::<sqlx::types::Uuid, _>(col_name) {
return Value::String(v.to_string());
}
if let Ok(v) = row.try_get::<sqlx::types::BigDecimal, _>(col_name) {
return Value::String(v.to_string());
}
if let Ok(v) = row.try_get::<Vec<u8>, _>(col_name) {
use base64::Engine as _;
return Value::String(base64::engine::general_purpose::STANDARD.encode(v));
}
Value::Null
}
fn resolve_query(
config: &PostgresSourceConfig,
context: &std::collections::HashMap<String, Value>,
) -> (String, Vec<Value>) {
if context.is_empty() {
(config.query.clone(), Vec::new())
} else {
faucet_core::util::substitute_context_bind_params(
&config.query,
context,
config.params.len() + 1,
|i| format!("${i}"),
)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum NumberBind {
I64,
U64,
F64,
}
fn classify_number(n: &serde_json::Number) -> NumberBind {
if n.is_i64() {
NumberBind::I64
} else if n.is_u64() {
NumberBind::U64
} else {
NumberBind::F64
}
}
fn bind_params<'q>(
mut query: sqlx::query::Query<'q, sqlx::Postgres, sqlx::postgres::PgArguments>,
config_params: &'q [Value],
bind_values: &'q [Value],
) -> sqlx::query::Query<'q, sqlx::Postgres, sqlx::postgres::PgArguments> {
for value in config_params.iter().chain(bind_values) {
query = match value {
Value::String(s) => query.bind(s.clone()),
Value::Number(n) => match classify_number(n) {
NumberBind::I64 => query.bind(n.as_i64().unwrap()),
NumberBind::U64 => query.bind(n.as_u64().unwrap() as i64),
NumberBind::F64 => query.bind(n.as_f64().unwrap_or(0.0)),
},
Value::Bool(b) => query.bind(*b),
Value::Null => query.bind(None::<String>),
_ => query.bind(value.to_string()),
};
}
query
}
fn row_to_json(row: &sqlx::postgres::PgRow) -> Value {
let mut map = serde_json::Map::new();
for col in row.columns() {
let name = col.name().to_string();
let value = pg_value_to_json(row, &name);
map.insert(name, value);
}
Value::Object(map)
}
#[async_trait]
impl faucet_core::Source for PostgresSource {
async fn fetch_with_context(
&self,
context: &std::collections::HashMap<String, serde_json::Value>,
) -> Result<Vec<Value>, FaucetError> {
let (query_str, bind_values) = resolve_query(&self.config, context);
let query_str = self.shard_wrap(query_str);
let query = bind_params(sqlx::query(&query_str), &self.config.params, &bind_values);
let rows = query
.fetch_all(&self.pool)
.await
.map_err(|e| FaucetError::Config(format!("PostgreSQL query failed: {e}")))?;
let records: Vec<Value> = rows.iter().map(row_to_json).collect();
tracing::info!(rows = records.len(), query = %self.config.query, "PostgreSQL source fetch complete");
Ok(records)
}
fn stream_pages<'a>(
&'a self,
context: &'a std::collections::HashMap<String, Value>,
_batch_size: usize,
) -> Pin<Box<dyn Stream<Item = Result<StreamPage, FaucetError>> + Send + 'a>> {
let batch_size = self.config.batch_size;
Box::pin(async_stream::try_stream! {
let (query_str, bind_values) = resolve_query(&self.config, context);
let query_str = self.shard_wrap(query_str);
let query = bind_params(
sqlx::query(&query_str),
&self.config.params,
&bind_values,
);
let mut rows = query.fetch(&self.pool);
let chunk = if batch_size == 0 { usize::MAX } else { batch_size };
let initial_capacity = if batch_size == 0 { 1024 } else { batch_size };
let mut buffer: Vec<Value> = Vec::with_capacity(initial_capacity);
let mut total = 0usize;
while let Some(row) = rows
.try_next()
.await
.map_err(|e| FaucetError::Config(format!("PostgreSQL query failed: {e}")))?
{
buffer.push(row_to_json(&row));
if buffer.len() >= chunk {
let page = std::mem::replace(&mut buffer, Vec::with_capacity(initial_capacity));
total += page.len();
yield StreamPage { records: page, bookmark: None };
}
}
if !buffer.is_empty() {
total += buffer.len();
yield StreamPage { records: buffer, bookmark: None };
}
tracing::info!(
rows = total,
batch_size,
query = %self.config.query,
"PostgreSQL source stream complete",
);
})
}
fn config_schema(&self) -> serde_json::Value {
serde_json::to_value(faucet_core::schema_for!(PostgresSourceConfig))
.expect("schema serialization")
}
fn dataset_uri(&self) -> String {
format!(
"{}?query={}",
faucet_core::redact_uri_credentials(&self.config.connection_url),
self.config.query
)
}
fn is_shardable(&self) -> bool {
self.config.shard.is_some()
}
async fn enumerate_shards(&self, target: usize) -> Result<Vec<ShardSpec>, FaucetError> {
let Some(shard_cfg) = &self.config.shard else {
return Ok(vec![ShardSpec::whole()]);
};
let key = quote_ident(&shard_cfg.key);
let bounds_sql = format!(
"SELECT MIN({key})::int8 AS lo, MAX({key})::int8 AS hi \
FROM ({inner}) AS _faucet_bounds",
inner = self.config.query
);
let row = bind_params(sqlx::query(&bounds_sql), &self.config.params, &[])
.fetch_one(&self.pool)
.await
.map_err(|e| {
FaucetError::Source(format!(
"postgres: failed to compute shard bounds for key {:?} \
(it must be an integer-typed column): {e}",
shard_cfg.key
))
})?;
let lo: Option<i64> = row.try_get("lo").map_err(|e| {
FaucetError::Source(format!("postgres: shard bounds decode failed: {e}"))
})?;
let hi: Option<i64> = row.try_get("hi").map_err(|e| {
FaucetError::Source(format!("postgres: shard bounds decode failed: {e}"))
})?;
match (lo, hi) {
(Some(lo), Some(hi)) => Ok(plan_pk_shards(&shard_cfg.key, lo, hi, target)),
_ => Ok(vec![ShardSpec::whole()]),
}
}
async fn apply_shard(&self, shard: &ShardSpec) -> Result<(), FaucetError> {
let bounds = if shard.is_whole() {
None
} else {
Some(ShardBounds::from_spec(shard).ok_or_else(|| {
FaucetError::Source(format!(
"postgres: invalid shard descriptor: {}",
shard.descriptor
))
})?)
};
*self.applied_shard.lock().expect("shard mutex poisoned") = bounds;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn new_rejects_out_of_range_batch_size() {
let mut config = PostgresSourceConfig::new("postgres://localhost/test", "SELECT 1");
config.batch_size = faucet_core::MAX_BATCH_SIZE + 1;
match PostgresSource::new(config).await {
Err(faucet_core::FaucetError::Config(m)) => {
assert!(m.contains("batch_size"), "got: {m}")
}
_ => panic!("expected a batch_size Config error"),
}
}
fn num(v: serde_json::Value) -> serde_json::Number {
match v {
serde_json::Value::Number(n) => n,
_ => panic!("not a number"),
}
}
#[test]
fn classify_small_int_is_i64() {
assert_eq!(
classify_number(&num(serde_json::json!(42))),
NumberBind::I64
);
assert_eq!(
classify_number(&num(serde_json::json!(-7))),
NumberBind::I64
);
assert_eq!(classify_number(&num(serde_json::json!(0))), NumberBind::I64);
}
#[test]
fn classify_above_2_pow_53_stays_i64_not_f64() {
let v = 9_007_199_254_740_993i64; assert_eq!(classify_number(&num(serde_json::json!(v))), NumberBind::I64);
}
#[test]
fn classify_i64_max_is_i64() {
assert_eq!(
classify_number(&num(serde_json::json!(i64::MAX))),
NumberBind::I64
);
assert_eq!(
classify_number(&num(serde_json::json!(i64::MIN))),
NumberBind::I64
);
}
#[test]
fn classify_above_i64_max_is_u64() {
let v: u64 = i64::MAX as u64 + 1;
assert_eq!(classify_number(&num(serde_json::json!(v))), NumberBind::U64);
assert_eq!(
classify_number(&num(serde_json::json!(u64::MAX))),
NumberBind::U64
);
}
#[test]
fn classify_float_is_f64() {
assert_eq!(
classify_number(&num(serde_json::json!(3.5))),
NumberBind::F64
);
assert_eq!(
classify_number(&num(serde_json::json!(-0.5))),
NumberBind::F64
);
}
#[test]
fn plan_pk_shards_covers_full_range_without_gaps_or_overlap() {
let shards = plan_pk_shards("id", 0, 99, 4);
assert_eq!(shards.len(), 4);
let mut expected_lo = 0i64;
for (i, s) in shards.iter().enumerate() {
let d = &s.descriptor;
assert_eq!(d["key"], "id");
assert_eq!(d["lo"].as_i64().unwrap(), expected_lo);
let hi = d["hi"].as_i64().unwrap();
let first = i == 0;
let last = i == shards.len() - 1;
assert_eq!(d["lo_unbounded"].as_bool().unwrap(), first);
assert_eq!(d["hi_unbounded"].as_bool().unwrap(), last);
expected_lo = hi; }
}
#[test]
fn plan_pk_shards_never_more_shards_than_values() {
let shards = plan_pk_shards("pk", 5, 7, 10);
assert!(shards.len() <= 3, "got {} shards", shards.len());
assert!(
shards[0].descriptor["lo_unbounded"].as_bool().unwrap(),
"first shard is unbounded below"
);
assert!(
shards.last().unwrap().descriptor["hi_unbounded"]
.as_bool()
.unwrap(),
"last shard is unbounded above"
);
}
#[test]
fn plan_pk_shards_single_value_one_shard() {
let shards = plan_pk_shards("id", 42, 42, 8);
assert_eq!(shards.len(), 1);
assert!(shards[0].descriptor["lo_unbounded"].as_bool().unwrap());
assert!(shards[0].descriptor["hi_unbounded"].as_bool().unwrap());
}
#[test]
fn plan_pk_shards_target_zero_treated_as_one() {
let shards = plan_pk_shards("id", 0, 9, 0);
assert_eq!(shards.len(), 1);
assert_eq!(shards[0].descriptor["hi"].as_i64().unwrap(), 9);
}
#[test]
fn shard_bounds_wrap_builds_half_open_predicate() {
let spec = ShardSpec::new(
"1",
serde_json::json!({"key": "id", "lo": 100, "hi": 200, "lo_unbounded": false, "hi_unbounded": false}),
);
let b = ShardBounds::from_spec(&spec).unwrap();
let sql = b.wrap("SELECT * FROM t");
assert!(sql.contains("(SELECT * FROM t) AS _faucet_shard"));
assert!(sql.contains(r#""id" >= 100"#), "got: {sql}");
assert!(
sql.contains(r#""id" < 200"#),
"half-open upper bound: {sql}"
);
}
#[test]
fn shard_bounds_wrap_first_shard_has_no_lower_bound() {
let spec = ShardSpec::new(
"0",
serde_json::json!({"key": "id", "lo": 0, "hi": 100, "lo_unbounded": true, "hi_unbounded": false}),
);
let b = ShardBounds::from_spec(&spec).unwrap();
let sql = b.wrap("SELECT * FROM t");
assert!(sql.contains(r#""id" < 100"#), "upper bound present: {sql}");
assert!(!sql.contains(">="), "first shard has no lower floor: {sql}");
}
#[test]
fn shard_bounds_wrap_last_shard_has_no_upper_bound() {
let spec = ShardSpec::new(
"2",
serde_json::json!({"key": "id", "lo": 200, "hi": 300, "lo_unbounded": false, "hi_unbounded": true}),
);
let b = ShardBounds::from_spec(&spec).unwrap();
let sql = b.wrap("SELECT * FROM t");
assert!(sql.contains(r#""id" >= 200"#), "lower bound present: {sql}");
assert!(
!sql.contains(" < ") && !sql.contains("<="),
"last shard has no upper bound: {sql}"
);
}
#[test]
fn shard_bounds_quotes_key_against_injection() {
let spec = ShardSpec::new(
"0",
serde_json::json!({"key": "weird\"; DROP", "lo": 0, "hi": 1, "lo_unbounded": false, "hi_unbounded": false}),
);
let b = ShardBounds::from_spec(&spec).unwrap();
let sql = b.wrap("SELECT 1");
assert!(
sql.contains(r#""weird""; DROP""#),
"key must be quoted: {sql}"
);
}
#[test]
fn shard_bounds_from_spec_rejects_malformed_descriptor() {
let spec = ShardSpec::new("0", serde_json::json!({"key": "id"})); assert!(ShardBounds::from_spec(&spec).is_none());
assert!(ShardBounds::from_spec(&ShardSpec::whole()).is_none());
}
#[test]
fn exactly_one_shard_includes_null() {
let shards = plan_pk_shards("id", 0, 99, 5);
let null_owners: Vec<usize> = shards
.iter()
.enumerate()
.filter(|(_, s)| s.descriptor["include_null"].as_bool().unwrap_or(false))
.map(|(i, _)| i)
.collect();
assert_eq!(
null_owners,
vec![shards.len() - 1],
"exactly the last shard owns NULL keys"
);
}
#[test]
fn single_shard_plan_still_owns_null() {
let shards = plan_pk_shards("id", 7, 7, 4);
assert_eq!(shards.len(), 1);
assert!(shards[0].descriptor["include_null"].as_bool().unwrap());
}
#[test]
fn last_shard_wrap_emits_is_null_clause() {
let shards = plan_pk_shards("id", 0, 99, 3);
let last = ShardBounds::from_spec(shards.last().unwrap()).unwrap();
let sql = last.wrap("SELECT * FROM t");
assert!(
sql.contains(r#""id" IS NULL"#),
"last shard must match NULL keys: {sql}"
);
assert!(sql.contains(" OR "), "NULL clause OR'd with range: {sql}");
}
#[test]
fn non_last_shard_wrap_omits_is_null_clause() {
let shards = plan_pk_shards("id", 0, 99, 3);
let first = ShardBounds::from_spec(&shards[0]).unwrap();
let sql = first.wrap("SELECT * FROM t");
assert!(
!sql.contains("IS NULL"),
"non-last shard must not match NULL keys: {sql}"
);
}
#[test]
fn predicate_coverage_complete_and_non_overlapping() {
let (min, max, target) = (0i64, 19i64, 4usize);
let bounds: Vec<ShardBounds> = plan_pk_shards("k", min, max, target)
.iter()
.map(|s| ShardBounds::from_spec(s).unwrap())
.collect();
let matches_key = |b: &ShardBounds, key: i64| -> bool {
let lower = b.lo_unbounded || key >= b.lo;
let upper = b.hi_unbounded || key < b.hi;
lower && upper
};
for key in (min - 50)..=(max + 50) {
let matches = bounds.iter().filter(|b| matches_key(b, key)).count();
assert_eq!(matches, 1, "key {key} matched {matches} shards (want 1)");
}
let null_matches = bounds.iter().filter(|b| b.include_null).count();
assert_eq!(null_matches, 1, "NULL keys must match exactly one shard");
}
#[test]
fn single_shard_wrap_selects_whole_dataset_including_null() {
let shards = plan_pk_shards("id", 7, 7, 1);
assert_eq!(shards.len(), 1);
let b = ShardBounds::from_spec(&shards[0]).unwrap();
let sql = b.wrap("SELECT * FROM t");
assert!(sql.contains("WHERE TRUE"), "whole-dataset predicate: {sql}");
assert!(!sql.contains(">="), "no bounds on a lone shard: {sql}");
}
#[test]
fn dataset_uri_strips_credentials() {
let redacted = faucet_core::redact_uri_credentials("postgres://u:p@h:5432/db");
let uri = format!("{}?query={}", redacted, "SELECT 1");
assert_eq!(uri, "postgres://h:5432/db?query=SELECT 1");
}
}