use std::sync::Arc;
use rmcp::handler::server::router::tool::ToolRouter;
use rmcp::handler::server::wrapper::Parameters;
use rmcp::model::{
CallToolResult, ContentBlock, Implementation, ProtocolVersion, ServerCapabilities, ServerInfo,
};
use rmcp::service::RequestContext;
use rmcp::transport::stdio;
use rmcp::{
ErrorData as McpError, RoleServer, ServerHandler, ServiceExt, schemars, tool, tool_handler,
tool_router,
};
use serde::Deserialize;
use tokio::sync::Mutex;
use crate::auth::auth_manager::{AuthManager, header_location_for};
use crate::core::config_schema::Config;
use crate::core::errors::McpifyError;
use crate::data::store::{cached_store_connection, get_endpoint};
use crate::http::auth_extractor::extract_request_credentials;
use crate::tools::call_tool::call_operation;
use crate::tools::get_tool::get_operation;
use crate::tools::search_tool::search_operations;
fn default_search_limit() -> usize {
5
}
#[derive(Debug, Deserialize, schemars::JsonSchema)]
pub struct SearchArgs {
pub query: String,
#[serde(default = "default_search_limit")]
pub limit: usize,
}
#[derive(Debug, Deserialize, schemars::JsonSchema)]
pub struct GetArgs {
pub operation_id: String,
}
fn default_call_arguments() -> serde_json::Value {
serde_json::json!({})
}
#[derive(Debug, Deserialize, schemars::JsonSchema)]
pub struct CallArgs {
pub operation_id: String,
#[serde(default = "default_call_arguments")]
pub arguments: serde_json::Value,
}
#[derive(Clone)]
pub struct McpifyServer {
api_version: String,
config: Config,
auth_manager: Arc<Mutex<AuthManager>>,
tool_router: ToolRouter<McpifyServer>,
}
#[tool_router]
impl McpifyServer {
pub fn new(api_version: String, config: Config, auth_manager: Arc<Mutex<AuthManager>>) -> Self {
Self {
api_version,
config,
auth_manager,
tool_router: Self::tool_router(),
}
}
#[tool(
description = "Semantic search for GitHub v3 REST API operations using a natural-language query."
)]
async fn search(
&self,
Parameters(args): Parameters<SearchArgs>,
) -> Result<CallToolResult, McpError> {
let api_version = self.api_version.clone();
self.run_tool("search", async move {
let conn = cached_store_connection(&api_version)?.lock().unwrap();
search_operations(&conn, &args.query, args.limit)
})
.await
}
#[tool(
description = "Return the schema, path, method, and documentation for a specific GitHub v3 REST API operationId."
)]
async fn get(&self, Parameters(args): Parameters<GetArgs>) -> Result<CallToolResult, McpError> {
let api_version = self.api_version.clone();
self.run_tool("get", async move {
let conn = cached_store_connection(&api_version)?.lock().unwrap();
get_operation(&conn, &args.operation_id)
})
.await
}
#[tool(
description = "Validate arguments, invoke a live GitHub v3 REST API API operation, and validate the response."
)]
async fn call(
&self,
Parameters(args): Parameters<CallArgs>,
context: RequestContext<RoleServer>,
) -> Result<CallToolResult, McpError> {
let api_version = self.api_version.clone();
let config = self.config.clone();
let auth_manager = self.auth_manager.clone();
let request_credentials = context
.extensions
.get::<axum::http::request::Parts>()
.and_then(|parts| {
let (header_location, header_name) = header_location_for(config.auth_method);
extract_request_credentials(&parts.headers, header_location, header_name).ok()
});
self.run_tool("call", async move {
let endpoint = {
let conn = cached_store_connection(&api_version)?.lock().unwrap();
get_endpoint(&conn, &args.operation_id)?.ok_or_else(|| {
McpifyError::NotFound(format!("unknown operationId '{}'", args.operation_id))
})?
};
let mut auth_manager = auth_manager.lock().await;
call_operation(
&endpoint,
&config,
&mut auth_manager,
&args.operation_id,
args.arguments,
request_credentials.as_ref(),
)
.await
})
.await
}
}
impl McpifyServer {
async fn run_tool<F>(&self, tool_name: &str, fut: F) -> Result<CallToolResult, McpError>
where
F: std::future::Future<Output = anyhow::Result<serde_json::Value>>,
{
match fut.await {
Ok(value) => {
let text =
serde_json::to_string_pretty(&value).unwrap_or_else(|_| value.to_string());
Ok(CallToolResult::success(vec![ContentBlock::text(text)]))
}
Err(err) => {
tracing::error!(tool = tool_name, error = %err, "tool execution failed");
Ok(CallToolResult::error(vec![ContentBlock::text(
err.to_string(),
)]))
}
}
}
}
#[tool_handler(router = self.tool_router.clone())]
impl ServerHandler for McpifyServer {
fn get_info(&self) -> ServerInfo {
ServerInfo::new(ServerCapabilities::builder().enable_tools().build())
.with_server_info(Implementation::from_build_env())
.with_protocol_version(ProtocolVersion::V_2024_11_05)
.with_instructions(
"Exposes exactly 3 tools -- search, get, call -- backed by an embedded \
semantic database, so you never need the full API surface in context."
.to_string(),
)
}
}
pub async fn connect_stdio<S>(server: S) -> anyhow::Result<()>
where
S: rmcp::ServerHandler,
{
let running = server.serve(stdio()).await?;
tracing::info!("MCP server connected over stdio");
running.waiting().await?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::config_schema::AuthMethod;
use crate::data::store::list_endpoints;
use rmcp::model::CallToolRequestParams;
#[derive(Debug, Clone, Default)]
struct TestClient;
impl rmcp::ClientHandler for TestClient {}
fn server() -> McpifyServer {
let config: Config = serde_json::from_value(serde_json::json!({
"url": "https://api.example.test",
"auth_method": "pat"
}))
.unwrap();
McpifyServer::new(
"gh-2026-03-10".to_string(),
config,
Arc::new(Mutex::new(AuthManager::new(AuthMethod::Pat))),
)
}
#[test]
fn argument_defaults_match_the_public_tool_contract() {
assert_eq!(default_search_limit(), 5);
assert_eq!(default_call_arguments(), serde_json::json!({}));
let search: SearchArgs = serde_json::from_value(serde_json::json!({
"query": "find an operation"
}))
.unwrap();
assert_eq!(search.limit, 5);
let call: CallArgs = serde_json::from_value(serde_json::json!({
"operation_id": "an-operation"
}))
.unwrap();
assert_eq!(call.arguments, serde_json::json!({}));
}
#[tokio::test]
async fn search_and_get_return_mcp_content_envelopes() {
let server = server();
let search = server
.search(Parameters(SearchArgs {
query: "find an operation".to_string(),
limit: 2,
}))
.await
.unwrap();
assert_eq!(search.is_error, Some(false));
let operation_id = {
let conn = cached_store_connection("gh-2026-03-10").unwrap();
let conn = conn.lock().unwrap();
list_endpoints(&conn).unwrap()[0].operation_id.clone()
};
let get = server
.get(Parameters(GetArgs { operation_id }))
.await
.unwrap();
assert_eq!(get.is_error, Some(false));
let missing = server
.get(Parameters(GetArgs {
operation_id: "definitely-unknown-operation".to_string(),
}))
.await
.unwrap();
assert_eq!(missing.is_error, Some(true));
}
#[tokio::test]
async fn run_tool_formats_successes_and_failures_consistently() {
let server = server();
let success = server
.run_tool("coverage", async { Ok(serde_json::json!({ "ok": true })) })
.await
.unwrap();
assert_eq!(success.is_error, Some(false));
let failure = server
.run_tool("coverage", async { anyhow::bail!("coverage failure") })
.await
.unwrap();
assert_eq!(failure.is_error, Some(true));
}
#[test]
fn server_info_advertises_the_generated_tool_surface() {
let info = server().get_info();
assert_eq!(info.protocol_version, ProtocolVersion::V_2024_11_05);
assert!(info.capabilities.tools.is_some());
assert!(info.instructions.unwrap().contains("search, get, call"));
}
#[tokio::test]
async fn mcp_protocol_routes_search_get_and_call_requests() {
let (server_transport, client_transport) = tokio::io::duplex(64 * 1024);
let server_task = tokio::spawn(async move {
server().serve(server_transport).await?.waiting().await?;
anyhow::Ok(())
});
let client = TestClient.serve(client_transport).await.unwrap();
let tools = client.list_all_tools().await.unwrap();
assert_eq!(
tools
.iter()
.map(|tool| tool.name.as_ref())
.collect::<Vec<_>>(),
["call", "get", "search"]
);
let search = client
.call_tool(
CallToolRequestParams::new("search").with_arguments(
serde_json::json!({ "query": "find an operation", "limit": 1 })
.as_object()
.unwrap()
.clone(),
),
)
.await
.unwrap();
assert_eq!(search.is_error, Some(false));
let operation_id = {
let conn = cached_store_connection("gh-2026-03-10").unwrap();
let conn = conn.lock().unwrap();
list_endpoints(&conn).unwrap()[0].operation_id.clone()
};
let get = client
.call_tool(
CallToolRequestParams::new("get").with_arguments(
serde_json::json!({ "operation_id": operation_id })
.as_object()
.unwrap()
.clone(),
),
)
.await
.unwrap();
assert_eq!(get.is_error, Some(false));
let call = client
.call_tool(
CallToolRequestParams::new("call").with_arguments(
serde_json::json!({ "operation_id": "definitely-unknown", "arguments": {} })
.as_object()
.unwrap()
.clone(),
),
)
.await
.unwrap();
assert_eq!(call.is_error, Some(true));
drop(client);
tokio::time::timeout(std::time::Duration::from_secs(2), server_task)
.await
.unwrap()
.unwrap()
.unwrap();
}
}