use super::traits::{Service, ServiceError, ServiceResult};
use crate::core::models::query::{NativeQuery, QueryResult};
use crate::core::models::DatasetQuery;
use crate::repository::query::QueryRepository;
use async_trait::async_trait;
use serde_json::Value;
use std::collections::HashMap;
use std::sync::Arc;
#[async_trait]
pub trait QueryService: Service {
async fn execute_dataset_query(&self, query: DatasetQuery) -> ServiceResult<QueryResult>;
async fn execute_native_query(
&self,
database_id: i32,
query: NativeQuery,
) -> ServiceResult<QueryResult>;
async fn execute_raw_query(&self, query: Value) -> ServiceResult<QueryResult>;
async fn execute_pivot_query(&self, query: Value) -> ServiceResult<QueryResult>;
async fn execute_sql_with_params(
&self,
database_id: i32,
sql: &str,
params: HashMap<String, serde_json::Value>,
) -> ServiceResult<QueryResult>;
async fn execute_sql(&self, database_id: i32, sql: &str) -> ServiceResult<QueryResult>;
async fn export_query(&self, format: &str, query: Value) -> ServiceResult<Vec<u8>>;
async fn validate_query(&self, sql: &str) -> ServiceResult<()>;
async fn validate_dataset_query(&self, query: &DatasetQuery) -> ServiceResult<()>;
}
pub struct HttpQueryService {
repository: Arc<dyn QueryRepository>,
}
impl HttpQueryService {
pub fn new(repository: Arc<dyn QueryRepository>) -> Self {
Self { repository }
}
fn validate_sql(&self, sql: &str) -> ServiceResult<()> {
if sql.trim().is_empty() {
return Err(ServiceError::Validation(
"SQL query cannot be empty".to_string(),
));
}
let sql_upper = sql.to_uppercase();
let dangerous_keywords = [
"DROP", "DELETE", "TRUNCATE", "ALTER", "CREATE", "GRANT", "REVOKE",
];
for keyword in &dangerous_keywords {
if sql_upper.contains(keyword) {
if !sql_upper.contains(&format!("--{}", keyword))
&& !sql_upper.contains(&format!("/*{}*/", keyword))
{
return Err(ServiceError::BusinessRule(format!(
"Query contains potentially dangerous operation: {}",
keyword
)));
}
}
}
Ok(())
}
}
#[async_trait]
impl Service for HttpQueryService {
fn name(&self) -> &str {
"QueryService"
}
}
#[async_trait]
impl QueryService for HttpQueryService {
async fn execute_dataset_query(&self, query: DatasetQuery) -> ServiceResult<QueryResult> {
self.validate_dataset_query(&query).await?;
self.repository
.execute_dataset_query(query)
.await
.map_err(ServiceError::from)
}
async fn execute_native_query(
&self,
database_id: i32,
query: NativeQuery,
) -> ServiceResult<QueryResult> {
self.validate_sql(&query.query)?;
self.repository
.execute_native_query(database_id, query)
.await
.map_err(ServiceError::from)
}
async fn execute_raw_query(&self, query: Value) -> ServiceResult<QueryResult> {
if !query.is_object() {
return Err(ServiceError::Validation(
"Query must be a JSON object".to_string(),
));
}
self.repository
.execute_raw_query(query)
.await
.map_err(ServiceError::from)
}
async fn execute_pivot_query(&self, query: Value) -> ServiceResult<QueryResult> {
if !query.is_object() {
return Err(ServiceError::Validation(
"Query must be a JSON object".to_string(),
));
}
self.repository
.execute_pivot_query(query)
.await
.map_err(ServiceError::from)
}
async fn execute_sql_with_params(
&self,
database_id: i32,
sql: &str,
params: HashMap<String, serde_json::Value>,
) -> ServiceResult<QueryResult> {
self.validate_sql(sql)?;
let query = NativeQuery::builder(sql).with_params(params).build();
self.repository
.execute_native_query(database_id, query)
.await
.map_err(ServiceError::from)
}
async fn execute_sql(&self, database_id: i32, sql: &str) -> ServiceResult<QueryResult> {
self.validate_sql(sql)?;
let query = NativeQuery::builder(sql).build();
self.repository
.execute_native_query(database_id, query)
.await
.map_err(ServiceError::from)
}
async fn export_query(&self, format: &str, query: Value) -> ServiceResult<Vec<u8>> {
let valid_formats = ["csv", "json", "xlsx"];
if !valid_formats.contains(&format) {
return Err(ServiceError::Validation(format!(
"Invalid export format: {}. Must be one of: csv, json, xlsx",
format
)));
}
if !query.is_object() {
return Err(ServiceError::Validation(
"Query must be a JSON object".to_string(),
));
}
self.repository
.export_query(format, query)
.await
.map_err(ServiceError::from)
}
async fn validate_query(&self, sql: &str) -> ServiceResult<()> {
self.validate_sql(sql)
}
async fn validate_dataset_query(&self, query: &DatasetQuery) -> ServiceResult<()> {
if query.database.0 < 1 {
return Err(ServiceError::Validation(
"Invalid database ID: must be positive".to_string(),
));
}
if query.query.is_null() {
return Err(ServiceError::Validation(
"Query content cannot be empty".to_string(),
));
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::models::query::QueryStatus;
use crate::core::models::MetabaseId;
use crate::repository::query::{MockQueryRepository, QueryRepository};
use serde_json::json;
use std::sync::Arc;
#[tokio::test]
async fn test_execute_dataset_query() {
let repository = Arc::new(MockQueryRepository::new()) as Arc<dyn QueryRepository>;
let service = HttpQueryService::new(repository);
let query = DatasetQuery {
database: MetabaseId(1),
query_type: "native".to_string(),
query: json!({"query": "SELECT * FROM users"}),
parameters: None,
constraints: None,
};
let result = service.execute_dataset_query(query).await;
assert!(result.is_ok());
let query_result = result.unwrap();
assert_eq!(query_result.status, QueryStatus::Completed);
}
#[tokio::test]
async fn test_validate_sql() {
let repository = Arc::new(MockQueryRepository::new()) as Arc<dyn QueryRepository>;
let service = HttpQueryService::new(repository);
assert!(service.validate_sql("SELECT * FROM users").is_ok());
assert!(service.validate_sql("").is_err());
assert!(service.validate_sql("DROP TABLE users").is_err());
}
#[tokio::test]
async fn test_export_query() {
let repository = Arc::new(MockQueryRepository::new()) as Arc<dyn QueryRepository>;
let service = HttpQueryService::new(repository);
let query = json!({
"database": 1,
"type": "native",
"native": {"query": "SELECT * FROM products"}
});
let result = service.export_query("csv", query.clone()).await;
assert!(result.is_ok());
let result = service.export_query("invalid", query).await;
assert!(result.is_err());
}
}