Skip to main content

synd_persistence/sqlite/feed_registry/subscription/
mod.rs

1use 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    // Editing attributes keeps subscribed_at.
41    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;