systemprompt_users/repository/banned_ip/
queries.rs1use 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}