use crate::{
AppConfig,
ai::rig::generate_sql_query,
db::{DatabaseInfo, DbPool, PoolHandler, QueryResult, TableInfo, TableSchema},
error::AppError,
state::AppState,
};
use axum::{
Json,
extract::{Path, State},
};
use serde::{Deserialize, Serialize};
use serde_json::{Value, json};
use std::sync::Arc;
use tracing::{info, instrument};
#[derive(Serialize, Clone, Debug)]
pub struct FullSchema {
pub databases: Vec<DatabaseSchema>,
}
#[derive(Serialize, Clone, Debug)]
pub struct DatabaseSchema {
pub name: String,
pub db_type: String,
pub tables: Vec<TableSchema>,
}
#[derive(Deserialize, Debug)]
pub struct GenerateQueryRequest {
pub db_name: String,
pub prompt: String,
}
#[derive(Serialize)]
pub struct GenerateQueryResponse {
pub query: String,
}
#[derive(Deserialize)]
pub struct ExecuteQueryRequest {
pub db_name: String,
pub query: String,
pub limit: Option<usize>,
}
#[derive(Serialize, Debug)]
pub struct ApiQueryResult {
result: Value, message: Option<String>, affected_rows: Option<i64>, plan: Option<Value>, #[serde(rename = "executionTime")] execution_time: f64, }
pub async fn ping() -> Json<Value> {
Json(json!({ "message": "pong" }))
}
pub async fn list_databases(State(state): State<AppState>) -> Json<Vec<DatabaseInfo>> {
let databases_info: Vec<DatabaseInfo> = state
.config
.databases
.iter()
.map(|db_config| DatabaseInfo {
name: db_config.name.clone(),
db_type: db_config.db_type.to_string(), })
.collect();
Json(databases_info)
}
pub async fn list_tables(
State(state): State<AppState>,
Path(db_name): Path<String>,
) -> Result<Json<Vec<TableInfo>>, AppError> {
let pools = state.pools.pin_owned();
let pool = pools
.get(&db_name)
.ok_or_else(|| AppError::NotFound(format!("Database '{}' not found", db_name)))?;
let tables = pool.list_tables().await?;
Ok(Json(tables))
}
pub async fn get_table_schema(
State(state): State<AppState>,
Path((db_name, table_name)): Path<(String, String)>,
) -> Result<Json<TableSchema>, AppError> {
let pools = state.pools.pin_owned();
let pool = pools
.get(&db_name)
.ok_or_else(|| AppError::NotFound(format!("Database '{}' not found", db_name)))?;
let schema = pool.get_table_schema(&table_name).await?;
Ok(Json(schema))
}
pub async fn execute_query(
State(state): State<AppState>,
Json(payload): Json<ExecuteQueryRequest>,
) -> Result<Json<ApiQueryResult>, AppError> {
let db_name = payload.db_name;
let limit = payload.limit;
let pools = state.pools.pin_owned();
let pool = pools
.get(&db_name)
.ok_or_else(|| AppError::NotFound(format!("Database '{}' not found", db_name)))?;
let query_result: QueryResult = pool.execute_query(&payload.query, limit).await?;
let api_response = ApiQueryResult {
result: query_result.data,
message: None,
affected_rows: None,
plan: query_result.plan,
execution_time: query_result.execution_time.as_secs_f64(),
};
Ok(Json(api_response))
}
pub async fn gen_query(
State(state): State<AppState>,
Json(payload): Json<GenerateQueryRequest>,
) -> Result<Json<GenerateQueryResponse>, AppError> {
info!(
"Received request to generate query for database: {}",
payload.db_name
);
let Json(schema) = get_full_schema(State(state.clone())).await?;
let generated_sql = generate_sql_query(
&state.openai_client,
&payload.db_name,
&schema,
&payload.prompt,
)
.await?;
Ok(Json(GenerateQueryResponse {
query: generated_sql,
}))
}
const SCHEMA_CACHE_KEY: &str = "full_schema";
#[instrument(skip(pools, config))] async fn fetch_full_schema_impl(
pools: Arc<papaya::HashMap<String, DbPool>>,
config: &AppConfig,
) -> Result<FullSchema, AppError> {
info!("Fetching full schema from databases...");
let mut database_schemas = Vec::new();
for db_config in &config.databases {
let db_name = &db_config.name;
info!(database = %db_name, "Fetching schema for database");
let result = async {
let pools_map = pools.pin_owned();
let pool = pools_map.get(db_name).ok_or_else(|| {
AppError::NotFound(format!("Pool not found for configured DB: {}", db_name))
})?;
let tables_info = pool.list_tables().await?;
let mut table_schemas = Vec::with_capacity(tables_info.len());
for table_info in tables_info {
info!(database = %db_name, table = %table_info.name, "Fetching schema for table");
match pool.get_table_schema(&table_info.name).await {
Ok(schema) => table_schemas.push(schema),
Err(e) => {
tracing::error!(
database = %db_name,
table = %table_info.name,
error = ?e,
"Failed to fetch schema for table, skipping."
);
}
}
}
Result::<_, AppError>::Ok(DatabaseSchema {
name: db_name.clone(),
db_type: db_config.db_type.to_string(),
tables: table_schemas,
})
}
.await;
match result {
Ok(db_schema) => database_schemas.push(db_schema),
Err(e) => {
tracing::error!(database = %db_name, error = ?e, "Failed to fetch schema for database, skipping.");
}
}
}
info!(
"Finished fetching schemas ({} successful).",
database_schemas.len()
);
Ok(FullSchema {
databases: database_schemas,
})
}
pub async fn get_full_schema(State(state): State<AppState>) -> Result<Json<FullSchema>, AppError> {
let cached_result_arc = state
.schema_cache
.get_with(SCHEMA_CACHE_KEY.to_string(), async {
let pools = Arc::clone(&state.pools);
let result = fetch_full_schema_impl(pools, &state.config).await;
Arc::new(result)
})
.await;
match &*cached_result_arc {
Ok(schema) => Ok(Json(schema.clone())), Err(e) => Err(e.clone_internal_error()), }
}
impl AppError {
fn clone_internal_error(&self) -> AppError {
match self {
AppError::Auth(e) => AppError::Auth((*e).clone()), AppError::Database(_) => AppError::Database(sqlx::Error::PoolClosed), AppError::UnsupportedDatabaseType(s) => AppError::UnsupportedDatabaseType(s.clone()),
AppError::Config(_) => {
AppError::Config(config::ConfigError::NotFound("cached config error".into()))
} AppError::NotFound(s) => AppError::NotFound(s.clone()),
AppError::NotImplemented(s) => AppError::NotImplemented(s.clone()),
AppError::BadRequest(s) => AppError::BadRequest(s.clone()),
AppError::SqlParsingError(s) => AppError::SqlParsingError(s.clone()),
AppError::InvalidQueryResult(s) => AppError::InvalidQueryResult(s.clone()),
AppError::AiError(e) => AppError::AiError((*e).clone()),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{
AppConfig,
config::DatabaseConfig,
db::{ColumnInfo, ColumnType, DatabaseType, TableType},
state::AppState,
};
use axum::{Json, extract::State};
#[derive(Deserialize)]
struct User {
id: i32,
name: String,
email: String,
#[allow(dead_code)]
password: String,
}
#[tokio::test]
async fn test_list_databases() {
let mock_db_config1 = DatabaseConfig {
name: "mock_db1".to_string(),
db_type: DatabaseType::Postgres,
conn_string: "postgresql://user:pass@host:port/db1".to_string(),
};
let mock_db_config2 = DatabaseConfig {
name: "mock_db2".to_string(),
db_type: DatabaseType::Mysql,
conn_string: "mysql://user:pass@host:port/db2".to_string(),
};
let mock_config = AppConfig {
server_addr: "127.0.0.1:8080".to_string(),
databases: vec![mock_db_config1, mock_db_config2],
jwt_secret: "test_secret".to_string(),
allowed_origin: "*".to_string(),
};
let state = AppState::new_for_test(mock_config);
let Json(response) = list_databases(State(state)).await;
assert_eq!(response.len(), 2);
assert_eq!(response[0].name, "mock_db1");
assert_eq!(response[0].db_type, "postgres"); assert_eq!(response[1].name, "mock_db2");
assert_eq!(response[1].db_type, "mysql"); }
#[tokio::test]
async fn test_list_tables() {
let state = AppState::new(AppConfig::load().unwrap()).await.unwrap();
let Json(response) = list_tables(State(state), Path("users".to_string()))
.await
.unwrap();
println!("response: {:?}", response);
assert_eq!(response.len(), 5);
assert_eq!(response[0].name, "public.repositories");
assert_eq!(response[0].table_type, TableType::Table);
}
#[tokio::test]
async fn test_get_table_schema() {
let state = AppState::new(AppConfig::load().unwrap()).await.unwrap();
let Json(response) = get_table_schema(
State(state),
Path(("users".to_string(), "repository_members".to_string())),
)
.await
.unwrap();
assert_eq!(response.columns.len(), 3);
assert_eq!(response.columns[0].name, "id");
assert_eq!(response.columns[0].data_type, ColumnType::Integer);
assert!(!response.columns[0].is_nullable);
assert!(response.columns[0].is_pk);
assert!(response.columns[0].is_unique);
assert_eq!(response.columns[1].name, "repository_id");
assert_eq!(response.columns[1].data_type, ColumnType::Integer);
assert!(!response.columns[1].is_nullable);
assert!(!response.columns[1].is_pk);
assert!(response.columns[1].is_unique);
assert_eq!(
response.columns[1].fk_table,
Some("repositories".to_string())
);
assert_eq!(response.columns[1].fk_column, Some("id".to_string()));
assert_eq!(response.columns[2].name, "user_id");
assert_eq!(response.columns[2].data_type, ColumnType::Integer);
assert!(!response.columns[2].is_nullable);
assert!(!response.columns[2].is_pk);
assert!(response.columns[2].is_unique);
assert_eq!(response.columns[2].fk_table, Some("users".to_string()));
assert_eq!(response.columns[2].fk_column, Some("id".to_string()));
}
#[tokio::test]
async fn test_execute_query() {
let state = AppState::new(AppConfig::load().unwrap()).await.unwrap();
let Json(data) = execute_query(
State(state),
Json(ExecuteQueryRequest {
db_name: "users".to_string(),
query: "SELECT * FROM users".to_string(),
limit: None,
}),
)
.await
.unwrap();
println!("data: {:?}", data);
let users: Vec<User> = serde_json::from_value(data.result).unwrap();
assert_eq!(users[0].id, 1);
assert_eq!(users[0].name, "Alice Johnson");
assert_eq!(users[0].email, "alice@example.com");
}
#[ignore]
#[tokio::test]
async fn test_gen_query_placeholder() {
let state = AppState::new(AppConfig::load().unwrap()).await.unwrap();
let payload = GenerateQueryRequest {
db_name: "users".to_string(),
prompt: "show me all users".to_string(),
};
let result = gen_query(State(state), Json(payload)).await;
assert!(result.is_ok());
let Json(res) = result.unwrap();
assert!(res.query.contains("SELECT"));
assert!(res.query.contains("FROM"));
}
#[tokio::test]
async fn test_gen_query_handler_success() {
let state = AppState::new(AppConfig::load().unwrap()).await.unwrap();
let mock_db_schema = DatabaseSchema {
name: "test_db".to_string(),
db_type: "postgresql".to_string(),
tables: vec![TableSchema {
table_name: "items".to_string(),
columns: vec![ColumnInfo {
name: "id".to_string(),
data_type: ColumnType::Integer,
is_nullable: false,
is_pk: true,
is_unique: false,
fk_table: None,
fk_column: None,
}],
}],
};
let mock_full_schema = FullSchema {
databases: vec![mock_db_schema],
};
state
.schema_cache
.insert(
SCHEMA_CACHE_KEY.to_string(),
Arc::new(Ok(mock_full_schema.clone())), )
.await;
let _payload = GenerateQueryRequest {
db_name: "test_db".to_string(), prompt: "show me all items".to_string(),
};
let mock_generated_sql = "SELECT * FROM items;".to_string();
let result: Result<Json<GenerateQueryResponse>, AppError> =
Ok(Json(GenerateQueryResponse {
query: mock_generated_sql,
}));
assert!(result.is_ok());
#[allow(clippy::unnecessary_literal_unwrap)]
let Json(response) = result.unwrap();
assert_eq!(response.query, "SELECT * FROM items;");
}
#[tokio::test]
async fn test_gen_query_handler_ai_error() {
let state = AppState::new(AppConfig::load().unwrap()).await.unwrap();
let mock_db_schema = DatabaseSchema {
name: "test_db".to_string(),
db_type: "postgresql".to_string(),
tables: vec![TableSchema {
table_name: "items".to_string(),
columns: vec![ColumnInfo {
name: "id".to_string(),
data_type: ColumnType::Integer,
is_nullable: false,
is_pk: true,
is_unique: false,
fk_table: None,
fk_column: None,
}],
}],
};
let mock_full_schema = FullSchema {
databases: vec![mock_db_schema],
};
state
.schema_cache
.insert(
SCHEMA_CACHE_KEY.to_string(),
Arc::new(Ok(mock_full_schema.clone())),
)
.await;
let _payload = GenerateQueryRequest {
db_name: "test_db".to_string(),
prompt: "some failing prompt".to_string(),
};
let mock_ai_error = AppError::AiError("AI failed to generate query".to_string());
let result: Result<Json<GenerateQueryResponse>, AppError> =
Err(mock_ai_error.clone_internal_error());
assert!(result.is_err());
match result.err().unwrap() {
AppError::AiError(msg) => {
assert_eq!(msg, "AI failed to generate query");
}
e => panic!("Expected AiError, got {:?}", e),
}
}
}