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}