1use crate::config::PostgresSourceConfig;
4use async_trait::async_trait;
5use faucet_core::shard::{
6 PkShardBounds, ShardSpec, parse_pk_shard, pk_bounds_query, pk_shards_from_bounds,
7};
8use faucet_core::util::quote_ident;
9use faucet_core::{FaucetError, Stream, StreamPage};
10use futures::TryStreamExt;
11use serde_json::Value;
12use sqlx::postgres::PgPoolOptions;
13use sqlx::{Column, PgPool, Row};
14use std::pin::Pin;
15use std::sync::Mutex;
16
17pub struct PostgresSource {
19 config: PostgresSourceConfig,
20 pool: PgPool,
21 applied_shard: Mutex<Option<PkShardBounds>>,
25}
26
27impl PostgresSource {
28 pub async fn new(config: PostgresSourceConfig) -> Result<Self, FaucetError> {
30 faucet_core::validate_batch_size(config.batch_size)?;
31
32 let pool = PgPoolOptions::new()
33 .max_connections(config.max_connections)
34 .connect(&config.connection_url)
35 .await
36 .map_err(|e| FaucetError::Config(format!("PostgreSQL connection failed: {e}")))?;
37
38 Ok(Self {
39 config,
40 pool,
41 applied_shard: Mutex::new(None),
42 })
43 }
44
45 fn shard_wrap(&self, query: String) -> String {
47 match &*self.applied_shard.lock().expect("shard mutex poisoned") {
48 Some(bounds) => bounds.wrap(&query, quote_ident),
49 None => query,
50 }
51 }
52}
53
54fn pg_value_to_json(row: &sqlx::postgres::PgRow, col_name: &str) -> Value {
59 if let Ok(v) = row.try_get::<Value, _>(col_name) {
61 return v;
62 }
63
64 if let Ok(v) = row.try_get::<String, _>(col_name) {
66 return Value::String(v);
67 }
68 if let Ok(v) = row.try_get::<i64, _>(col_name) {
69 return Value::Number(v.into());
70 }
71 if let Ok(v) = row.try_get::<i32, _>(col_name) {
72 return Value::Number(v.into());
73 }
74 if let Ok(v) = row.try_get::<i16, _>(col_name) {
75 return Value::Number(v.into());
76 }
77 if let Ok(v) = row.try_get::<f64, _>(col_name) {
78 return serde_json::Number::from_f64(v)
79 .map(Value::Number)
80 .unwrap_or(Value::Null);
81 }
82 if let Ok(v) = row.try_get::<f32, _>(col_name) {
83 return serde_json::Number::from_f64(v as f64)
84 .map(Value::Number)
85 .unwrap_or(Value::Null);
86 }
87 if let Ok(v) = row.try_get::<bool, _>(col_name) {
88 return Value::Bool(v);
89 }
90
91 if let Ok(v) =
94 row.try_get::<sqlx::types::chrono::DateTime<sqlx::types::chrono::Utc>, _>(col_name)
95 {
96 return Value::String(v.to_rfc3339());
97 }
98 if let Ok(v) = row.try_get::<sqlx::types::chrono::NaiveDateTime, _>(col_name) {
99 return Value::String(v.to_string());
100 }
101 if let Ok(v) = row.try_get::<sqlx::types::chrono::NaiveDate, _>(col_name) {
102 return Value::String(v.to_string());
103 }
104 if let Ok(v) = row.try_get::<sqlx::types::chrono::NaiveTime, _>(col_name) {
105 return Value::String(v.to_string());
106 }
107 if let Ok(v) = row.try_get::<sqlx::types::Uuid, _>(col_name) {
109 return Value::String(v.to_string());
110 }
111 if let Ok(v) = row.try_get::<sqlx::types::BigDecimal, _>(col_name) {
113 return Value::String(v.to_string());
114 }
115 if let Ok(v) = row.try_get::<Vec<u8>, _>(col_name) {
117 use base64::Engine as _;
118 return Value::String(base64::engine::general_purpose::STANDARD.encode(v));
119 }
120
121 Value::Null
122}
123
124fn resolve_query(
127 config: &PostgresSourceConfig,
128 context: &std::collections::HashMap<String, Value>,
129) -> (String, Vec<Value>) {
130 if context.is_empty() {
131 (config.query.clone(), Vec::new())
132 } else {
133 faucet_core::util::substitute_context_bind_params(
134 &config.query,
135 context,
136 config.params.len() + 1,
137 |i| format!("${i}"),
138 )
139 }
140}
141
142#[derive(Debug, Clone, Copy, PartialEq, Eq)]
151enum NumberBind {
152 I64,
154 U64,
157 F64,
159}
160
161fn classify_number(n: &serde_json::Number) -> NumberBind {
167 if n.is_i64() {
168 NumberBind::I64
169 } else if n.is_u64() {
170 NumberBind::U64
171 } else {
172 NumberBind::F64
173 }
174}
175
176fn bind_params<'q>(
179 mut query: sqlx::query::Query<'q, sqlx::Postgres, sqlx::postgres::PgArguments>,
180 config_params: &'q [Value],
181 bind_values: &'q [Value],
182) -> sqlx::query::Query<'q, sqlx::Postgres, sqlx::postgres::PgArguments> {
183 for value in config_params.iter().chain(bind_values) {
190 query = match value {
191 Value::String(s) => query.bind(s.clone()),
192 Value::Number(n) => match classify_number(n) {
193 NumberBind::I64 => query.bind(n.as_i64().unwrap()),
195 NumberBind::U64 => query.bind(n.as_u64().unwrap() as i64),
199 NumberBind::F64 => query.bind(n.as_f64().unwrap_or(0.0)),
200 },
201 Value::Bool(b) => query.bind(*b),
202 Value::Null => query.bind(None::<String>),
203 _ => query.bind(value.to_string()),
204 };
205 }
206 query
207}
208
209fn row_to_json(row: &sqlx::postgres::PgRow) -> Value {
212 let mut map = serde_json::Map::new();
213 for col in row.columns() {
214 let name = col.name().to_string();
215 let value = pg_value_to_json(row, &name);
216 map.insert(name, value);
217 }
218 Value::Object(map)
219}
220
221#[async_trait]
222impl faucet_core::Source for PostgresSource {
223 async fn fetch_with_context(
224 &self,
225 context: &std::collections::HashMap<String, serde_json::Value>,
226 ) -> Result<Vec<Value>, FaucetError> {
227 let (query_str, bind_values) = resolve_query(&self.config, context);
228 let query_str = self.shard_wrap(query_str);
229 let query = bind_params(sqlx::query(&query_str), &self.config.params, &bind_values);
230
231 let rows = query
232 .fetch_all(&self.pool)
233 .await
234 .map_err(|e| FaucetError::Config(format!("PostgreSQL query failed: {e}")))?;
235
236 let records: Vec<Value> = rows.iter().map(row_to_json).collect();
237 tracing::info!(rows = records.len(), query = %self.config.query, "PostgreSQL source fetch complete");
238 Ok(records)
239 }
240
241 fn stream_pages<'a>(
254 &'a self,
255 context: &'a std::collections::HashMap<String, Value>,
256 _batch_size: usize,
257 ) -> Pin<Box<dyn Stream<Item = Result<StreamPage, FaucetError>> + Send + 'a>> {
258 let batch_size = self.config.batch_size;
259
260 Box::pin(async_stream::try_stream! {
261 let (query_str, bind_values) = resolve_query(&self.config, context);
262 let query_str = self.shard_wrap(query_str);
263 let query = bind_params(
264 sqlx::query(&query_str),
265 &self.config.params,
266 &bind_values,
267 );
268
269 let mut rows = query.fetch(&self.pool);
270 let chunk = if batch_size == 0 { usize::MAX } else { batch_size };
271 let initial_capacity = if batch_size == 0 { 1024 } else { batch_size };
272 let mut buffer: Vec<Value> = Vec::with_capacity(initial_capacity);
273 let mut total = 0usize;
274
275 while let Some(row) = rows
276 .try_next()
277 .await
278 .map_err(|e| FaucetError::Config(format!("PostgreSQL query failed: {e}")))?
279 {
280 buffer.push(row_to_json(&row));
281 if buffer.len() >= chunk {
282 let page = std::mem::replace(&mut buffer, Vec::with_capacity(initial_capacity));
283 total += page.len();
284 yield StreamPage { records: page, bookmark: None };
285 }
286 }
287 if !buffer.is_empty() {
288 total += buffer.len();
289 yield StreamPage { records: buffer, bookmark: None };
290 }
291
292 tracing::info!(
293 rows = total,
294 batch_size,
295 query = %self.config.query,
296 "PostgreSQL source stream complete",
297 );
298 })
299 }
300
301 fn config_schema(&self) -> serde_json::Value {
302 serde_json::to_value(faucet_core::schema_for!(PostgresSourceConfig))
303 .expect("schema serialization")
304 }
305
306 fn dataset_uri(&self) -> String {
307 format!(
308 "{}?query={}",
309 faucet_core::redact_uri_credentials(&self.config.connection_url),
310 self.config.query
311 )
312 }
313
314 fn is_shardable(&self) -> bool {
316 self.config.shard.is_some()
317 }
318
319 async fn enumerate_shards(&self, target: usize) -> Result<Vec<ShardSpec>, FaucetError> {
324 let Some(shard_cfg) = &self.config.shard else {
325 return Ok(vec![ShardSpec::whole()]);
326 };
327
328 let bounds_sql =
329 pk_bounds_query(&self.config.query, "e_ident(&shard_cfg.key), "BIGINT");
330 let row = bind_params(sqlx::query(&bounds_sql), &self.config.params, &[])
331 .fetch_one(&self.pool)
332 .await
333 .map_err(|e| {
334 FaucetError::Source(format!(
335 "postgres: failed to compute shard bounds for key {:?} \
336 (it must be an integer-typed column): {e}",
337 shard_cfg.key
338 ))
339 })?;
340
341 let lo: Option<i64> = row.try_get("lo").map_err(|e| {
342 FaucetError::Source(format!("postgres: shard bounds decode failed: {e}"))
343 })?;
344 let hi: Option<i64> = row.try_get("hi").map_err(|e| {
345 FaucetError::Source(format!("postgres: shard bounds decode failed: {e}"))
346 })?;
347 Ok(pk_shards_from_bounds(&shard_cfg.key, lo, hi, target))
348 }
349
350 async fn apply_shard(&self, shard: &ShardSpec) -> Result<(), FaucetError> {
353 *self.applied_shard.lock().expect("shard mutex poisoned") =
354 parse_pk_shard(shard, "postgres")?;
355 Ok(())
356 }
357}
358
359#[cfg(test)]
360mod tests {
361 use super::*;
362 use faucet_core::shard::plan_pk_shards;
363
364 type ShardBounds = PkShardBounds;
368
369 #[tokio::test]
370 async fn new_rejects_out_of_range_batch_size() {
371 let mut config = PostgresSourceConfig::new("postgres://localhost/test", "SELECT 1");
372 config.batch_size = faucet_core::MAX_BATCH_SIZE + 1;
373 match PostgresSource::new(config).await {
374 Err(faucet_core::FaucetError::Config(m)) => {
375 assert!(m.contains("batch_size"), "got: {m}")
376 }
377 _ => panic!("expected a batch_size Config error"),
378 }
379 }
380
381 fn num(v: serde_json::Value) -> serde_json::Number {
384 match v {
385 serde_json::Value::Number(n) => n,
386 _ => panic!("not a number"),
387 }
388 }
389
390 #[test]
391 fn classify_small_int_is_i64() {
392 assert_eq!(
393 classify_number(&num(serde_json::json!(42))),
394 NumberBind::I64
395 );
396 assert_eq!(
397 classify_number(&num(serde_json::json!(-7))),
398 NumberBind::I64
399 );
400 assert_eq!(classify_number(&num(serde_json::json!(0))), NumberBind::I64);
401 }
402
403 #[test]
404 fn classify_above_2_pow_53_stays_i64_not_f64() {
405 let v = 9_007_199_254_740_993i64; assert_eq!(classify_number(&num(serde_json::json!(v))), NumberBind::I64);
409 }
410
411 #[test]
412 fn classify_i64_max_is_i64() {
413 assert_eq!(
414 classify_number(&num(serde_json::json!(i64::MAX))),
415 NumberBind::I64
416 );
417 assert_eq!(
418 classify_number(&num(serde_json::json!(i64::MIN))),
419 NumberBind::I64
420 );
421 }
422
423 #[test]
424 fn classify_above_i64_max_is_u64() {
425 let v: u64 = i64::MAX as u64 + 1;
427 assert_eq!(classify_number(&num(serde_json::json!(v))), NumberBind::U64);
428 assert_eq!(
429 classify_number(&num(serde_json::json!(u64::MAX))),
430 NumberBind::U64
431 );
432 }
433
434 #[test]
435 fn classify_float_is_f64() {
436 assert_eq!(
437 classify_number(&num(serde_json::json!(3.5))),
438 NumberBind::F64
439 );
440 assert_eq!(
441 classify_number(&num(serde_json::json!(-0.5))),
442 NumberBind::F64
443 );
444 }
445
446 #[test]
449 fn plan_pk_shards_covers_full_range_without_gaps_or_overlap() {
450 let shards = plan_pk_shards("id", 0, 99, 4);
451 assert_eq!(shards.len(), 4);
452 let mut expected_lo = 0i64;
454 for (i, s) in shards.iter().enumerate() {
455 let d = &s.descriptor;
456 assert_eq!(d["key"], "id");
457 assert_eq!(d["lo"].as_i64().unwrap(), expected_lo);
458 let hi = d["hi"].as_i64().unwrap();
459 let first = i == 0;
460 let last = i == shards.len() - 1;
461 assert_eq!(d["lo_unbounded"].as_bool().unwrap(), first);
462 assert_eq!(d["hi_unbounded"].as_bool().unwrap(), last);
463 expected_lo = hi; }
465 }
466
467 #[test]
468 fn plan_pk_shards_never_more_shards_than_values() {
469 let shards = plan_pk_shards("pk", 5, 7, 10);
471 assert!(shards.len() <= 3, "got {} shards", shards.len());
472 assert!(
473 shards[0].descriptor["lo_unbounded"].as_bool().unwrap(),
474 "first shard is unbounded below"
475 );
476 assert!(
477 shards.last().unwrap().descriptor["hi_unbounded"]
478 .as_bool()
479 .unwrap(),
480 "last shard is unbounded above"
481 );
482 }
483
484 #[test]
485 fn plan_pk_shards_single_value_one_shard() {
486 let shards = plan_pk_shards("id", 42, 42, 8);
487 assert_eq!(shards.len(), 1);
488 assert!(shards[0].descriptor["lo_unbounded"].as_bool().unwrap());
490 assert!(shards[0].descriptor["hi_unbounded"].as_bool().unwrap());
491 }
492
493 #[test]
494 fn plan_pk_shards_target_zero_treated_as_one() {
495 let shards = plan_pk_shards("id", 0, 9, 0);
496 assert_eq!(shards.len(), 1);
497 assert_eq!(shards[0].descriptor["hi"].as_i64().unwrap(), 9);
498 }
499
500 #[test]
501 fn shard_bounds_wrap_builds_half_open_predicate() {
502 let spec = ShardSpec::new(
504 "1",
505 serde_json::json!({"key": "id", "lo": 100, "hi": 200, "lo_unbounded": false, "hi_unbounded": false}),
506 );
507 let b = ShardBounds::from_spec(&spec).unwrap();
508 let sql = b.wrap("SELECT * FROM t", quote_ident);
509 assert!(sql.contains("(SELECT * FROM t) AS _faucet_shard"));
510 assert!(sql.contains(r#""id" >= 100"#), "got: {sql}");
511 assert!(
512 sql.contains(r#""id" < 200"#),
513 "half-open upper bound: {sql}"
514 );
515 }
516
517 #[test]
518 fn shard_bounds_wrap_first_shard_has_no_lower_bound() {
519 let spec = ShardSpec::new(
522 "0",
523 serde_json::json!({"key": "id", "lo": 0, "hi": 100, "lo_unbounded": true, "hi_unbounded": false}),
524 );
525 let b = ShardBounds::from_spec(&spec).unwrap();
526 let sql = b.wrap("SELECT * FROM t", quote_ident);
527 assert!(sql.contains(r#""id" < 100"#), "upper bound present: {sql}");
528 assert!(!sql.contains(">="), "first shard has no lower floor: {sql}");
529 }
530
531 #[test]
532 fn shard_bounds_wrap_last_shard_has_no_upper_bound() {
533 let spec = ShardSpec::new(
536 "2",
537 serde_json::json!({"key": "id", "lo": 200, "hi": 300, "lo_unbounded": false, "hi_unbounded": true}),
538 );
539 let b = ShardBounds::from_spec(&spec).unwrap();
540 let sql = b.wrap("SELECT * FROM t", quote_ident);
541 assert!(sql.contains(r#""id" >= 200"#), "lower bound present: {sql}");
542 assert!(
543 !sql.contains(" < ") && !sql.contains("<="),
544 "last shard has no upper bound: {sql}"
545 );
546 }
547
548 #[test]
549 fn shard_bounds_quotes_key_against_injection() {
550 let spec = ShardSpec::new(
551 "0",
552 serde_json::json!({"key": "weird\"; DROP", "lo": 0, "hi": 1, "lo_unbounded": false, "hi_unbounded": false}),
553 );
554 let b = ShardBounds::from_spec(&spec).unwrap();
555 let sql = b.wrap("SELECT 1", quote_ident);
556 assert!(
558 sql.contains(r#""weird""; DROP""#),
559 "key must be quoted: {sql}"
560 );
561 }
562
563 #[test]
564 fn shard_bounds_from_spec_rejects_malformed_descriptor() {
565 let spec = ShardSpec::new("0", serde_json::json!({"key": "id"})); assert!(ShardBounds::from_spec(&spec).is_none());
567 assert!(ShardBounds::from_spec(&ShardSpec::whole()).is_none());
568 }
569
570 #[test]
573 fn exactly_one_shard_includes_null() {
574 let shards = plan_pk_shards("id", 0, 99, 5);
575 let null_owners: Vec<usize> = shards
576 .iter()
577 .enumerate()
578 .filter(|(_, s)| s.descriptor["include_null"].as_bool().unwrap_or(false))
579 .map(|(i, _)| i)
580 .collect();
581 assert_eq!(
582 null_owners,
583 vec![shards.len() - 1],
584 "exactly the last shard owns NULL keys"
585 );
586 }
587
588 #[test]
589 fn single_shard_plan_still_owns_null() {
590 let shards = plan_pk_shards("id", 7, 7, 4);
592 assert_eq!(shards.len(), 1);
593 assert!(shards[0].descriptor["include_null"].as_bool().unwrap());
594 }
595
596 #[test]
597 fn last_shard_wrap_emits_is_null_clause() {
598 let shards = plan_pk_shards("id", 0, 99, 3);
599 let last = ShardBounds::from_spec(shards.last().unwrap()).unwrap();
600 let sql = last.wrap("SELECT * FROM t", quote_ident);
601 assert!(
602 sql.contains(r#""id" IS NULL"#),
603 "last shard must match NULL keys: {sql}"
604 );
605 assert!(sql.contains(" OR "), "NULL clause OR'd with range: {sql}");
606 }
607
608 #[test]
609 fn non_last_shard_wrap_omits_is_null_clause() {
610 let shards = plan_pk_shards("id", 0, 99, 3);
611 let first = ShardBounds::from_spec(&shards[0]).unwrap();
613 let sql = first.wrap("SELECT * FROM t", quote_ident);
614 assert!(
615 !sql.contains("IS NULL"),
616 "non-last shard must not match NULL keys: {sql}"
617 );
618 }
619
620 #[test]
625 fn predicate_coverage_complete_and_non_overlapping() {
626 let (min, max, target) = (0i64, 19i64, 4usize);
627 let bounds: Vec<ShardBounds> = plan_pk_shards("k", min, max, target)
628 .iter()
629 .map(|s| ShardBounds::from_spec(s).unwrap())
630 .collect();
631
632 let matches_key = |b: &ShardBounds, key: i64| -> bool {
635 let lower = b.lo_unbounded || key >= b.lo;
636 let upper = b.hi_unbounded || key < b.hi;
637 lower && upper
638 };
639
640 for key in (min - 50)..=(max + 50) {
644 let matches = bounds.iter().filter(|b| matches_key(b, key)).count();
645 assert_eq!(matches, 1, "key {key} matched {matches} shards (want 1)");
646 }
647
648 let null_matches = bounds.iter().filter(|b| b.include_null).count();
650 assert_eq!(null_matches, 1, "NULL keys must match exactly one shard");
651 }
652
653 #[test]
654 fn single_shard_wrap_selects_whole_dataset_including_null() {
655 let shards = plan_pk_shards("id", 7, 7, 1);
657 assert_eq!(shards.len(), 1);
658 let b = ShardBounds::from_spec(&shards[0]).unwrap();
659 let sql = b.wrap("SELECT * FROM t", quote_ident);
660 assert!(sql.contains("WHERE TRUE"), "whole-dataset predicate: {sql}");
661 assert!(!sql.contains(">="), "no bounds on a lone shard: {sql}");
662 }
663
664 #[test]
667 fn dataset_uri_strips_credentials() {
668 let redacted = faucet_core::redact_uri_credentials("postgres://u:p@h:5432/db");
671 let uri = format!("{}?query={}", redacted, "SELECT 1");
672 assert_eq!(uri, "postgres://h:5432/db?query=SELECT 1");
673 }
674
675 fn lazy_source(config: PostgresSourceConfig) -> PostgresSource {
679 let pool = PgPoolOptions::new()
680 .acquire_timeout(std::time::Duration::from_millis(200))
682 .connect_lazy(&config.connection_url)
683 .expect("lazy pool");
684 PostgresSource {
685 config,
686 pool,
687 applied_shard: Mutex::new(None),
688 }
689 }
690
691 #[tokio::test]
692 async fn apply_shard_then_shard_wrap_narrows_query() {
693 use faucet_core::Source as _;
694 let mut config =
695 PostgresSourceConfig::new("postgres://u@127.0.0.1:1/db", "SELECT * FROM t");
696 config.shard = Some(crate::config::ShardConfig { key: "id".into() });
697 let source = lazy_source(config);
698 assert!(source.is_shardable());
699
700 assert_eq!(source.shard_wrap("SELECT 1".into()), "SELECT 1");
702 source
703 .apply_shard(&faucet_core::ShardSpec::whole())
704 .await
705 .unwrap();
706 assert_eq!(source.shard_wrap("SELECT 1".into()), "SELECT 1");
707
708 let spec = &plan_pk_shards("id", 0, 99, 2)[0];
710 source.apply_shard(spec).await.unwrap();
711 let wrapped = source.shard_wrap("SELECT * FROM t".into());
712 assert!(wrapped.contains(r#""id""#), "got: {wrapped}");
713 assert!(wrapped.contains("_faucet_shard"), "got: {wrapped}");
714
715 let err = source.enumerate_shards(4).await.unwrap_err();
718 assert!(
719 err.to_string().contains("shard bounds"),
720 "expected bounds-probe error, got: {err}"
721 );
722 }
723}