Skip to main content

pylon_client/
listen.rs

1//
2// This source file is part of the Pylon open source project.
3//
4// Copyright (c) 2026 Jaldis B.V.
5//
6// Licensed under the MIT OR Apache-2.0 license (the "License");
7// you may not use this file except in compliance with the License.
8// You may obtain a copy of the License at
9//
10//     https://opensource.org/licenses/MIT
11//     https://www.apache.org/licenses/LICENSE-2.0
12//
13// Unless required by applicable law or agreed to in writing, software
14// distributed under the License is distributed on an "AS IS" BASIS,
15// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
16// See the License for the specific language governing permissions and
17// limitations under the License.
18//
19
20//! `Client::listen()` — subscribes to a schema-declared `Channel` over a
21//! dedicated (non-pooled) LISTEN/NOTIFY connection and decodes each payload
22//! into the crate's generic `Value`. Mirrors `pylon/client.py`'s own
23//! `Client.listen()` (there, a typed async generator; here, a `recv()`-based
24//! handle — this crate has no `Stream`/async-generator precedent to build
25//! on, and `recv()` matches `tokio::sync::mpsc::Receiver`'s own idiom
26//! closely enough not to need one).
27
28use 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
35/// A live subscription to one `Channel`, returned by [`crate::Client::listen`].
36///
37/// Holds a dedicated (non-pooled) connection for as long as it's alive —
38/// `LISTEN` is per-session, so sharing a pooled connection would leak the
39/// subscription onto whatever unrelated query later borrows that same
40/// connection back out of the pool. Dropping this closes the connection
41/// and ends the server-side subscription with it.
42pub struct ChannelListener {
43    rx: mpsc::UnboundedReceiver<Result<Value>>,
44    _conn: PgListener,
45    /// Set once a payload has failed to decode. The subscription is over at
46    /// that point; further `recv()` calls report end-of-stream rather than
47    /// resuming, so a caller can't accidentally keep consuming a channel it
48    /// has already seen foreign data on.
49    ended: bool,
50}
51
52impl ChannelListener {
53    /// Waits for the next decoded payload. Returns `None` once the dedicated
54    /// connection closes, or once a payload has failed to decode.
55    ///
56    /// A malformed payload yields `Err(Error::MalformedPayload(_))` **and
57    /// ends the subscription** — every later call returns `None`. A channel
58    /// is a database-wide name, so anything can publish to it; failing hard
59    /// keeps foreign data on a typed channel visible instead of silently
60    /// dropped, and not resuming keeps that failure from being papered over
61    /// by the next good message.
62    ///
63    /// A consumer that needs to stay subscribed past a bad message
64    /// re-subscribes itself:
65    ///
66    /// ```ignore
67    /// loop {
68    ///     let mut sub = client.listen("UserUpdates").await?;
69    ///     while let Some(payload) = sub.recv().await {
70    ///         match payload {
71    ///             Ok(v) => handle(v),
72    ///             Err(e) => { log(e); break; }  // re-listen on the next pass
73    ///         }
74    ///     }
75    /// }
76    /// ```
77    ///
78    /// This matches `Client.listen()` on the Python side, where the same
79    /// failure raises out of the `async for` and ends the generator.
80    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        // The receiving end only ever drops once `ChannelListener` itself
102        // does (which also drops `_conn`, ending this callback's own
103        // background task) — a send failure here means that already
104        // happened moments ago; nothing left to report it to.
105        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
132/// Decodes NOTIFY's raw text payload as PostgreSQL's own `<pg_type>::text`
133/// cast would have rendered it (see `notify()`'s SQL emission — the
134/// payload is always literally cast to `text` before being sent).
135///
136/// Covers every base type `PG_TYPE_MAP` maps a built-in scalar to, except
137/// `interval` and `bytea`, whose text encodings are non-trivial to parse
138/// correctly and uncommon as a pub/sub payload — both pass through as the
139/// raw string rather than risk a wrong decode.
140///
141/// Deliberately identical in coverage to
142/// `pylon.schema._channels._decode_scalar_text` on the Python side: the same
143/// `Channel` declaration has to yield the same type through either client,
144/// and this used to fall back to a raw string for every temporal type while
145/// Python decoded them, so a `Channel(datetime)` produced a `datetime` in
146/// Python and a bare string in Rust.
147fn 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        // Kept as its canonical string form, same as `Value::Decimal`'s own
159        // documented convention elsewhere — no parsing needed.
160        "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
170/// Days from the Unix epoch to PostgreSQL's (2000-01-01), the offset between
171/// what `chrono` counts from and what `Value::Date`/`Timestamp` store.
172const 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    // Postgres omits the fractional part when it's zero.
189    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
203/// `timestamptz` renders with a numeric UTC offset (`+00`, `-04:30`), which
204/// `%#z` accepts in all the widths Postgres emits.
205fn 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
232/// Like `decode_scalar_text`, but for a value already parsed out of an
233/// Object channel's JSON payload — `serde_json` already turned a JSON
234/// number/bool/null into the right shape, so only the types `to_jsonb()`
235/// renders as a JSON *string* (uuid, numeric, and the same date/time/etc.
236/// carve-out `decode_scalar_text` documents) need any further decoding here.
237fn 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    // ── Temporal parity with `_channels._decode_scalar_text` ──────────
272    //
273    // These four used to fall through to `Value::Str` here while the Python
274    // client decoded them, so one `Channel` declaration produced two
275    // different types depending on which client read it.
276
277    #[test]
278    fn decodes_a_date_to_pg_epoch_days() {
279        // 2000-01-01 is the PG epoch itself, so day 0.
280        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        // Postgres only prints the fractional part when it is non-zero.
293        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        // Same instant, written in a different zone.
315        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"), // missing the offset
332        ] {
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        // Matches the Python side's own carve-out — deliberately not parsed.
344        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}