1use super::{build_in_clause, map_db_error, params_refs, parse_datetime_row};
4use chrono::Utc;
5use r2d2::Pool;
6use r2d2_sqlite::SqliteConnectionManager;
7use rusqlite::OptionalExtension;
8use stateset_core::{
9 AgentValidationRepository, AgentValidationRequest, AgentValidationResponse,
10 AgentValidationStatus, CommerceError, CreateAgentValidationRequest,
11 CreateAgentValidationResponse, Result, ValidationSummary,
12};
13use uuid::Uuid;
14
15#[derive(Debug)]
17pub struct SqliteAgentValidationRepository {
18 pool: Pool<SqliteConnectionManager>,
19}
20
21impl SqliteAgentValidationRepository {
22 #[must_use]
23 pub const fn new(pool: Pool<SqliteConnectionManager>) -> Self {
24 Self { pool }
25 }
26
27 fn conn(&self) -> Result<r2d2::PooledConnection<SqliteConnectionManager>> {
28 self.pool.get().map_err(|e| CommerceError::DatabaseError(e.to_string()))
29 }
30
31 fn row_to_request(row: &rusqlite::Row<'_>) -> rusqlite::Result<AgentValidationRequest> {
32 Ok(AgentValidationRequest {
33 request_hash: row.get("request_hash")?,
34 agent_registry: row.get("agent_registry")?,
35 agent_id: row.get("agent_id")?,
36 validator_address: row.get("validator_address")?,
37 request_uri: row.get("request_uri")?,
38 created_at: parse_datetime_row(
39 &row.get::<_, String>("created_at")?,
40 "agent_validation_request",
41 "created_at",
42 )?,
43 })
44 }
45
46 fn row_to_response(row: &rusqlite::Row<'_>) -> rusqlite::Result<AgentValidationResponse> {
47 Ok(AgentValidationResponse {
48 id: Uuid::parse_str(&row.get::<_, String>("id")?).map_err(|e| {
49 rusqlite::Error::FromSqlConversionFailure(
50 0,
51 rusqlite::types::Type::Text,
52 Box::new(e),
53 )
54 })?,
55 request_hash: row.get("request_hash")?,
56 agent_registry: row.get("agent_registry")?,
57 agent_id: row.get("agent_id")?,
58 validator_address: row.get("validator_address")?,
59 response: row.get::<_, i64>("response")? as u8,
60 response_uri: row.get("response_uri")?,
61 response_hash: row.get("response_hash")?,
62 tag: row.get("tag")?,
63 created_at: parse_datetime_row(
64 &row.get::<_, String>("created_at")?,
65 "agent_validation_response",
66 "created_at",
67 )?,
68 })
69 }
70}
71
72impl AgentValidationRepository for SqliteAgentValidationRepository {
73 fn request_validation(
74 &self,
75 input: CreateAgentValidationRequest,
76 ) -> Result<AgentValidationRequest> {
77 let conn = self.conn()?;
78 let now = Utc::now();
79 let request_hash = input.request_hash.clone();
80
81 conn.execute(
82 "INSERT INTO agent_validation_requests (
83 request_hash, agent_registry, agent_id, validator_address, request_uri, created_at
84 ) VALUES (?, ?, ?, ?, ?, ?)",
85 rusqlite::params![
86 request_hash,
87 input.agent_registry,
88 input.agent_id,
89 input.validator_address,
90 input.request_uri,
91 now.to_rfc3339(),
92 ],
93 )
94 .map_err(map_db_error)?;
95
96 let mut stmt = conn
97 .prepare("SELECT * FROM agent_validation_requests WHERE request_hash = ?")
98 .map_err(map_db_error)?;
99 stmt.query_row([request_hash], Self::row_to_request).map_err(map_db_error)
100 }
101
102 fn respond_validation(
103 &self,
104 request_hash: &str,
105 input: CreateAgentValidationResponse,
106 ) -> Result<AgentValidationResponse> {
107 if input.response > 100 {
108 return Err(CommerceError::ValidationError(
109 "validation response must be between 0 and 100".to_string(),
110 ));
111 }
112
113 let conn = self.conn()?;
114 let now = Utc::now();
115 let id = Uuid::new_v4();
116
117 let request: AgentValidationRequest = conn
118 .query_row(
119 "SELECT * FROM agent_validation_requests WHERE request_hash = ?",
120 [request_hash],
121 Self::row_to_request,
122 )
123 .optional()
124 .map_err(map_db_error)?
125 .ok_or(CommerceError::NotFound)?;
126
127 conn.execute(
128 "INSERT INTO agent_validation_responses (
129 id, request_hash, agent_registry, agent_id, validator_address,
130 response, response_uri, response_hash, tag, created_at
131 ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",
132 rusqlite::params![
133 id.to_string(),
134 request_hash,
135 request.agent_registry,
136 request.agent_id,
137 request.validator_address,
138 i64::from(input.response),
139 input.response_uri,
140 input.response_hash,
141 input.tag,
142 now.to_rfc3339(),
143 ],
144 )
145 .map_err(map_db_error)?;
146
147 let mut stmt = conn
148 .prepare("SELECT * FROM agent_validation_responses WHERE id = ?")
149 .map_err(map_db_error)?;
150
151 stmt.query_row([id.to_string()], Self::row_to_response).map_err(map_db_error)
152 }
153
154 fn get_validation_status(&self, request_hash: &str) -> Result<Option<AgentValidationStatus>> {
155 let conn = self.conn()?;
156 let request: Option<AgentValidationRequest> = conn
157 .query_row(
158 "SELECT * FROM agent_validation_requests WHERE request_hash = ?",
159 [request_hash],
160 Self::row_to_request,
161 )
162 .optional()
163 .map_err(map_db_error)?;
164
165 let request = match request {
166 Some(req) => req,
167 None => return Ok(None),
168 };
169
170 let mut stmt = conn
171 .prepare(
172 "SELECT response, response_hash, tag, created_at
173 FROM agent_validation_responses
174 WHERE request_hash = ?
175 ORDER BY created_at DESC
176 LIMIT 1",
177 )
178 .map_err(map_db_error)?;
179
180 let row: Option<(i64, Option<String>, Option<String>, String)> = stmt
181 .query_row([request_hash], |row| {
182 Ok((row.get(0)?, row.get(1)?, row.get(2)?, row.get(3)?))
183 })
184 .optional()
185 .map_err(map_db_error)?;
186
187 let (response, response_hash, tag, created_at) = match row {
188 Some(val) => val,
189 None => return Ok(None),
190 };
191
192 let last_update =
193 parse_datetime_row(&created_at, "agent_validation_response", "created_at")
194 .map_err(map_db_error)?;
195
196 Ok(Some(AgentValidationStatus {
197 validator_address: request.validator_address,
198 agent_registry: request.agent_registry,
199 agent_id: request.agent_id,
200 response: response as u8,
201 response_hash,
202 tag,
203 last_update,
204 }))
205 }
206
207 fn get_summary(
208 &self,
209 agent_registry: &str,
210 agent_id: &str,
211 validator_addresses: Option<Vec<String>>,
212 tag: Option<String>,
213 ) -> Result<ValidationSummary> {
214 let conn = self.conn()?;
215 let mut conditions = vec!["agent_registry = ?".to_string(), "agent_id = ?".to_string()];
216 let mut params: Vec<Box<dyn rusqlite::ToSql>> =
217 vec![Box::new(agent_registry.to_string()), Box::new(agent_id.to_string())];
218
219 if let Some(validators) = validator_addresses {
220 if !validators.is_empty() {
221 let placeholders = build_in_clause(validators.len());
222 conditions.push(format!("validator_address IN ({placeholders})"));
223 for validator in validators {
224 params.push(Box::new(validator));
225 }
226 }
227 }
228
229 if let Some(tag_val) = tag {
230 conditions.push("tag = ?".to_string());
231 params.push(Box::new(tag_val));
232 }
233
234 let sql = format!(
235 "SELECT response FROM agent_validation_responses WHERE {}",
236 conditions.join(" AND ")
237 );
238
239 let param_refs = params_refs(¶ms);
240 let mut stmt = conn.prepare(&sql).map_err(map_db_error)?;
241 let mut rows = stmt.query(rusqlite::params_from_iter(param_refs)).map_err(map_db_error)?;
242
243 let mut count: u64 = 0;
244 let mut sum: u64 = 0;
245 while let Some(row) = rows.next().map_err(map_db_error)? {
246 let response: i64 = row.get(0).map_err(map_db_error)?;
247 count += 1;
248 sum += response as u64;
249 }
250
251 if count == 0 {
252 return Ok(ValidationSummary { count: 0, average_response: 0 });
253 }
254
255 Ok(ValidationSummary { count, average_response: (sum / count) as u8 })
256 }
257
258 fn get_agent_validations(&self, agent_registry: &str, agent_id: &str) -> Result<Vec<String>> {
259 let conn = self.conn()?;
260 let mut stmt = conn
261 .prepare(
262 "SELECT request_hash FROM agent_validation_requests
263 WHERE agent_registry = ? AND agent_id = ?",
264 )
265 .map_err(map_db_error)?;
266
267 let rows = stmt
268 .query_map([agent_registry, agent_id], |row| row.get::<_, String>(0))
269 .map_err(map_db_error)?;
270
271 let mut results = Vec::new();
272 for row in rows {
273 results.push(row.map_err(map_db_error)?);
274 }
275 Ok(results)
276 }
277
278 fn get_validator_requests(&self, validator_address: &str) -> Result<Vec<String>> {
279 let conn = self.conn()?;
280 let mut stmt = conn
281 .prepare(
282 "SELECT request_hash FROM agent_validation_requests
283 WHERE validator_address = ?",
284 )
285 .map_err(map_db_error)?;
286
287 let rows = stmt
288 .query_map([validator_address], |row| row.get::<_, String>(0))
289 .map_err(map_db_error)?;
290
291 let mut results = Vec::new();
292 for row in rows {
293 results.push(row.map_err(map_db_error)?);
294 }
295 Ok(results)
296 }
297}