Skip to main content

stateset_db/sqlite/
agent_validation.rs

1//! SQLite validation registry repository implementation
2
3use 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/// SQLite implementation of `AgentValidationRepository`
16#[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(&params);
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}