1use pylon_core::schema::{ChannelDescriptor, ChannelPayload, SchemaDescriptor};
29use pylon_pgcon::PgListener;
30use tokio::sync::mpsc;
31
32use crate::error::{Error, Result};
33use crate::value::{Object, Value};
34
35pub struct ChannelListener {
43 rx: mpsc::UnboundedReceiver<Result<Value>>,
44 _conn: PgListener,
45 ended: bool,
50}
51
52impl ChannelListener {
53 pub async fn recv(&mut self) -> Option<Result<Value>> {
81 if self.ended {
82 return None;
83 }
84 let item = self.rx.recv().await;
85 if matches!(item, Some(Err(_))) {
86 self.ended = true;
87 }
88 item
89 }
90}
91
92pub(crate) async fn listen(dsn: &str, schema: &SchemaDescriptor, channel: &str) -> Result<ChannelListener> {
93 let ch = schema
94 .find_channel(channel)
95 .cloned()
96 .ok_or_else(|| Error::UnknownChannel(channel.to_string()))?;
97 let wire_name = ch.wire_name.clone();
98
99 let (tx, rx) = mpsc::unbounded_channel();
100 let conn = PgListener::connect(dsn, move |n| {
101 let _ = tx.send(decode_payload(&ch, n.payload()));
106 })
107 .await
108 .map_err(Error::Db)?;
109 conn.listen(&wire_name).await.map_err(Error::Db)?;
110
111 Ok(ChannelListener {
112 rx,
113 _conn: conn,
114 ended: false,
115 })
116}
117
118fn decode_payload(ch: &ChannelDescriptor, raw: &str) -> Result<Value> {
119 match &ch.payload {
120 ChannelPayload::Type(_) => decode_uuid(raw),
121 ChannelPayload::Scalar(pg_type) => decode_scalar_text(raw, pg_type),
122 ChannelPayload::Object(fields) => decode_object_payload(fields, raw),
123 }
124}
125
126fn decode_uuid(text: &str) -> Result<Value> {
127 text.parse::<uuid::Uuid>()
128 .map(Value::Uuid)
129 .map_err(|e| Error::MalformedPayload(format!("invalid uuid {text:?}: {e}")))
130}
131
132fn decode_scalar_text(text: &str, pg_type: &str) -> Result<Value> {
148 match pg_type {
149 "uuid" => decode_uuid(text),
150 "int2" | "int4" | "int8" => text
151 .parse::<i64>()
152 .map(Value::Int64)
153 .map_err(|e| Error::MalformedPayload(format!("invalid integer {text:?}: {e}"))),
154 "float4" | "float8" => text
155 .parse::<f64>()
156 .map(Value::Float64)
157 .map_err(|e| Error::MalformedPayload(format!("invalid float {text:?}: {e}"))),
158 "numeric" => Ok(Value::Decimal(text.to_string())),
161 "boolean" => Ok(Value::Bool(text == "true" || text == "t")),
162 "date" => parse_date(text),
163 "time" => parse_time(text),
164 "timestamp" => parse_timestamp(text).map(Value::Timestamp),
165 "timestamptz" => parse_timestamptz(text).map(Value::Timestamptz),
166 _ => Ok(Value::Str(text.to_string())),
167 }
168}
169
170const PG_EPOCH_DAYS_FROM_UNIX: i64 = 10_957;
173const PG_EPOCH_MICROS_FROM_UNIX: i64 = PG_EPOCH_DAYS_FROM_UNIX * 86_400 * 1_000_000;
174
175fn malformed(kind: &str, text: &str, e: impl std::fmt::Display) -> Error {
176 Error::MalformedPayload(format!("invalid {kind} {text:?}: {e}"))
177}
178
179fn parse_date(text: &str) -> Result<Value> {
180 let d = chrono::NaiveDate::parse_from_str(text, "%Y-%m-%d").map_err(|e| malformed("date", text, e))?;
181 let days = d
182 .signed_duration_since(chrono::NaiveDate::from_ymd_opt(1970, 1, 1).expect("valid epoch date"))
183 .num_days();
184 Ok(Value::Date((days - PG_EPOCH_DAYS_FROM_UNIX) as i32))
185}
186
187fn parse_time(text: &str) -> Result<Value> {
188 let t = chrono::NaiveTime::parse_from_str(text, "%H:%M:%S%.f").map_err(|e| malformed("time", text, e))?;
190 let micros = t
191 .signed_duration_since(chrono::NaiveTime::from_hms_opt(0, 0, 0).expect("valid midnight"))
192 .num_microseconds()
193 .ok_or_else(|| Error::MalformedPayload(format!("time out of range: {text:?}")))?;
194 Ok(Value::Time(micros))
195}
196
197fn parse_timestamp(text: &str) -> Result<i64> {
198 let ts = chrono::NaiveDateTime::parse_from_str(text, "%Y-%m-%d %H:%M:%S%.f")
199 .map_err(|e| malformed("timestamp", text, e))?;
200 Ok(ts.and_utc().timestamp_micros() - PG_EPOCH_MICROS_FROM_UNIX)
201}
202
203fn parse_timestamptz(text: &str) -> Result<i64> {
206 let ts = chrono::DateTime::parse_from_str(text, "%Y-%m-%d %H:%M:%S%.f%#z")
207 .map_err(|e| malformed("timestamptz", text, e))?;
208 Ok(ts.timestamp_micros() - PG_EPOCH_MICROS_FROM_UNIX)
209}
210
211fn decode_object_payload(declared_fields: &[(String, String)], raw_payload: &str) -> Result<Value> {
212 let parsed: serde_json::Value =
213 serde_json::from_str(raw_payload).map_err(|e| Error::MalformedPayload(format!("invalid JSON payload: {e}")))?;
214 let obj = parsed
215 .as_object()
216 .ok_or_else(|| Error::MalformedPayload("expected a JSON object payload".to_string()))?;
217
218 let mut fields = Vec::with_capacity(declared_fields.len());
219 for (name, pg_type) in declared_fields {
220 let json_value = obj
221 .get(name)
222 .ok_or_else(|| Error::MalformedPayload(format!("payload is missing declared field {name:?}")))?;
223 fields.push((name.clone(), decode_json_value(json_value, pg_type)?));
224 }
225 Ok(Value::Object(Object {
226 type_name: None,
227 fields,
228 implicit_id: false,
229 }))
230}
231
232fn decode_json_value(value: &serde_json::Value, pg_type: &str) -> Result<Value> {
238 match value {
239 serde_json::Value::Null => Ok(Value::Null),
240 serde_json::Value::Bool(b) => Ok(Value::Bool(*b)),
241 serde_json::Value::Number(n) => {
242 if pg_type == "numeric" {
243 Ok(Value::Decimal(n.to_string()))
244 } else if let Some(i) = n.as_i64() {
245 Ok(Value::Int64(i))
246 } else {
247 Ok(Value::Float64(n.as_f64().unwrap_or_default()))
248 }
249 }
250 serde_json::Value::String(s) => decode_scalar_text(s, pg_type),
251 other => Err(Error::MalformedPayload(format!(
252 "unexpected JSON shape {other:?} for a field of type '{pg_type}'"
253 ))),
254 }
255}
256
257#[cfg(test)]
258mod tests {
259 use super::*;
260
261 fn channel(payload: ChannelPayload) -> ChannelDescriptor {
262 ChannelDescriptor {
263 name: "X".into(),
264 module: "m".into(),
265 wire_name: "m__x".into(),
266 payload,
267 description: None,
268 }
269 }
270
271 #[test]
278 fn decodes_a_date_to_pg_epoch_days() {
279 assert_eq!(decode_scalar_text("2000-01-01", "date").unwrap(), Value::Date(0));
281 assert_eq!(decode_scalar_text("2000-01-02", "date").unwrap(), Value::Date(1));
282 assert_eq!(decode_scalar_text("1999-12-31", "date").unwrap(), Value::Date(-1));
283 }
284
285 #[test]
286 fn decodes_a_time_to_microseconds_since_midnight() {
287 assert_eq!(decode_scalar_text("00:00:00", "time").unwrap(), Value::Time(0));
288 assert_eq!(
289 decode_scalar_text("01:00:00", "time").unwrap(),
290 Value::Time(3_600_000_000)
291 );
292 assert_eq!(decode_scalar_text("00:00:00.5", "time").unwrap(), Value::Time(500_000));
294 }
295
296 #[test]
297 fn decodes_a_timestamp_to_pg_epoch_microseconds() {
298 assert_eq!(
299 decode_scalar_text("2000-01-01 00:00:00", "timestamp").unwrap(),
300 Value::Timestamp(0)
301 );
302 assert_eq!(
303 decode_scalar_text("2000-01-01 00:00:01.5", "timestamp").unwrap(),
304 Value::Timestamp(1_500_000)
305 );
306 }
307
308 #[test]
309 fn decodes_a_timestamptz_and_normalises_the_offset() {
310 assert_eq!(
311 decode_scalar_text("2000-01-01 00:00:00+00", "timestamptz").unwrap(),
312 Value::Timestamptz(0)
313 );
314 assert_eq!(
316 decode_scalar_text("2000-01-01 01:00:00+01", "timestamptz").unwrap(),
317 Value::Timestamptz(0)
318 );
319 assert_eq!(
320 decode_scalar_text("1999-12-31 23:30:00-00:30", "timestamptz").unwrap(),
321 Value::Timestamptz(0)
322 );
323 }
324
325 #[test]
326 fn a_malformed_temporal_payload_is_an_error_not_a_string() {
327 for (text, pg_type) in [
328 ("not-a-date", "date"),
329 ("25:99:99", "time"),
330 ("nope", "timestamp"),
331 ("2000-01-01 00:00:00", "timestamptz"), ] {
333 let err = decode_scalar_text(text, pg_type).unwrap_err();
334 assert!(
335 matches!(err, Error::MalformedPayload(_)),
336 "{pg_type} {text:?} should be MalformedPayload, got {err:?}"
337 );
338 }
339 }
340
341 #[test]
342 fn interval_and_bytea_still_pass_through_as_text() {
343 assert_eq!(
345 decode_scalar_text("1 day", "interval").unwrap(),
346 Value::Str("1 day".into())
347 );
348 assert_eq!(
349 decode_scalar_text("\\xdeadbeef", "bytea").unwrap(),
350 Value::Str("\\xdeadbeef".into())
351 );
352 }
353
354 #[test]
355 fn decode_scalar_uuid() {
356 let u = uuid::Uuid::parse_str("3fa85f64-5717-4562-b3fc-2c963f66afa6").unwrap();
357 assert_eq!(decode_scalar_text(&u.to_string(), "uuid").unwrap(), Value::Uuid(u));
358 }
359
360 #[test]
361 fn decode_scalar_integers() {
362 assert_eq!(decode_scalar_text("42", "int2").unwrap(), Value::Int64(42));
363 assert_eq!(decode_scalar_text("42", "int4").unwrap(), Value::Int64(42));
364 assert_eq!(decode_scalar_text("42", "int8").unwrap(), Value::Int64(42));
365 }
366
367 #[test]
368 fn decode_scalar_floats() {
369 assert_eq!(decode_scalar_text("0.5", "float4").unwrap(), Value::Float64(0.5));
370 assert_eq!(decode_scalar_text("0.5", "float8").unwrap(), Value::Float64(0.5));
371 }
372
373 #[test]
374 fn decode_scalar_numeric_kept_as_string() {
375 assert_eq!(
376 decode_scalar_text("123.456", "numeric").unwrap(),
377 Value::Decimal("123.456".to_string())
378 );
379 }
380
381 #[test]
382 fn decode_scalar_boolean() {
383 assert_eq!(decode_scalar_text("true", "boolean").unwrap(), Value::Bool(true));
384 assert_eq!(decode_scalar_text("false", "boolean").unwrap(), Value::Bool(false));
385 }
386
387 #[test]
388 fn decode_scalar_text_passthrough() {
389 assert_eq!(
390 decode_scalar_text("hello", "text").unwrap(),
391 Value::Str("hello".to_string())
392 );
393 }
394
395 #[test]
396 fn decode_scalar_unsupported_type_falls_back_to_raw_string() {
397 assert_eq!(
398 decode_scalar_text("1 day 02:00:00", "interval").unwrap(),
399 Value::Str("1 day 02:00:00".to_string())
400 );
401 }
402
403 #[test]
404 fn decode_scalar_rejects_malformed_uuid() {
405 let err = decode_scalar_text("not-a-uuid", "uuid").unwrap_err();
406 assert!(matches!(err, Error::MalformedPayload(_)), "got: {err:?}");
407 }
408
409 #[test]
410 fn decode_scalar_rejects_malformed_integer() {
411 let err = decode_scalar_text("not-an-int", "int8").unwrap_err();
412 assert!(matches!(err, Error::MalformedPayload(_)), "got: {err:?}");
413 }
414
415 #[test]
416 fn decode_payload_type_kind_is_a_uuid() {
417 let u = uuid::Uuid::parse_str("3fa85f64-5717-4562-b3fc-2c963f66afa6").unwrap();
418 let ch = channel(ChannelPayload::Type("m::Widget".into()));
419 assert_eq!(decode_payload(&ch, &u.to_string()).unwrap(), Value::Uuid(u));
420 }
421
422 #[test]
423 fn decode_payload_scalar_kind() {
424 let ch = channel(ChannelPayload::Scalar("text".into()));
425 assert_eq!(decode_payload(&ch, "hello").unwrap(), Value::Str("hello".to_string()));
426 }
427
428 #[test]
429 fn decode_payload_object_kind() {
430 let u = uuid::Uuid::parse_str("3fa85f64-5717-4562-b3fc-2c963f66afa6").unwrap();
431 let ch = channel(ChannelPayload::Object(vec![
432 ("doc_id".into(), "uuid".into()),
433 ("score".into(), "float8".into()),
434 ]));
435 let payload = format!(r#"{{"doc_id": "{u}", "score": 0.5}}"#);
436 let Value::Object(obj) = decode_payload(&ch, &payload).unwrap() else {
437 panic!("expected Object")
438 };
439 assert_eq!(obj.get("doc_id"), Some(&Value::Uuid(u)));
440 assert_eq!(obj.get("score"), Some(&Value::Float64(0.5)));
441 }
442
443 #[test]
444 fn decode_payload_object_kind_rejects_malformed_json() {
445 let ch = channel(ChannelPayload::Object(vec![("doc_id".into(), "uuid".into())]));
446 let err = decode_payload(&ch, "not json at all").unwrap_err();
447 assert!(matches!(err, Error::MalformedPayload(_)), "got: {err:?}");
448 }
449
450 #[test]
451 fn decode_payload_object_kind_rejects_missing_field() {
452 let ch = channel(ChannelPayload::Object(vec![("doc_id".into(), "uuid".into())]));
453 let err = decode_payload(&ch, "{}").unwrap_err();
454 assert!(matches!(err, Error::MalformedPayload(_)), "got: {err:?}");
455 }
456
457 #[tokio::test]
458 async fn listen_rejects_an_unknown_channel() {
459 let schema = SchemaDescriptor::default();
460 let err = listen("postgresql://ignored", &schema, "NoSuchChannel")
461 .await
462 .err()
463 .expect("expected an error");
464 assert!(matches!(err, Error::UnknownChannel(_)), "got: {err:?}");
465 }
466}