use std::sync::{
Arc,
atomic::{AtomicU64, Ordering},
};
use fraiseql_core::{
db::traits::DatabaseAdapter,
runtime::Executor,
schema::CompiledSchema,
security::{OidcValidator, SecurityContext},
};
use rmcp::{
ServerHandler,
model::{
CallToolRequestParams, CallToolResult, ListToolsResult, ServerCapabilities, ServerInfo,
Tool,
},
service::RequestContext,
};
use super::{McpConfig, executor::error_result};
pub(crate) fn extract_bearer(headers: &http::HeaderMap) -> Option<String> {
let value = headers.get(http::header::AUTHORIZATION)?.to_str().ok()?;
let token = value.strip_prefix("Bearer ")?.trim();
if token.is_empty() {
None
} else {
Some(token.to_string())
}
}
pub static MCP_TOOL_CALLS_TOTAL: AtomicU64 = AtomicU64::new(0);
pub static MCP_TOOL_ERRORS_TOTAL: AtomicU64 = AtomicU64::new(0);
pub fn mcp_tool_calls_total() -> u64 {
MCP_TOOL_CALLS_TOTAL.load(Ordering::Relaxed)
}
pub fn mcp_tool_errors_total() -> u64 {
MCP_TOOL_ERRORS_TOTAL.load(Ordering::Relaxed)
}
pub struct FraiseQLMcpService<A: DatabaseAdapter> {
schema: Arc<CompiledSchema>,
executor: Arc<Executor<A>>,
tools: Vec<Tool>,
config: McpConfig,
oidc_validator: Option<Arc<OidcValidator>>,
}
impl<A: DatabaseAdapter> FraiseQLMcpService<A> {
#[must_use]
pub fn new(schema: Arc<CompiledSchema>, executor: Arc<Executor<A>>, config: McpConfig) -> Self {
let tools = super::tools::schema_to_tools(&schema, &config);
Self {
schema,
executor,
tools,
config,
oidc_validator: None,
}
}
#[must_use]
pub fn with_oidc_validator(mut self, validator: Option<Arc<OidcValidator>>) -> Self {
self.oidc_validator = validator;
self
}
async fn authenticate(
&self,
token: Option<String>,
request_id: String,
) -> Result<Option<SecurityContext>, CallToolResult> {
let Some(validator) = self.oidc_validator.as_ref() else {
return Ok(None); };
let Some(token) = token else {
return Ok(None); };
match validator.validate_token(&token).await {
Ok(user) => Ok(Some(SecurityContext::from_user(&user, request_id))),
Err(e) => {
tracing::warn!(error = %e, "MCP token validation failed");
Err(error_result("Invalid or expired authentication token"))
},
}
}
}
impl<A: DatabaseAdapter + Clone + Send + Sync + 'static> ServerHandler for FraiseQLMcpService<A> {
fn get_info(&self) -> ServerInfo {
ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
.with_instructions("FraiseQL GraphQL database — query and mutate via MCP tools")
}
fn list_tools(
&self,
_request: Option<rmcp::model::PaginatedRequestParams>,
_context: RequestContext<rmcp::RoleServer>,
) -> impl std::future::Future<Output = Result<ListToolsResult, rmcp::ErrorData>> + Send + '_
{
let result = ListToolsResult {
tools: self.tools.clone(),
next_cursor: None,
meta: None,
};
std::future::ready(Ok(result))
}
fn call_tool(
&self,
request: CallToolRequestParams,
context: RequestContext<rmcp::RoleServer>,
) -> impl std::future::Future<Output = Result<CallToolResult, rmcp::ErrorData>> + Send + '_
{
let tool_name = request.name.to_string();
let arguments = request.arguments;
let require_auth = self.config.require_auth;
let request_id = context.id.to_string();
let token = context
.extensions
.get::<http::request::Parts>()
.and_then(|parts| extract_bearer(&parts.headers));
async move {
MCP_TOOL_CALLS_TOTAL.fetch_add(1, Ordering::Relaxed);
tracing::info!(tool = %tool_name, "MCP tool call");
let security_context = match self.authenticate(token, request_id).await {
Ok(ctx) => ctx,
Err(err_result) => {
MCP_TOOL_ERRORS_TOTAL.fetch_add(1, Ordering::Relaxed);
return Ok(err_result);
},
};
let result = super::executor::call_tool(
&tool_name,
arguments.as_ref(),
&self.schema,
&self.executor,
security_context.as_ref(),
require_auth,
)
.await;
if result.is_error == Some(true) {
MCP_TOOL_ERRORS_TOTAL.fetch_add(1, Ordering::Relaxed);
}
Ok(result)
}
}
fn get_tool(&self, name: &str) -> Option<Tool> {
self.tools.iter().find(|t| t.name == name).cloned()
}
}