use crate::db::ConnectionManager;
use crate::error::AppError;
use rmcp::handler::server::tool::ToolRouter;
use rmcp::handler::server::wrapper::Parameters;
use rmcp::model::*;
use rmcp::{tool, tool_handler, tool_router};
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use serde_json::Value as JsonValue;
use std::sync::Arc;
use tracing::{info, warn};
#[derive(Clone)]
pub struct SqlxMcpServer {
conn_manager: Arc<ConnectionManager>,
tool_router: ToolRouter<Self>,
}
#[derive(Debug, Deserialize, Serialize, JsonSchema)]
pub struct QueryParams {
pub query: String,
#[serde(default)]
pub params: Vec<JsonValue>,
#[serde(default)]
pub connection: Option<String>,
}
#[derive(Debug, Deserialize, Serialize, JsonSchema)]
pub struct InsertParams {
pub query: String,
#[serde(default)]
pub params: Vec<JsonValue>,
#[serde(default)]
pub connection: Option<String>,
}
#[derive(Debug, Deserialize, Serialize, JsonSchema)]
pub struct UpdateParams {
pub query: String,
#[serde(default)]
pub params: Vec<JsonValue>,
#[serde(default)]
pub connection: Option<String>,
}
#[derive(Debug, Deserialize, Serialize, JsonSchema)]
pub struct DeleteParams {
pub query: String,
#[serde(default)]
pub params: Vec<JsonValue>,
#[serde(default)]
pub connection: Option<String>,
}
#[derive(Debug, Deserialize, Serialize, JsonSchema)]
pub struct DescribeTableParams {
pub table: String,
#[serde(default)]
pub connection: Option<String>,
}
#[derive(Debug, Deserialize, Serialize, JsonSchema)]
pub struct ListTablesParams {
#[serde(default)]
pub connection: Option<String>,
}
#[derive(Debug, Deserialize, Serialize, JsonSchema)]
pub struct HealthCheckParams {
#[serde(default)]
pub connection: Option<String>,
}
#[tool_router]
impl SqlxMcpServer {
pub fn new(conn_manager: ConnectionManager) -> Self {
Self {
conn_manager: Arc::new(conn_manager),
tool_router: Self::tool_router(),
}
}
#[tool(
name = "db_list_connections",
description = "List all configured database connections with their engine types (MySQL, PostgreSQL, SQLite)."
)]
async fn list_connections(&self) -> Result<CallToolResult, ErrorData> {
info!("Executing list_connections tool");
let connections = self.conn_manager.list_connections();
let response = serde_json::json!({
"connections": connections,
"count": connections.len()
});
Ok(CallToolResult::success(vec![Content::text(
serde_json::to_string_pretty(&response).unwrap_or_else(|_| response.to_string()),
)]))
}
#[tool(
name = "db_query",
description = "Execute a SELECT query on a database. Use ? placeholders for parameters to prevent SQL injection (automatically converted to $1,$2,... for PostgreSQL). Returns results as JSON array. Specify 'connection' to choose a database, or omit to use default."
)]
async fn query(&self, params: Parameters<QueryParams>) -> Result<CallToolResult, ErrorData> {
info!("Executing query tool");
match self
.conn_manager
.query(
params.0.connection.as_deref(),
¶ms.0.query,
params.0.params.clone(),
)
.await
{
Ok(results) => {
let json = serde_json::to_string_pretty(&results).unwrap_or_else(|e| {
format!("{{\"error\": \"Failed to serialize results: {}\"}}", e)
});
Ok(CallToolResult::success(vec![Content::text(json)]))
}
Err(e) => {
warn!("Query failed: {}", e);
Err(app_error_to_mcp(&e))
}
}
}
#[tool(
name = "db_insert",
description = "Execute an INSERT query on a database. Use ? placeholders for parameters. Returns the last insert ID. For PostgreSQL, automatically appends RETURNING id if not present. Specify 'connection' to choose a database."
)]
async fn insert(&self, params: Parameters<InsertParams>) -> Result<CallToolResult, ErrorData> {
info!("Executing insert tool");
match self
.conn_manager
.insert(
params.0.connection.as_deref(),
¶ms.0.query,
params.0.params.clone(),
)
.await
{
Ok(last_id) => {
let response = serde_json::json!({
"success": true,
"last_insert_id": last_id
});
Ok(CallToolResult::success(vec![Content::text(
response.to_string(),
)]))
}
Err(e) => {
warn!("Insert failed: {}", e);
Err(app_error_to_mcp(&e))
}
}
}
#[tool(
name = "db_update",
description = "Execute an UPDATE query on a database. Use ? placeholders for parameters. Returns the number of affected rows. Specify 'connection' to choose a database."
)]
async fn update(&self, params: Parameters<UpdateParams>) -> Result<CallToolResult, ErrorData> {
info!("Executing update tool");
match self
.conn_manager
.update(
params.0.connection.as_deref(),
¶ms.0.query,
params.0.params.clone(),
)
.await
{
Ok(rows_affected) => {
let response = serde_json::json!({
"success": true,
"rows_affected": rows_affected
});
Ok(CallToolResult::success(vec![Content::text(
response.to_string(),
)]))
}
Err(e) => {
warn!("Update failed: {}", e);
Err(app_error_to_mcp(&e))
}
}
}
#[tool(
name = "db_delete",
description = "Execute a DELETE query on a database. Use ? placeholders for parameters. Returns the number of affected rows. Specify 'connection' to choose a database."
)]
async fn delete(&self, params: Parameters<DeleteParams>) -> Result<CallToolResult, ErrorData> {
info!("Executing delete tool");
match self
.conn_manager
.delete(
params.0.connection.as_deref(),
¶ms.0.query,
params.0.params.clone(),
)
.await
{
Ok(rows_affected) => {
let response = serde_json::json!({
"success": true,
"rows_affected": rows_affected
});
Ok(CallToolResult::success(vec![Content::text(
response.to_string(),
)]))
}
Err(e) => {
warn!("Delete failed: {}", e);
Err(app_error_to_mcp(&e))
}
}
}
#[tool(
name = "db_list_tables",
description = "List all tables in the database. Works with MySQL (SHOW TABLES), PostgreSQL (pg_tables), and SQLite (sqlite_master). Specify 'connection' to choose a database."
)]
async fn list_tables(
&self,
params: Parameters<ListTablesParams>,
) -> Result<CallToolResult, ErrorData> {
info!("Executing list_tables tool");
match self
.conn_manager
.list_tables(params.0.connection.as_deref())
.await
{
Ok(tables) => {
let response = serde_json::json!({
"tables": tables,
"count": tables.len()
});
Ok(CallToolResult::success(vec![Content::text(
serde_json::to_string_pretty(&response).unwrap_or_else(|_| response.to_string()),
)]))
}
Err(e) => {
warn!("List tables failed: {}", e);
Err(app_error_to_mcp(&e))
}
}
}
#[tool(
name = "db_describe_table",
description = "Describe the structure of a database table. Returns column names, types, keys, and other metadata. Works across MySQL, PostgreSQL, and SQLite. Specify 'connection' to choose a database."
)]
async fn describe_table(
&self,
params: Parameters<DescribeTableParams>,
) -> Result<CallToolResult, ErrorData> {
info!("Executing describe_table tool for: {}", params.0.table);
match self
.conn_manager
.describe_table(params.0.connection.as_deref(), ¶ms.0.table)
.await
{
Ok(columns) => {
let response = serde_json::json!({
"table": params.0.table,
"columns": columns
});
Ok(CallToolResult::success(vec![Content::text(
serde_json::to_string_pretty(&response).unwrap_or_else(|_| response.to_string()),
)]))
}
Err(e) => {
warn!("Describe table failed: {}", e);
Err(app_error_to_mcp(&e))
}
}
}
#[tool(
name = "db_health_check",
description = "Check database connectivity and health status. Specify 'connection' to check a specific database, or omit to check the default."
)]
async fn health_check(
&self,
params: Parameters<HealthCheckParams>,
) -> Result<CallToolResult, ErrorData> {
info!("Executing health_check tool");
match self
.conn_manager
.health_check(params.0.connection.as_deref())
.await
{
Ok(_) => {
let engine = self
.conn_manager
.get_engine(params.0.connection.as_deref())
.map(|e| e.to_string())
.unwrap_or_else(|_| "unknown".to_string());
let response = serde_json::json!({
"status": "healthy",
"connected": true,
"engine": engine
});
Ok(CallToolResult::success(vec![Content::text(
response.to_string(),
)]))
}
Err(e) => {
let response = serde_json::json!({
"status": "unhealthy",
"connected": false,
"error": e.to_string()
});
Ok(CallToolResult::success(vec![Content::text(
response.to_string(),
)]))
}
}
}
}
#[tool_handler(router = self.tool_router)]
impl rmcp::ServerHandler for SqlxMcpServer {
fn get_info(&self) -> ServerInfo {
ServerInfo {
protocol_version: ProtocolVersion::V_2024_11_05,
server_info: Implementation {
name: env!("CARGO_PKG_NAME").to_string(),
version: env!("CARGO_PKG_VERSION").to_string(),
..Default::default()
},
capabilities: ServerCapabilities::builder().enable_tools().build(),
instructions: Some(
r#"SQLx MCP Server - Multi-database operations via Model Context Protocol.
Supported databases: MySQL, PostgreSQL, SQLite
Available tools:
- db_list_connections: List all configured database connections
- db_query: Execute SELECT queries with parameterized inputs
- db_insert: Execute INSERT queries, returns last insert ID
- db_update: Execute UPDATE queries, returns affected rows
- db_delete: Execute DELETE queries, returns affected rows
- db_list_tables: List all tables in the database
- db_describe_table: Get table structure and column info
- db_health_check: Check database connectivity
All tools accept a 'connection' parameter to specify which database to use.
If omitted, the default connection is used.
Security features:
- All queries use parameterized statements (? placeholders, converted to $1,$2,... for PostgreSQL)
- Dangerous operations (DROP, TRUNCATE, ALTER, etc.) are blocked
- Multiple statements in single query are blocked
- Query type validation (SELECT for query, INSERT for insert, etc.)
Example usage:
- Query: {"query": "SELECT * FROM users WHERE id = ?", "params": [1], "connection": "mydb"}
- Insert: {"query": "INSERT INTO users (name) VALUES (?)", "params": ["John"]}"#
.to_string(),
),
}
}
}
fn app_error_to_mcp(error: &AppError) -> ErrorData {
error.to_mcp_error()
}