uqa_client/notifications/
request.rs1use 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#[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 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 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 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}