synd_persistence/sqlite/feed_registry/subscription/
mod.rs1use std::str::FromStr;
2
3use chrono::{DateTime, Utc};
4use sqlx::{Sqlite, Transaction};
5use synd_feed::types::{Category, FeedUrl, Requirement};
6use synd_registry::{
7 FeedSubscriptionAttrs, RegistryDbResult, SubscriberId, Subscription, SubscriptionKey,
8 crawl::target_list::{FeedSubscriptions, SubscriptionPolicy},
9 db::SubscriptionDb,
10 query::{Subscriptions, SubscriptionsQuery},
11};
12
13use super::{
14 SqliteRegistryTx, codec,
15 error::{DecodeResultExt, IntoDbResult, SqliteResult},
16 feed,
17 pagination::PageLimit,
18};
19
20const SUBSCRIPTION_SELECT_COLUMNS: &str = r#"
21s.subscriber_id AS subscriber_id,
22f.url AS feed_url,
23s.requirement AS requirement,
24s.category AS category,
25s.crawl_policy_json AS crawl_policy_json,
26s.subscribed_at AS subscribed_at
27"#;
28
29async fn upsert(
30 tx: &mut Transaction<'_, Sqlite>,
31 subscription: &SubscriptionKey,
32 attrs: FeedSubscriptionAttrs,
33 now: DateTime<Utc>,
34) -> SqliteResult<()> {
35 let requirement = attrs.requirement.map(|r| r.to_string());
36 let category = attrs.category.map(|c| c.to_string());
37 let policy_json = codec::encode_crawl_policy_json(attrs.crawl_policy)?;
38 let feed_pk = feed::upsert_pk(tx, &subscription.feed_url).await?;
39
40 sqlx::query(
42 r#"
43 INSERT INTO feed_subscription (
44 subscriber_id,
45 feed_pk,
46 requirement,
47 category,
48 crawl_policy_json,
49 subscribed_at
50 )
51 VALUES (?, ?, ?, ?, ?, ?)
52 ON CONFLICT(subscriber_id, feed_pk) DO UPDATE SET
53 requirement = excluded.requirement,
54 category = excluded.category,
55 crawl_policy_json = excluded.crawl_policy_json
56 "#,
57 )
58 .bind(subscription.subscriber_id.as_str())
59 .bind(feed_pk)
60 .bind(requirement)
61 .bind(category)
62 .bind(policy_json)
63 .bind(now)
64 .execute(&mut **tx)
65 .await?;
66
67 Ok(())
68}
69
70async fn delete(
71 tx: &mut Transaction<'_, Sqlite>,
72 subscriber_id: &SubscriberId,
73 feed_url: &FeedUrl,
74) -> SqliteResult<()> {
75 sqlx::query(
76 r#"
77 DELETE FROM feed_subscription
78 WHERE subscriber_id = ?
79 AND feed_pk = (
80 SELECT pk
81 FROM feed
82 WHERE url = ?
83 )
84 "#,
85 )
86 .bind(subscriber_id.as_str())
87 .bind(feed_url.as_str())
88 .execute(&mut **tx)
89 .await?;
90
91 Ok(())
92}
93
94async fn contains(
95 tx: &mut Transaction<'_, Sqlite>,
96 subscriber_id: &SubscriberId,
97 feed_url: &FeedUrl,
98) -> SqliteResult<bool> {
99 let row = sqlx::query(
100 r#"
101 SELECT 1 AS found
102 FROM feed_subscription AS s
103 INNER JOIN feed AS f
104 ON f.pk = s.feed_pk
105 WHERE s.subscriber_id = ? AND f.url = ?
106 LIMIT 1
107 "#,
108 )
109 .bind(subscriber_id.as_str())
110 .bind(feed_url.as_str())
111 .fetch_optional(&mut **tx)
112 .await?;
113
114 Ok(row.is_some())
115}
116
117async fn list(
118 tx: &mut Transaction<'_, Sqlite>,
119 query: SubscriptionsQuery,
120) -> SqliteResult<Subscriptions> {
121 let page_limit = PageLimit::new(query.first);
122 let rows = if let Some(after) = query.after {
123 let sql = format!(
124 r#"
125 SELECT {SUBSCRIPTION_SELECT_COLUMNS}
126 FROM feed_subscription AS s
127 INNER JOIN feed AS f
128 ON f.pk = s.feed_pk
129 WHERE s.subscriber_id = ? AND f.url > ?
130 ORDER BY f.url
131 LIMIT ?
132 "#
133 );
134 sqlx::query_as::<_, SubscriptionRow>(&sql)
135 .bind(query.subscriber_id.as_str())
136 .bind(after)
137 .bind(page_limit.sql_limit())
138 .fetch_all(&mut **tx)
139 .await
140 } else {
141 let sql = format!(
142 r#"
143 SELECT {SUBSCRIPTION_SELECT_COLUMNS}
144 FROM feed_subscription AS s
145 INNER JOIN feed AS f
146 ON f.pk = s.feed_pk
147 WHERE s.subscriber_id = ?
148 ORDER BY f.url
149 LIMIT ?
150 "#
151 );
152 sqlx::query_as::<_, SubscriptionRow>(&sql)
153 .bind(query.subscriber_id.as_str())
154 .bind(page_limit.sql_limit())
155 .fetch_all(&mut **tx)
156 .await
157 }?;
158
159 let mut nodes = rows
160 .into_iter()
161 .map(SubscriptionRow::into_subscription)
162 .collect::<SqliteResult<Vec<_>>>()?;
163 let has_next_page = page_limit.truncate_overfetch(&mut nodes);
164 let end_cursor = nodes.last().map(|sub| sub.feed_url.to_string());
165
166 Ok(Subscriptions::from_subscriptions(
167 nodes,
168 has_next_page,
169 end_cursor,
170 ))
171}
172
173async fn load_for_feed(
174 tx: &mut Transaction<'_, Sqlite>,
175 feed_url: &FeedUrl,
176) -> SqliteResult<FeedSubscriptions> {
177 let rows = sqlx::query_as::<_, FeedSubscriptionRow>(
178 r#"
179 SELECT
180 s.subscriber_id AS subscriber_id,
181 f.url AS feed_url,
182 s.crawl_policy_json AS crawl_policy_json
183 FROM feed_subscription AS s
184 INNER JOIN feed AS f
185 ON f.pk = s.feed_pk
186 WHERE f.url = ?
187 ORDER BY s.subscriber_id
188 "#,
189 )
190 .bind(feed_url.as_str())
191 .fetch_all(&mut **tx)
192 .await?;
193
194 let subscriptions = rows
195 .into_iter()
196 .map(FeedSubscriptionRow::into_policy)
197 .collect::<SqliteResult<Vec<_>>>()?;
198
199 Ok(FeedSubscriptions::new(feed_url.clone(), subscriptions))
200}
201
202#[derive(sqlx::FromRow)]
203struct SubscriptionRow {
204 subscriber_id: String,
205 feed_url: String,
206 requirement: Option<String>,
207 category: Option<String>,
208 crawl_policy_json: String,
209 subscribed_at: DateTime<Utc>,
210}
211
212impl SubscriptionRow {
213 fn into_subscription(self) -> SqliteResult<Subscription> {
214 Ok(Subscription {
215 subscriber_id: SubscriberId::new(self.subscriber_id),
216 feed_url: FeedUrl::parse(&self.feed_url).decode()?,
217 requirement: self
218 .requirement
219 .as_deref()
220 .map(Requirement::from_str)
221 .transpose()
222 .decode()?,
223 category: self.category.map(Category::new).transpose().decode()?,
224 crawl_policy: codec::decode_crawl_policy_json(&self.crawl_policy_json)?,
225 subscribed_at: self.subscribed_at,
226 })
227 }
228}
229
230#[derive(sqlx::FromRow)]
231struct FeedSubscriptionRow {
232 subscriber_id: String,
233 feed_url: String,
234 crawl_policy_json: String,
235}
236
237impl FeedSubscriptionRow {
238 fn into_policy(self) -> SqliteResult<SubscriptionPolicy> {
239 let subscription = SubscriptionKey::new(
240 SubscriberId::new(self.subscriber_id),
241 FeedUrl::parse(&self.feed_url).decode()?,
242 );
243
244 Ok(SubscriptionPolicy::new(
245 subscription,
246 codec::decode_crawl_policy_json(&self.crawl_policy_json)?,
247 ))
248 }
249}
250
251impl SubscriptionDb for SqliteRegistryTx<'_> {
252 async fn upsert_subscription(
253 &mut self,
254 subscription: &SubscriptionKey,
255 attrs: FeedSubscriptionAttrs,
256 now: DateTime<Utc>,
257 ) -> RegistryDbResult<()> {
258 upsert(&mut self.tx, subscription, attrs, now).await.db()
259 }
260
261 async fn delete_subscription(
262 &mut self,
263 subscriber_id: &SubscriberId,
264 feed_url: &FeedUrl,
265 ) -> RegistryDbResult<()> {
266 delete(&mut self.tx, subscriber_id, feed_url).await.db()
267 }
268
269 async fn has_subscription(
270 &mut self,
271 subscriber_id: &SubscriberId,
272 feed_url: &FeedUrl,
273 ) -> RegistryDbResult<bool> {
274 contains(&mut self.tx, subscriber_id, feed_url).await.db()
275 }
276
277 async fn list_subscriptions(
278 &mut self,
279 query: SubscriptionsQuery,
280 ) -> RegistryDbResult<Subscriptions> {
281 list(&mut self.tx, query).await.db()
282 }
283
284 async fn load_feed_subscriptions(
285 &mut self,
286 feed_url: &FeedUrl,
287 ) -> RegistryDbResult<FeedSubscriptions> {
288 load_for_feed(&mut self.tx, feed_url).await.db()
289 }
290}
291
292#[cfg(test)]
293mod tests;