use std::sync::{
Arc,
atomic::{AtomicU64, Ordering},
};
use fraiseql_core::{
schema::CompiledSchema,
security::{AuthMiddleware, AuthRequest, AuthenticatedUser, OidcValidator, SecurityContext},
};
use rmcp::{
ServerHandler,
model::{
CallToolRequestParams, CallToolResponse, CallToolResult, GetPromptRequestParams,
GetPromptResponse, GetPromptResult, ListPromptsResult, ListResourceTemplatesResult,
ListResourcesResult, ListToolsResult, PaginatedRequestParams, ReadResourceRequestParams,
ReadResourceResponse, ReadResourceResult, ResourceContents, ServerCapabilities,
ServerConfig, Tool,
},
service::RequestContext,
};
use super::{McpConfig, executor::error_result};
use crate::routes::graphql::AppState;
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 enum McpTokenValidator {
Oidc(Arc<OidcValidator>),
Hs256(Arc<AuthMiddleware>),
}
impl std::fmt::Debug for McpTokenValidator {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Oidc(_) => f.write_str("McpTokenValidator::Oidc"),
Self::Hs256(_) => f.write_str("McpTokenValidator::Hs256"),
}
}
}
impl McpTokenValidator {
async fn validate(&self, token: &str) -> Result<AuthenticatedUser, String> {
match self {
Self::Oidc(v) => v.validate_token(token).await.map_err(|e| e.to_string()),
Self::Hs256(m) => {
let req = AuthRequest::new(Some(format!("Bearer {token}")));
m.validate_request(&req).map_err(|e| e.to_string())
},
}
}
}
pub struct FraiseQLMcpService {
state: AppState,
schema: Arc<CompiledSchema>,
tools: Vec<Tool>,
config: McpConfig,
validator: Option<McpTokenValidator>,
session_state: Option<Arc<fraiseql_auth::session_state::SessionState>>,
}
impl FraiseQLMcpService {
#[must_use]
pub fn new(state: AppState, config: McpConfig) -> Self {
let schema = Arc::new(state.executor().schema().clone());
let tools = super::tools::schema_to_tools(&schema, &config);
Self {
state,
schema,
tools,
config,
validator: None,
session_state: None,
}
}
#[must_use]
pub fn with_session_state(
mut self,
store: Option<Arc<fraiseql_auth::session_state::SessionState>>,
) -> Self {
self.session_state = store;
self
}
#[must_use]
pub fn with_token_validator(mut self, validator: Option<McpTokenValidator>) -> Self {
self.validator = validator;
self
}
#[allow(clippy::result_large_err)]
async fn authenticate(
&self,
token: Option<String>,
request_id: String,
) -> Result<Option<SecurityContext>, CallToolResult> {
let Some(validator) = self.validator.as_ref() else {
return Ok(None); };
let Some(token) = token else {
return Ok(None); };
match validator.validate(&token).await {
Ok(user) => {
let mut ctx = crate::extractors::build_security_context(
&user,
request_id,
Some(self.schema.tenant_claim()),
)
.with_transport("mcp");
match crate::identity::resolve_request_identity(
self.state.identity_resolver.as_deref(),
Some(&mut ctx),
)
.await
{
crate::identity::EnrichmentOutcome::Proceed => Ok(Some(ctx)),
crate::identity::EnrichmentOutcome::Denied => {
Err(error_result(crate::identity::EnrichmentOutcome::DENIED_MESSAGE))
},
crate::identity::EnrichmentOutcome::Unavailable => {
Err(error_result(crate::identity::EnrichmentOutcome::UNAVAILABLE_MESSAGE))
},
}
},
Err(e) => {
tracing::warn!(error = %e, "MCP token validation failed");
Err(error_result("Invalid or expired authentication token"))
},
}
}
#[doc(hidden)] pub async fn call_tool_authenticated(
&self,
tool_name: &str,
arguments: Option<&serde_json::Map<String, serde_json::Value>>,
token: Option<String>,
request_id: String,
headers: &axum::http::HeaderMap,
) -> CallToolResult {
use crate::routes::graphql::tenant_dispatch;
let security_context = match self.authenticate(token, request_id).await {
Ok(ctx) => ctx,
Err(err_result) => return err_result,
};
let security_context = security_context.as_ref();
let sanitizer = &self.state.error_sanitizer;
let tenant_key =
match tenant_dispatch::resolve_tenant_key(&self.state, security_context, headers) {
Ok(key) => key,
Err(e) => return error_result(&sanitize(sanitizer, &e)),
};
let dispatch = match tenant_dispatch::dispatch_to_tenant(&self.state, tenant_key.as_deref())
{
Ok(d) => d,
Err(e) => return error_result(&sanitize(sanitizer, &e)),
};
let thread = self
.config
.session_state
.then_some(self.session_state.as_ref())
.flatten()
.and_then(|store| {
super::session::thread_key(security_context, headers).map(|key| (store, key))
});
let prior = match thread {
Some((store, ref key)) => super::session::read_context(store, key).await,
None => Vec::new(),
};
let mut result = super::executor::call_tool(
tool_name,
arguments,
&super::executor::McpCallContext {
schema: &self.schema,
executor: &dispatch.executor,
config: &self.config,
security_context,
error_sanitizer: sanitizer,
},
)
.await;
if let Some((store, key)) = thread {
if result.is_error != Some(true) {
super::session::record_call(store, &key, tool_name, arguments, prior.clone()).await;
let mut calls = prior;
calls.push(serde_json::json!({ "tool": tool_name }));
super::session::attach_context(&mut result, &key, &calls);
}
}
result
}
}
impl FraiseQLMcpService {
#[doc(hidden)] pub async fn read_resource_authenticated(
&self,
uri: &str,
token: Option<String>,
request_id: String,
headers: &axum::http::HeaderMap,
) -> Result<ReadResourceResult, rmcp::ErrorData> {
let Some(name) = super::resources::query_name_from_uri(uri) else {
return Err(rmcp::ErrorData::resource_not_found(
format!("Unknown resource URI: {uri}"),
None,
));
};
let result = self.call_tool_authenticated(name, None, token, request_id, headers).await;
if result.is_error == Some(true) {
let detail = first_text(&result).unwrap_or_else(|| "resource read refused".to_string());
return Err(rmcp::ErrorData::invalid_request(detail, None));
}
Ok(ReadResourceResult::new(vec![ResourceContents::TextResourceContents {
uri: uri.to_string(),
mime_type: Some("application/json".to_string()),
text: first_text(&result).unwrap_or_default(),
meta: None,
}]))
}
}
fn first_text(result: &CallToolResult) -> Option<String> {
result.content.first()?.as_text().map(|t| t.text.clone())
}
pub(crate) fn sanitize(
sanitizer: &crate::config::error_sanitization::ErrorSanitizer,
error: &fraiseql_error::FraiseQLError,
) -> String {
sanitizer
.sanitize(crate::error::GraphQLError::from_fraiseql_error(error))
.message
}
impl ServerHandler for FraiseQLMcpService {
fn get_info(&self) -> ServerConfig {
ServerConfig::new(
ServerCapabilities::builder()
.enable_tools()
.enable_resources()
.enable_prompts()
.build(),
)
.with_instructions(
"FraiseQL GraphQL database — query and mutate via MCP tools. Each readable query is \
also a Resource at fraiseql://query/{name}; reading one runs the same operation \
under the same authentication and tenant scoping as calling its tool.",
)
}
fn list_tools(
&self,
_request: Option<rmcp::model::PaginatedRequestParams>,
_context: RequestContext<rmcp::RoleServer>,
) -> impl std::future::Future<Output = Result<ListToolsResult, rmcp::ErrorData>> + Send + '_
{
std::future::ready(Ok(ListToolsResult::with_all_items(self.tools.clone())))
}
fn call_tool(
&self,
request: CallToolRequestParams,
context: RequestContext<rmcp::RoleServer>,
) -> impl std::future::Future<Output = Result<CallToolResponse, rmcp::ErrorData>> + Send + '_
{
let tool_name = request.name.to_string();
let arguments = request.arguments;
let request_id = context.id.to_string();
let parts = context.extensions.get::<http::request::Parts>();
let token = parts.and_then(|parts| extract_bearer(&parts.headers));
let headers = parts.map(|parts| parts.headers.clone()).unwrap_or_default();
async move {
MCP_TOOL_CALLS_TOTAL.fetch_add(1, Ordering::Relaxed);
tracing::info!(tool = %tool_name, "MCP tool call");
let result = self
.call_tool_authenticated(
&tool_name,
arguments.as_ref(),
token,
request_id,
&headers,
)
.await;
if result.is_error == Some(true) {
MCP_TOOL_ERRORS_TOTAL.fetch_add(1, Ordering::Relaxed);
}
Ok(result.into())
}
}
fn list_resources(
&self,
_request: Option<PaginatedRequestParams>,
_context: RequestContext<rmcp::RoleServer>,
) -> impl std::future::Future<Output = Result<ListResourcesResult, rmcp::ErrorData>> + Send + '_
{
let result = ListResourcesResult::with_all_items(super::resources::schema_to_resources(
&self.schema,
&self.config,
));
std::future::ready(Ok(result))
}
fn list_resource_templates(
&self,
_request: Option<PaginatedRequestParams>,
_context: RequestContext<rmcp::RoleServer>,
) -> impl std::future::Future<Output = Result<ListResourceTemplatesResult, rmcp::ErrorData>>
+ Send
+ '_ {
let result = ListResourceTemplatesResult::with_all_items(
super::resources::schema_to_resource_templates(&self.schema, &self.config),
);
std::future::ready(Ok(result))
}
fn read_resource(
&self,
request: ReadResourceRequestParams,
context: RequestContext<rmcp::RoleServer>,
) -> impl std::future::Future<Output = Result<ReadResourceResponse, rmcp::ErrorData>> + Send + '_
{
let request_id = context.id.to_string();
let parts = context.extensions.get::<http::request::Parts>();
let token = parts.and_then(|parts| extract_bearer(&parts.headers));
let headers = parts.map(|parts| parts.headers.clone()).unwrap_or_default();
let uri = request.uri;
async move {
self.read_resource_authenticated(&uri, token, request_id, &headers)
.await
.map(Into::into)
}
}
fn list_prompts(
&self,
_request: Option<PaginatedRequestParams>,
_context: RequestContext<rmcp::RoleServer>,
) -> impl std::future::Future<Output = Result<ListPromptsResult, rmcp::ErrorData>> + Send + '_
{
let result = ListPromptsResult::with_all_items(super::resources::schema_to_prompts(
&self.schema,
&self.config,
));
std::future::ready(Ok(result))
}
fn get_prompt(
&self,
request: GetPromptRequestParams,
_context: RequestContext<rmcp::RoleServer>,
) -> impl std::future::Future<Output = Result<GetPromptResponse, rmcp::ErrorData>> + Send + '_
{
let rendered = super::resources::render_prompt(
&request.name,
request.arguments.as_ref(),
&self.schema,
&self.config,
);
std::future::ready(match rendered {
Some((description, messages)) => {
Ok(GetPromptResult::new(messages).with_description(description).into())
},
None => Err(rmcp::ErrorData::invalid_params(
format!("Unknown prompt: {}", request.name),
None,
)),
})
}
fn get_tool(&self, name: &str) -> Option<Tool> {
self.tools.iter().find(|t| t.name == name).cloned()
}
}