Skip to main content

systemprompt_users/repository/banned_ip/
queries.rs

1//! Banned-IP lookup and mutation queries.
2//!
3//! Copyright (c) systemprompt.io — Business Source License 1.1.
4//! See <https://systemprompt.io> for licensing details.
5
6use crate::error::Result;
7use systemprompt_identifiers::SessionId;
8
9use super::BannedIpRepository;
10use super::types::{BanDuration, BanIpParams, BanIpWithMetadataParams, BannedIp};
11
12impl BannedIpRepository {
13    pub async fn is_banned(&self, ip_address: &str) -> Result<bool> {
14        let result = sqlx::query_scalar!(
15            r#"
16            SELECT EXISTS(
17                SELECT 1 FROM banned_ips
18                WHERE ip_address = $1
19                  AND (expires_at IS NULL OR expires_at > CURRENT_TIMESTAMP)
20            ) as "exists!"
21            "#,
22            ip_address
23        )
24        .fetch_one(&*self.write_pool)
25        .await?;
26
27        Ok(result)
28    }
29
30    pub async fn find_ban(&self, ip_address: &str) -> Result<Option<BannedIp>> {
31        let row = sqlx::query_as!(
32            BannedIp,
33            r#"
34            SELECT
35                ip_address,
36                reason,
37                banned_at,
38                expires_at,
39                ban_count,
40                last_offense_path,
41                last_user_agent,
42                is_permanent,
43                source_fingerprint,
44                ban_source,
45                associated_session_ids as "associated_session_ids: Vec<SessionId>"
46            FROM banned_ips
47            WHERE ip_address = $1
48              AND (expires_at IS NULL OR expires_at > CURRENT_TIMESTAMP)
49            "#,
50            ip_address
51        )
52        .fetch_optional(&*self.write_pool)
53        .await?;
54
55        Ok(row)
56    }
57
58    pub async fn ban_ip(&self, params: BanIpParams<'_>) -> Result<()> {
59        let expires_at = params.duration.to_expiry();
60        let is_permanent = matches!(params.duration, BanDuration::Permanent);
61
62        sqlx::query!(
63            r#"
64            INSERT INTO banned_ips (
65                ip_address, reason, expires_at, is_permanent,
66                source_fingerprint, ban_source
67            )
68            VALUES ($1, $2, $3, $4, $5, $6)
69            ON CONFLICT (ip_address) DO UPDATE SET
70                reason = $2,
71                expires_at = CASE
72                    WHEN banned_ips.is_permanent THEN banned_ips.expires_at
73                    ELSE COALESCE($3, banned_ips.expires_at)
74                END,
75                ban_count = banned_ips.ban_count + 1,
76                is_permanent = banned_ips.is_permanent OR $4,
77                source_fingerprint = COALESCE($5, banned_ips.source_fingerprint),
78                ban_source = $6
79            "#,
80            params.ip_address,
81            params.reason,
82            expires_at,
83            is_permanent,
84            params.source_fingerprint,
85            params.ban_source
86        )
87        .execute(&*self.write_pool)
88        .await?;
89
90        Ok(())
91    }
92
93    pub async fn ban_ip_with_metadata(&self, params: BanIpWithMetadataParams<'_>) -> Result<()> {
94        let expires_at = params.duration.to_expiry();
95        let is_permanent = matches!(params.duration, BanDuration::Permanent);
96        let session_ids: Option<Vec<String>> = params.session_id.map(|s| vec![s.to_string()]);
97
98        sqlx::query!(
99            r#"
100            INSERT INTO banned_ips (
101                ip_address, reason, expires_at, is_permanent,
102                source_fingerprint, ban_source, last_offense_path,
103                last_user_agent, associated_session_ids
104            )
105            VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9)
106            ON CONFLICT (ip_address) DO UPDATE SET
107                reason = $2,
108                expires_at = CASE
109                    WHEN banned_ips.is_permanent THEN banned_ips.expires_at
110                    ELSE COALESCE($3, banned_ips.expires_at)
111                END,
112                ban_count = banned_ips.ban_count + 1,
113                is_permanent = banned_ips.is_permanent OR $4,
114                source_fingerprint = COALESCE($5, banned_ips.source_fingerprint),
115                ban_source = $6,
116                last_offense_path = COALESCE($7, banned_ips.last_offense_path),
117                last_user_agent = COALESCE($8, banned_ips.last_user_agent),
118                associated_session_ids = CASE
119                    WHEN $9::TEXT[] IS NOT NULL
120                    THEN array_cat(COALESCE(banned_ips.associated_session_ids, '{}'::TEXT[]), $9)
121                    ELSE banned_ips.associated_session_ids
122                END
123            "#,
124            params.ip_address,
125            params.reason,
126            expires_at,
127            is_permanent,
128            params.source_fingerprint,
129            params.ban_source,
130            params.offense_path,
131            params.user_agent,
132            session_ids.as_deref()
133        )
134        .execute(&*self.write_pool)
135        .await?;
136
137        Ok(())
138    }
139
140    pub async fn unban_ip(&self, ip_address: &str) -> Result<bool> {
141        let result = sqlx::query!(
142            r#"
143            DELETE FROM banned_ips
144            WHERE ip_address = $1
145            "#,
146            ip_address
147        )
148        .execute(&*self.write_pool)
149        .await?;
150
151        Ok(result.rows_affected() > 0)
152    }
153
154    pub async fn cleanup_expired(&self) -> Result<u64> {
155        let result = sqlx::query!(
156            r#"
157            DELETE FROM banned_ips
158            WHERE expires_at IS NOT NULL
159              AND expires_at < CURRENT_TIMESTAMP
160              AND NOT is_permanent
161            "#
162        )
163        .execute(&*self.write_pool)
164        .await?;
165
166        Ok(result.rows_affected())
167    }
168}