use crate::advice_common::{AdviceType, AiAdviceAuditRecord, BenefitEstimate};
use serde::{Deserialize, Serialize};
use sqlparser::ast::{
Expr, JoinConstraint, JoinOperator, OrderByExpr, SetExpr, Statement, TableFactor,
TableWithJoins,
};
use sqlparser::dialect::GenericDialect;
use sqlparser::parser::Parser;
use thiserror::Error;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum IndexType {
BTree,
Hash,
Gin,
Brin,
}
impl IndexType {
pub fn ddl_keyword(&self) -> &str {
match self {
IndexType::BTree => "BTREE",
IndexType::Hash => "HASH",
IndexType::Gin => "GIN",
IndexType::Brin => "BRIN",
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct QueryPattern {
pub sql_template: String,
pub frequency: u64,
pub columns_accessed: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SlowQueryLog {
pub sql: String,
pub execution_time_ms: u64,
pub timestamp: i64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct IndexSuggestion {
pub index_columns: Vec<String>,
pub index_type: IndexType,
pub ddl_text: String,
pub expected_benefit: BenefitEstimate,
pub evidence: Vec<QueryPattern>,
}
#[derive(Debug, Error)]
pub enum IndexError {
#[error("SQL parse error: {0}")]
ParseError(String),
#[error("No query patterns provided")]
NoQueryPatterns,
#[error("LLM service unavailable: {0}")]
LlmServiceUnavailable(String),
}
pub struct IndexAdvisor {
llm_enabled: bool,
}
impl Default for IndexAdvisor {
fn default() -> Self {
Self::new()
}
}
impl IndexAdvisor {
pub fn new() -> Self {
Self { llm_enabled: false }
}
pub fn with_llm(mut self) -> Self {
self.llm_enabled = true;
self
}
pub async fn suggest(
&self,
query_patterns: &[QueryPattern],
slow_queries: &[SlowQueryLog],
) -> Result<Vec<IndexSuggestion>, IndexError> {
if query_patterns.is_empty() {
return Err(IndexError::NoQueryPatterns);
}
let mut suggestions = Vec::new();
let dialect = GenericDialect {};
for pattern in query_patterns {
let parsed = Parser::parse_sql(&dialect, &pattern.sql_template);
if parsed.is_err() {
continue;
}
let statements = parsed.unwrap();
for stmt in &statements {
if let Some((table, filter_cols, join_cols, order_cols)) =
Self::extract_query_info(stmt)
{
let mut candidate_cols = Vec::new();
candidate_cols.extend(filter_cols);
candidate_cols.extend(join_cols);
if candidate_cols.is_empty() && order_cols.is_empty() {
continue;
}
candidate_cols.sort();
candidate_cols.dedup();
let index_type = IndexType::BTree;
let col_list = candidate_cols.join(", ");
let idx_name = format!("idx_{}_{}", table, candidate_cols.join("_"));
let ddl_text = format!(
"CREATE {} INDEX {} ON {} ({})",
index_type.ddl_keyword(),
idx_name,
table,
col_list
);
let total_frequency: u64 = query_patterns.iter().map(|p| p.frequency).sum();
let speedup_ratio = if pattern.frequency > 0 && total_frequency > 0 {
1.0 + (pattern.frequency as f64 / total_frequency as f64) * 10.0
} else {
1.0
};
let confidence = if self.llm_enabled { 0.85 } else { 0.7 };
let uncertain = slow_queries.is_empty();
let benefit = if uncertain {
BenefitEstimate::uncertain(speedup_ratio, confidence)
} else {
BenefitEstimate::certain(speedup_ratio, confidence)
};
suggestions.push(IndexSuggestion {
index_columns: candidate_cols.clone(),
index_type: index_type.clone(),
ddl_text,
expected_benefit: benefit,
evidence: vec![pattern.clone()],
});
}
}
}
suggestions.sort_by(|a, b| {
b.expected_benefit
.speedup_ratio
.partial_cmp(&a.expected_benefit.speedup_ratio)
.unwrap_or(std::cmp::Ordering::Equal)
});
Ok(suggestions)
}
pub fn audit_record(&self, confidence: f32) -> AiAdviceAuditRecord {
if self.llm_enabled {
AiAdviceAuditRecord::from_llm(AdviceType::Index, confidence, "gpt-4o-mini")
} else {
AiAdviceAuditRecord::from_rule(AdviceType::Index, confidence)
}
}
fn extract_query_info(
stmt: &Statement,
) -> Option<(String, Vec<String>, Vec<String>, Vec<String>)> {
let query = match stmt {
Statement::Query(q) => q.as_ref(),
_ => return None,
};
let select = match &*query.body {
SetExpr::Select(s) => s,
_ => return None,
};
let table = Self::extract_table_name(&select.from)?;
let filter_cols = Self::extract_filter_columns(&select.selection);
let join_cols = Self::extract_join_columns(&select.from);
let order_cols = Self::extract_order_columns(&query.order_by);
Some((table, filter_cols, join_cols, order_cols))
}
fn extract_table_name(from: &[TableWithJoins]) -> Option<String> {
if from.is_empty() {
return None;
}
match &from[0].relation {
TableFactor::Table { name, .. } => {
Some(name.0.last().map(|i| i.value.clone()).unwrap_or_default())
}
_ => None,
}
}
fn extract_filter_columns(selection: &Option<Expr>) -> Vec<String> {
let mut cols = Vec::new();
if let Some(expr) = selection {
Self::collect_columns(expr, &mut cols);
}
cols
}
fn extract_join_columns(from: &[TableWithJoins]) -> Vec<String> {
let mut cols = Vec::new();
for table_with_joins in from {
for join in &table_with_joins.joins {
match &join.join_operator {
JoinOperator::Inner(constraint)
| JoinOperator::LeftOuter(constraint)
| JoinOperator::RightOuter(constraint)
| JoinOperator::FullOuter(constraint) => {
if let JoinConstraint::On(expr) = constraint {
Self::collect_columns(expr, &mut cols);
}
}
_ => {}
}
}
}
cols
}
fn extract_order_columns(order_by: &[OrderByExpr]) -> Vec<String> {
let mut cols = Vec::new();
for ob in order_by {
Self::collect_columns(&ob.expr, &mut cols);
}
cols
}
fn collect_columns(expr: &Expr, cols: &mut Vec<String>) {
match expr {
Expr::Identifier(ident) => {
cols.push(ident.value.clone());
}
Expr::CompoundIdentifier(idents) => {
if let Some(last) = idents.last() {
cols.push(last.value.clone());
}
}
Expr::BinaryOp { left, right, .. } => {
Self::collect_columns(left, cols);
Self::collect_columns(right, cols);
}
Expr::InList { expr, .. } => {
Self::collect_columns(expr, cols);
}
Expr::Like { expr, pattern, .. } => {
Self::collect_columns(expr, cols);
Self::collect_columns(pattern, cols);
}
Expr::IsNull(expr) => {
Self::collect_columns(expr, cols);
}
Expr::Nested(expr) => {
Self::collect_columns(expr, cols);
}
Expr::UnaryOp { expr, .. } => {
Self::collect_columns(expr, cols);
}
_ => {}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_suggest_index_for_where_clause() {
let advisor = IndexAdvisor::new();
let patterns = vec![QueryPattern {
sql_template: "SELECT * FROM users WHERE email = $1".to_string(),
frequency: 100,
columns_accessed: vec!["email".to_string()],
}];
let result = advisor.suggest(&patterns, &[]).await.unwrap();
assert!(!result.is_empty());
assert!(result[0].ddl_text.contains("CREATE"));
assert!(result[0].ddl_text.contains("users"));
assert!(result[0].index_columns.contains(&"email".to_string()));
}
#[tokio::test]
async fn test_suggest_index_for_join() {
let advisor = IndexAdvisor::new();
let patterns = vec![QueryPattern {
sql_template: "SELECT * FROM orders o JOIN users u ON o.user_id = u.id".to_string(),
frequency: 50,
columns_accessed: vec!["user_id".to_string(), "id".to_string()],
}];
let result = advisor.suggest(&patterns, &[]).await.unwrap();
assert!(!result.is_empty());
assert!(result[0].index_columns.contains(&"user_id".to_string()));
}
#[tokio::test]
async fn test_suggest_index_for_order_by() {
let advisor = IndexAdvisor::new();
let patterns = vec![QueryPattern {
sql_template: "SELECT * FROM users ORDER BY created_at".to_string(),
frequency: 30,
columns_accessed: vec!["created_at".to_string()],
}];
let result = advisor.suggest(&patterns, &[]).await.unwrap();
assert!(!result.is_empty());
}
#[tokio::test]
async fn test_no_query_patterns_error() {
let advisor = IndexAdvisor::new();
let result = advisor.suggest(&[], &[]).await;
assert!(matches!(result, Err(IndexError::NoQueryPatterns)));
}
#[tokio::test]
async fn test_ddl_not_executed() {
let advisor = IndexAdvisor::new();
let patterns = vec![QueryPattern {
sql_template: "SELECT * FROM users WHERE email = $1".to_string(),
frequency: 100,
columns_accessed: vec!["email".to_string()],
}];
let result = advisor.suggest(&patterns, &[]).await.unwrap();
for s in &result {
assert!(s.ddl_text.starts_with("CREATE"));
assert!(!s.ddl_text.contains("EXECUTE"));
}
}
#[tokio::test]
async fn test_benefit_estimate_uncertain_without_slow_queries() {
let advisor = IndexAdvisor::new();
let patterns = vec![QueryPattern {
sql_template: "SELECT * FROM users WHERE email = $1".to_string(),
frequency: 100,
columns_accessed: vec!["email".to_string()],
}];
let result = advisor.suggest(&patterns, &[]).await.unwrap();
assert!(result[0].expected_benefit.uncertain);
}
#[tokio::test]
async fn test_benefit_estimate_certain_with_slow_queries() {
let advisor = IndexAdvisor::new();
let patterns = vec![QueryPattern {
sql_template: "SELECT * FROM users WHERE email = $1".to_string(),
frequency: 100,
columns_accessed: vec!["email".to_string()],
}];
let slow_queries = vec![SlowQueryLog {
sql: "SELECT * FROM users WHERE email = $1".to_string(),
execution_time_ms: 500,
timestamp: 1000,
}];
let result = advisor.suggest(&patterns, &slow_queries).await.unwrap();
assert!(!result[0].expected_benefit.uncertain);
}
#[tokio::test]
async fn test_suggestions_sorted_by_benefit() {
let advisor = IndexAdvisor::new();
let patterns = vec![
QueryPattern {
sql_template: "SELECT * FROM users WHERE email = $1".to_string(),
frequency: 100,
columns_accessed: vec!["email".to_string()],
},
QueryPattern {
sql_template: "SELECT * FROM users WHERE name = $1".to_string(),
frequency: 10,
columns_accessed: vec!["name".to_string()],
},
];
let result = advisor.suggest(&patterns, &[]).await.unwrap();
if result.len() >= 2 {
assert!(
result[0].expected_benefit.speedup_ratio
>= result[1].expected_benefit.speedup_ratio
);
}
}
#[test]
fn test_index_type_ddl_keyword() {
assert_eq!(IndexType::BTree.ddl_keyword(), "BTREE");
assert_eq!(IndexType::Hash.ddl_keyword(), "HASH");
assert_eq!(IndexType::Gin.ddl_keyword(), "GIN");
assert_eq!(IndexType::Brin.ddl_keyword(), "BRIN");
}
#[test]
fn test_audit_record() {
let advisor = IndexAdvisor::new();
let record = advisor.audit_record(0.8);
assert_eq!(record.advice_type, AdviceType::Index);
assert_eq!(
record.source_engine,
crate::advice_common::AdviceSource::Rule
);
}
}