Skip to main content

ahri_tre_app/
query.rs

1use ahri_tre_tabular::ArrowTable;
2use ahri_tre_types::{StudyId, VersionId};
3use duckdb::{Connection, params_from_iter, types::Value};
4use serde::{Deserialize, Serialize};
5
6use crate::AppError;
7
8#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
9pub struct LakeQueryContext {
10    pub study_id: StudyId,
11    pub dataset_id: Option<VersionId>,
12    pub principal: Option<String>,
13}
14
15impl LakeQueryContext {
16    pub fn new(study_id: StudyId) -> Self {
17        Self {
18            study_id,
19            dataset_id: None,
20            principal: None,
21        }
22    }
23
24    pub fn for_dataset(mut self, dataset_id: VersionId) -> Self {
25        self.dataset_id = Some(dataset_id);
26        self
27    }
28
29    pub fn for_principal(mut self, principal: impl Into<String>) -> Self {
30        self.principal = Some(principal.into());
31        self
32    }
33}
34
35#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
36pub struct LakeRelation {
37    pub parts: Vec<String>,
38}
39
40impl LakeRelation {
41    pub fn try_new<I, S>(parts: I) -> Result<Self, AppError>
42    where
43        I: IntoIterator<Item = S>,
44        S: Into<String>,
45    {
46        let parts: Vec<String> = parts.into_iter().map(Into::into).collect();
47        if !(1..=3).contains(&parts.len()) {
48            return Err(AppError::Validation(
49                "lake query relation must contain table, schema.table, or catalog.schema.table"
50                    .to_string(),
51            ));
52        }
53        for part in &parts {
54            validate_identifier_part("lake query relation", part)?;
55        }
56        Ok(Self { parts })
57    }
58
59    pub fn table(table: impl Into<String>) -> Result<Self, AppError> {
60        Self::try_new([table.into()])
61    }
62
63    pub fn schema_table(
64        schema: impl Into<String>,
65        table: impl Into<String>,
66    ) -> Result<Self, AppError> {
67        Self::try_new([schema.into(), table.into()])
68    }
69
70    pub fn catalog_schema_table(
71        catalog: impl Into<String>,
72        schema: impl Into<String>,
73        table: impl Into<String>,
74    ) -> Result<Self, AppError> {
75        Self::try_new([catalog.into(), schema.into(), table.into()])
76    }
77
78    fn to_sql(&self) -> Result<String, AppError> {
79        Self::try_new(self.parts.clone())?;
80        self.parts
81            .iter()
82            .map(|part| quote_identifier("lake query relation", part))
83            .collect::<Result<Vec<_>, _>>()
84            .map(|parts| parts.join("."))
85    }
86}
87
88#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
89pub struct LakeQueryRequest {
90    pub context: LakeQueryContext,
91    pub relation: LakeRelation,
92    pub projection: Vec<String>,
93    pub filters: Vec<LakeQueryFilter>,
94    pub order_by: Vec<LakeQueryOrder>,
95    pub limit: Option<u64>,
96}
97
98impl LakeQueryRequest {
99    pub fn new(context: LakeQueryContext, relation: LakeRelation) -> Self {
100        Self {
101            context,
102            relation,
103            projection: Vec::new(),
104            filters: Vec::new(),
105            order_by: Vec::new(),
106            limit: None,
107        }
108    }
109
110    pub fn with_projection<I, S>(mut self, projection: I) -> Self
111    where
112        I: IntoIterator<Item = S>,
113        S: Into<String>,
114    {
115        self.projection = projection.into_iter().map(Into::into).collect();
116        self
117    }
118
119    pub fn with_filters<I>(mut self, filters: I) -> Self
120    where
121        I: IntoIterator<Item = LakeQueryFilter>,
122    {
123        self.filters = filters.into_iter().collect();
124        self
125    }
126
127    pub fn with_order_by<I>(mut self, order_by: I) -> Self
128    where
129        I: IntoIterator<Item = LakeQueryOrder>,
130    {
131        self.order_by = order_by.into_iter().collect();
132        self
133    }
134
135    pub fn with_limit(mut self, limit: u64) -> Self {
136        self.limit = Some(limit);
137        self
138    }
139}
140
141#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
142pub enum LakeQueryFilter {
143    Compare {
144        column: String,
145        operator: LakeQueryComparison,
146        value: LakeQueryValue,
147    },
148    IsNull {
149        column: String,
150    },
151    IsNotNull {
152        column: String,
153    },
154}
155
156impl LakeQueryFilter {
157    pub fn compare(
158        column: impl Into<String>,
159        operator: LakeQueryComparison,
160        value: LakeQueryValue,
161    ) -> Self {
162        Self::Compare {
163            column: column.into(),
164            operator,
165            value,
166        }
167    }
168
169    pub fn is_null(column: impl Into<String>) -> Self {
170        Self::IsNull {
171            column: column.into(),
172        }
173    }
174
175    pub fn is_not_null(column: impl Into<String>) -> Self {
176        Self::IsNotNull {
177            column: column.into(),
178        }
179    }
180}
181
182#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
183pub enum LakeQueryComparison {
184    Equal,
185    NotEqual,
186    LessThan,
187    LessThanOrEqual,
188    GreaterThan,
189    GreaterThanOrEqual,
190}
191
192#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
193pub enum LakeQueryValue {
194    Null,
195    Bool(bool),
196    Int64(i64),
197    Float64(f64),
198    Utf8(String),
199}
200
201#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
202pub struct LakeQueryOrder {
203    pub column: String,
204    pub direction: LakeQueryOrderDirection,
205}
206
207impl LakeQueryOrder {
208    pub fn asc(column: impl Into<String>) -> Self {
209        Self {
210            column: column.into(),
211            direction: LakeQueryOrderDirection::Asc,
212        }
213    }
214
215    pub fn desc(column: impl Into<String>) -> Self {
216        Self {
217            column: column.into(),
218            direction: LakeQueryOrderDirection::Desc,
219        }
220    }
221}
222
223#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
224pub enum LakeQueryOrderDirection {
225    Asc,
226    Desc,
227}
228
229pub trait LakeQueryAuthorizer: Send + Sync {
230    fn authorize_lake_query(&self, request: &LakeQueryRequest) -> Result<(), AppError>;
231}
232
233#[derive(Debug)]
234pub struct LakeQueryResult {
235    pub generated_sql: String,
236    pub table: ArrowTable,
237}
238
239pub fn execute_lake_query<A>(
240    connection: &Connection,
241    request: &LakeQueryRequest,
242    authorizer: &A,
243) -> Result<LakeQueryResult, AppError>
244where
245    A: LakeQueryAuthorizer + ?Sized,
246{
247    authorizer.authorize_lake_query(request)?;
248
249    let planned = plan_lake_query(request)?;
250    let mut statement = connection
251        .prepare(&planned.sql)
252        .map_err(|error| infrastructure_error("lake query prepare failed", error))?;
253    let mut arrow = statement
254        .query_arrow(params_from_iter(planned.parameters.iter()))
255        .map_err(|error| infrastructure_error("lake query execution failed", error))?;
256    let schema = arrow.get_schema();
257    let batches = arrow.by_ref().collect();
258    let table = ArrowTable::try_new(schema, batches).map_err(|error| {
259        infrastructure_error("lake query Arrow result validation failed", error)
260    })?;
261
262    Ok(LakeQueryResult {
263        generated_sql: planned.sql,
264        table,
265    })
266}
267
268#[derive(Debug)]
269struct PlannedLakeQuery {
270    sql: String,
271    parameters: Vec<Value>,
272}
273
274fn plan_lake_query(request: &LakeQueryRequest) -> Result<PlannedLakeQuery, AppError> {
275    let projection = if request.projection.is_empty() {
276        "*".to_string()
277    } else {
278        request
279            .projection
280            .iter()
281            .map(|column| quote_identifier("lake query projection", column))
282            .collect::<Result<Vec<_>, _>>()?
283            .join(", ")
284    };
285
286    let mut sql = format!("SELECT {projection} FROM {}", request.relation.to_sql()?);
287    let mut parameters = Vec::new();
288
289    if !request.filters.is_empty() {
290        let mut clauses = Vec::with_capacity(request.filters.len());
291        for filter in &request.filters {
292            clauses.push(filter_sql(filter, &mut parameters)?);
293        }
294        sql.push_str(" WHERE ");
295        sql.push_str(&clauses.join(" AND "));
296    }
297
298    if !request.order_by.is_empty() {
299        let clauses = request
300            .order_by
301            .iter()
302            .map(order_sql)
303            .collect::<Result<Vec<_>, _>>()?;
304        sql.push_str(" ORDER BY ");
305        sql.push_str(&clauses.join(", "));
306    }
307
308    if let Some(limit) = request.limit {
309        sql.push_str(" LIMIT ");
310        sql.push_str(&limit.to_string());
311    }
312
313    Ok(PlannedLakeQuery { sql, parameters })
314}
315
316fn filter_sql(filter: &LakeQueryFilter, parameters: &mut Vec<Value>) -> Result<String, AppError> {
317    match filter {
318        LakeQueryFilter::Compare {
319            column,
320            operator,
321            value,
322        } => {
323            if matches!(value, LakeQueryValue::Null) {
324                return Err(AppError::Validation(
325                    "lake query NULL comparisons must use is_null or is_not_null".to_string(),
326                ));
327            }
328            parameters.push(query_value_to_duckdb(value)?);
329            Ok(format!(
330                "{} {} ?",
331                quote_identifier("lake query filter", column)?,
332                comparison_sql(*operator)
333            ))
334        }
335        LakeQueryFilter::IsNull { column } => Ok(format!(
336            "{} IS NULL",
337            quote_identifier("lake query filter", column)?
338        )),
339        LakeQueryFilter::IsNotNull { column } => Ok(format!(
340            "{} IS NOT NULL",
341            quote_identifier("lake query filter", column)?
342        )),
343    }
344}
345
346fn order_sql(order: &LakeQueryOrder) -> Result<String, AppError> {
347    let direction = match order.direction {
348        LakeQueryOrderDirection::Asc => "ASC",
349        LakeQueryOrderDirection::Desc => "DESC",
350    };
351    Ok(format!(
352        "{} {direction}",
353        quote_identifier("lake query ordering", &order.column)?
354    ))
355}
356
357fn comparison_sql(operator: LakeQueryComparison) -> &'static str {
358    match operator {
359        LakeQueryComparison::Equal => "=",
360        LakeQueryComparison::NotEqual => "<>",
361        LakeQueryComparison::LessThan => "<",
362        LakeQueryComparison::LessThanOrEqual => "<=",
363        LakeQueryComparison::GreaterThan => ">",
364        LakeQueryComparison::GreaterThanOrEqual => ">=",
365    }
366}
367
368fn query_value_to_duckdb(value: &LakeQueryValue) -> Result<Value, AppError> {
369    match value {
370        LakeQueryValue::Null => Ok(Value::Null),
371        LakeQueryValue::Bool(value) => Ok(Value::Boolean(*value)),
372        LakeQueryValue::Int64(value) => Ok(Value::BigInt(*value)),
373        LakeQueryValue::Float64(value) if value.is_finite() => Ok(Value::Double(*value)),
374        LakeQueryValue::Float64(_) => Err(AppError::Validation(
375            "lake query floating point filter values must be finite".to_string(),
376        )),
377        LakeQueryValue::Utf8(value) => Ok(Value::Text(value.clone())),
378    }
379}
380
381fn quote_identifier(context: &str, value: &str) -> Result<String, AppError> {
382    validate_identifier_part(context, value)?;
383    Ok(format!("\"{}\"", value.replace('"', "\"\"")))
384}
385
386fn validate_identifier_part(context: &str, value: &str) -> Result<(), AppError> {
387    let trimmed = value.trim();
388    if trimmed.is_empty() {
389        return Err(AppError::Validation(format!(
390            "{context} identifier must not be empty"
391        )));
392    }
393
394    if value.len() != trimmed.len() {
395        return Err(AppError::Validation(format!(
396            "{context} identifier must not have leading or trailing whitespace"
397        )));
398    }
399
400    if value.chars().any(char::is_control) {
401        return Err(AppError::Validation(format!(
402            "{context} identifier must not contain control characters"
403        )));
404    }
405
406    Ok(())
407}
408
409fn infrastructure_error(context: &str, error: impl std::fmt::Display) -> AppError {
410    AppError::Infrastructure(format!("{context}: {error}"))
411}
412
413#[cfg(test)]
414mod tests {
415    use ahri_tre_tabular::{TabularColumn, record_batch_to_columns};
416    use ahri_tre_types::{StudyId, VersionId};
417
418    use super::*;
419
420    #[test]
421    fn governed_lake_query_supports_projection_filter_order_and_limit() {
422        let connection = fixture_connection();
423        let request = LakeQueryRequest::new(
424            query_context(),
425            LakeRelation::table("patients").expect("valid relation"),
426        )
427        .with_projection(["participant_id", "age"])
428        .with_filters([
429            LakeQueryFilter::compare(
430                "age",
431                LakeQueryComparison::GreaterThanOrEqual,
432                LakeQueryValue::Int64(18),
433            ),
434            LakeQueryFilter::compare(
435                "consented",
436                LakeQueryComparison::Equal,
437                LakeQueryValue::Bool(true),
438            ),
439        ])
440        .with_order_by([LakeQueryOrder::desc("age")])
441        .with_limit(1);
442
443        let result =
444            execute_lake_query(&connection, &request, &AllowQuery).expect("query should execute");
445
446        assert_eq!(
447            result.generated_sql,
448            "SELECT \"participant_id\", \"age\" FROM \"patients\" WHERE \"age\" >= ? AND \"consented\" = ? ORDER BY \"age\" DESC LIMIT 1"
449        );
450        assert_eq!(result.table.row_count(), 1);
451        assert_eq!(result.table.column_count(), 2);
452
453        let columns =
454            record_batch_to_columns(&result.table.batches()[0]).expect("columns should decode");
455        assert_eq!(
456            columns,
457            vec![
458                TabularColumn::Utf8 {
459                    name: "participant_id".to_string(),
460                    values: vec![Some("p2".to_string())],
461                },
462                TabularColumn::Int64 {
463                    name: "age".to_string(),
464                    values: vec![Some(44)],
465                },
466            ]
467        );
468    }
469
470    #[test]
471    fn lake_query_authorizer_can_block_execution_before_duckdb_runs() {
472        let connection = Connection::open_in_memory().expect("in-memory DuckDB should open");
473        let request = LakeQueryRequest::new(
474            query_context(),
475            LakeRelation::table("missing_table").expect("valid relation"),
476        );
477
478        let result = execute_lake_query(&connection, &request, &DenyQuery);
479
480        assert!(
481            matches!(result, Err(AppError::Validation(message)) if message == "study access denied")
482        );
483    }
484
485    #[test]
486    fn lake_query_rejects_null_comparison_without_is_null_operator() {
487        let connection = fixture_connection();
488        let request = LakeQueryRequest::new(
489            query_context(),
490            LakeRelation::table("patients").expect("valid relation"),
491        )
492        .with_filters([LakeQueryFilter::compare(
493            "age",
494            LakeQueryComparison::Equal,
495            LakeQueryValue::Null,
496        )]);
497
498        let result = execute_lake_query(&connection, &request, &AllowQuery);
499
500        assert!(
501            matches!(result, Err(AppError::Validation(message)) if message.contains("NULL comparisons"))
502        );
503    }
504
505    fn fixture_connection() -> Connection {
506        let connection = Connection::open_in_memory().expect("in-memory DuckDB should open");
507        connection
508            .execute_batch(
509                r#"
510                CREATE TABLE patients(
511                    participant_id VARCHAR,
512                    age BIGINT,
513                    consented BOOLEAN
514                );
515                INSERT INTO patients VALUES
516                    ('p1', 17, true),
517                    ('p2', 44, true),
518                    ('p3', 33, false);
519                "#,
520            )
521            .expect("fixture table should be created");
522        connection
523    }
524
525    fn query_context() -> LakeQueryContext {
526        LakeQueryContext::new(StudyId(Default::default()))
527            .for_dataset(VersionId(Default::default()))
528            .for_principal("analyst")
529    }
530
531    struct AllowQuery;
532
533    impl LakeQueryAuthorizer for AllowQuery {
534        fn authorize_lake_query(&self, _request: &LakeQueryRequest) -> Result<(), AppError> {
535            Ok(())
536        }
537    }
538
539    struct DenyQuery;
540
541    impl LakeQueryAuthorizer for DenyQuery {
542        fn authorize_lake_query(&self, _request: &LakeQueryRequest) -> Result<(), AppError> {
543            Err(AppError::Validation("study access denied".to_string()))
544        }
545    }
546}