Skip to main content

uqa_client/notifications/
request.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Exact channel sets, bounded request materialization and resume rejection.
8
9use super::{json, ProtocolError, MAX_NOTIFICATION_WIRE_BYTES};
10use serde::{
11    de::{self, DeserializeSeed, MapAccess, SeqAccess, Visitor},
12    Deserialize, Deserializer, Serialize,
13};
14use std::{fmt, num::NonZeroUsize};
15
16/// Immutable exact channel set, stored in bytewise order for bounded membership checks. Request order has no notification-order semantics.
17#[derive(Clone, PartialEq, Eq)]
18pub struct SubscriptionRequest {
19    channels: Vec<String>,
20}
21
22impl SubscriptionRequest {
23    pub fn new(channels: &[&str], max_channels: NonZeroUsize) -> Result<Self, ProtocolError> {
24        check_count(channels.len(), max_channels)?;
25        if channels.iter().any(|channel| !valid_channel(channel)) {
26            return Err(ProtocolError::InvalidChannels);
27        }
28        // Count escaped output before copying any caller-owned channel strings.
29        json::encode(
30            &RequestWire {
31                protocol_version: 1,
32                channels,
33            },
34            false,
35        )?;
36        let mut owned = Vec::new();
37        owned
38            .try_reserve_exact(channels.len())
39            .map_err(|_| ProtocolError::Allocation)?;
40        for channel in channels {
41            let mut value = String::new();
42            value
43                .try_reserve_exact(channel.len())
44                .map_err(|_| ProtocolError::Allocation)?;
45            value.push_str(channel);
46            owned.push(value);
47        }
48        Self::from_channels(owned)
49    }
50
51    /// Decode a bounded protocol request before registration. A nonempty resume header is rejected even when it contains whitespace.
52    pub fn from_json(
53        input: &[u8],
54        max_channels: NonZeroUsize,
55        last_event_id: Option<&[u8]>,
56    ) -> Result<Self, ProtocolError> {
57        if last_event_id.is_some_and(|value| !value.is_empty()) {
58            return Err(ProtocolError::ResumeUnsupported);
59        }
60        let input = json::validate(input)?;
61        let mut deserializer = serde_json::Deserializer::from_str(input);
62        let mut failure = None;
63        let request = RequestSeed {
64            max_channels,
65            failure: &mut failure,
66        }
67        .deserialize(&mut deserializer)
68        .map_err(|_| failure.unwrap_or(ProtocolError::InvalidFields))?;
69        deserializer.end().map_err(|_| ProtocolError::InvalidJSON)?;
70        Self::from_channels(request)
71    }
72
73    pub fn encode(&self) -> Result<Vec<u8>, ProtocolError> {
74        json::encode(
75            &RequestWire {
76                protocol_version: 1,
77                channels: &self.channels,
78            },
79            true,
80        )
81    }
82
83    pub fn channels(&self) -> &[String] {
84        &self.channels
85    }
86
87    pub(super) fn contains(&self, channel: &str) -> bool {
88        self.channels
89            .binary_search_by(|value| value.as_str().cmp(channel))
90            .is_ok()
91    }
92
93    fn from_channels(mut channels: Vec<String>) -> Result<Self, ProtocolError> {
94        if channels.is_empty() {
95            return Err(ProtocolError::InvalidChannels);
96        }
97        channels.sort_unstable();
98        if channels.windows(2).any(|pair| pair[0] == pair[1]) {
99            return Err(ProtocolError::InvalidChannels);
100        }
101        Ok(Self { channels })
102    }
103}
104
105impl fmt::Debug for SubscriptionRequest {
106    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
107        formatter
108            .debug_struct("SubscriptionRequest")
109            .field("channel_count", &self.channels.len())
110            .finish_non_exhaustive()
111    }
112}
113
114#[derive(Serialize)]
115struct RequestWire<T> {
116    protocol_version: u8,
117    channels: T,
118}
119
120fn check_count(count: usize, max: NonZeroUsize) -> Result<(), ProtocolError> {
121    if count == 0 {
122        return Err(ProtocolError::InvalidChannels);
123    }
124    if count > max.get() {
125        return Err(ProtocolError::ChannelLimit);
126    }
127    // Every admitted name needs at least three JSON bytes plus a separator or final bracket.
128    if count > MAX_NOTIFICATION_WIRE_BYTES / 4 {
129        return Err(ProtocolError::ByteLimit);
130    }
131    Ok(())
132}
133
134fn valid_channel(channel: &str) -> bool {
135    !channel.is_empty() && channel.len() <= 63 && !channel.contains('\0')
136}
137
138#[derive(Deserialize)]
139#[serde(field_identifier, rename_all = "snake_case")]
140enum Field {
141    ProtocolVersion,
142    Channels,
143}
144
145struct RequestSeed<'a> {
146    max_channels: NonZeroUsize,
147    failure: &'a mut Option<ProtocolError>,
148}
149
150impl<'de> DeserializeSeed<'de> for RequestSeed<'_> {
151    type Value = Vec<String>;
152    fn deserialize<D: Deserializer<'de>>(self, deserializer: D) -> Result<Self::Value, D::Error> {
153        deserializer.deserialize_map(self)
154    }
155}
156
157impl<'de> Visitor<'de> for RequestSeed<'_> {
158    type Value = Vec<String>;
159    fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
160        formatter.write_str("a notification request")
161    }
162    fn visit_map<A: MapAccess<'de>>(self, mut map: A) -> Result<Self::Value, A::Error> {
163        let mut version = None;
164        let mut channels = None;
165        while let Some(field) = map.next_key::<Field>()? {
166            match field {
167                Field::ProtocolVersion => {
168                    if version.is_some() {
169                        return Err(de::Error::duplicate_field("protocol_version"));
170                    }
171                    let value = map.next_value::<u64>()?;
172                    if value != 1 {
173                        *self.failure = Some(ProtocolError::UnsupportedVersion);
174                        return Err(de::Error::custom("unsupported notification version"));
175                    }
176                    version = Some(value);
177                }
178                Field::Channels => {
179                    if channels.is_some() {
180                        return Err(de::Error::duplicate_field("channels"));
181                    }
182                    channels = Some(map.next_value_seed(ChannelsSeed {
183                        maximum: self.max_channels,
184                        failure: self.failure,
185                    })?);
186                }
187            }
188        }
189        version.ok_or_else(|| de::Error::missing_field("protocol_version"))?;
190        channels.ok_or_else(|| de::Error::missing_field("channels"))
191    }
192}
193
194struct ChannelsSeed<'a> {
195    maximum: NonZeroUsize,
196    failure: &'a mut Option<ProtocolError>,
197}
198
199impl<'de> DeserializeSeed<'de> for ChannelsSeed<'_> {
200    type Value = Vec<String>;
201    fn deserialize<D: Deserializer<'de>>(self, deserializer: D) -> Result<Self::Value, D::Error> {
202        deserializer.deserialize_seq(self)
203    }
204}
205
206impl<'de> Visitor<'de> for ChannelsSeed<'_> {
207    type Value = Vec<String>;
208    fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
209        formatter.write_str("bounded notification channels")
210    }
211    fn visit_seq<A: SeqAccess<'de>>(self, mut seq: A) -> Result<Self::Value, A::Error> {
212        let mut channels = Vec::new();
213        while let Some(channel) = seq.next_element::<String>()? {
214            let validation = check_count(channels.len() + 1, self.maximum).and_then(|()| {
215                if valid_channel(&channel) {
216                    Ok(())
217                } else {
218                    Err(ProtocolError::InvalidChannels)
219                }
220            });
221            if let Err(error) = validation {
222                *self.failure = Some(error);
223                return Err(de::Error::custom("invalid notification channels"));
224            }
225            if channels.len() == channels.capacity() {
226                let next = channels
227                    .capacity()
228                    .saturating_mul(2)
229                    .max(4)
230                    .min(self.maximum.get())
231                    .min(MAX_NOTIFICATION_WIRE_BYTES / 4);
232                channels
233                    .try_reserve_exact(next - channels.len())
234                    .map_err(|_| {
235                        *self.failure = Some(ProtocolError::Allocation);
236                        de::Error::custom("notification allocation failed")
237                    })?;
238            }
239            channels.push(channel);
240        }
241        Ok(channels)
242    }
243}