Update: 将子项目从 submodule 转为完整内容
- 移除 GovAI, nomifun-tauri, 算力盒子 的 submodule 引用 - 添加所有子项目的完整源代码 - 保留原始 .git 为 .git.bak 备份
This commit is contained in:
@@ -0,0 +1,30 @@
|
||||
[package]
|
||||
name = "nomifun-mcp"
|
||||
version.workspace = true
|
||||
edition.workspace = true
|
||||
|
||||
[dependencies]
|
||||
nomifun-runtime.workspace = true
|
||||
nomifun-common.workspace = true
|
||||
nomifun-db.workspace = true
|
||||
nomifun-api-types.workspace = true
|
||||
nomifun-net.workspace = true
|
||||
async-trait.workspace = true
|
||||
dashmap.workspace = true
|
||||
axum.workspace = true
|
||||
dirs.workspace = true
|
||||
thiserror.workspace = true
|
||||
oauth2.workspace = true
|
||||
open.workspace = true
|
||||
toml.workspace = true
|
||||
tokio.workspace = true
|
||||
tracing.workspace = true
|
||||
reqwest.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
axum.workspace = true
|
||||
serde_json.workspace = true
|
||||
tempfile.workspace = true
|
||||
tokio = { workspace = true, features = ["macros", "rt"] }
|
||||
@@ -0,0 +1,246 @@
|
||||
use nomifun_common::McpSource;
|
||||
|
||||
use crate::error::McpError;
|
||||
use crate::types::McpServerTransport;
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// DetectedServer — lightweight server info from Agent CLI detection
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// A server configuration detected from an Agent CLI.
|
||||
///
|
||||
/// Returned by `McpAgentAdapter::detect_existing()`. Contains only the
|
||||
/// fields needed for diff comparison during sync operations (name +
|
||||
/// transport). The full `McpServer` model includes DB-level metadata
|
||||
/// (id, timestamps, etc.) that CLI detection cannot provide.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct DetectedServer {
|
||||
/// Server name as registered in the Agent CLI.
|
||||
pub name: String,
|
||||
/// Transport configuration detected from the Agent CLI.
|
||||
pub transport: McpServerTransport,
|
||||
/// Whether this detected MCP can be imported without extra intervention.
|
||||
pub importable: bool,
|
||||
/// Human-readable reason when the MCP is not currently importable.
|
||||
pub import_skip_reason: Option<String>,
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// McpAgentAdapter — trait for Agent CLI adapters
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Abstraction for AI Agent CLI MCP configuration management.
|
||||
///
|
||||
/// Each Agent CLI (Claude, Gemini, Qwen, etc.) implements this trait to
|
||||
/// provide detection, installation, and removal of MCP server configurations.
|
||||
///
|
||||
/// # Concurrency
|
||||
///
|
||||
/// Implementations do **not** need to handle concurrency internally.
|
||||
/// The sync service applies per-agent serialization locks before calling
|
||||
/// adapter methods.
|
||||
///
|
||||
/// # Error handling
|
||||
///
|
||||
/// Methods return `McpError` rather than `AppError` to keep the adapter
|
||||
/// layer independent of HTTP concerns.
|
||||
#[async_trait::async_trait]
|
||||
pub trait McpAgentAdapter: Send + Sync {
|
||||
/// Returns the agent source identifier (e.g., `McpSource::Claude`).
|
||||
fn source(&self) -> McpSource;
|
||||
|
||||
/// Checks whether the Agent CLI is installed on this machine.
|
||||
///
|
||||
/// Typically implemented via `which <cli-name>` or checking a known
|
||||
/// config directory.
|
||||
async fn is_installed(&self) -> Result<bool, McpError>;
|
||||
|
||||
/// Reads the currently configured MCP servers from this Agent CLI.
|
||||
///
|
||||
/// Returns an empty vec if the CLI is installed but has no MCP servers.
|
||||
/// Returns `McpError::AgentNotInstalled` if the CLI is not available.
|
||||
async fn detect_existing(&self) -> Result<Vec<DetectedServer>, McpError>;
|
||||
|
||||
/// Installs (or updates) an MCP server configuration in this Agent CLI.
|
||||
///
|
||||
/// The `name` and `transport` fields from `server` are used to configure
|
||||
/// the Agent CLI. If a server with the same name already exists in the
|
||||
/// CLI, it should be replaced.
|
||||
async fn install_server(&self, name: &str, transport: &McpServerTransport) -> Result<(), McpError>;
|
||||
|
||||
/// Removes an MCP server configuration from this Agent CLI by name.
|
||||
///
|
||||
/// Should be idempotent: removing a non-existent server is not an error.
|
||||
async fn remove_server(&self, name: &str) -> Result<(), McpError>;
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::collections::HashMap;
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
/// In-memory mock adapter for testing the trait interface.
|
||||
struct MockAdapter {
|
||||
source: McpSource,
|
||||
installed: bool,
|
||||
servers: Arc<Mutex<Vec<DetectedServer>>>,
|
||||
}
|
||||
|
||||
impl MockAdapter {
|
||||
fn new(source: McpSource, installed: bool) -> Self {
|
||||
Self {
|
||||
source,
|
||||
installed,
|
||||
servers: Arc::new(Mutex::new(Vec::new())),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl McpAgentAdapter for MockAdapter {
|
||||
fn source(&self) -> McpSource {
|
||||
self.source
|
||||
}
|
||||
|
||||
async fn is_installed(&self) -> Result<bool, McpError> {
|
||||
Ok(self.installed)
|
||||
}
|
||||
|
||||
async fn detect_existing(&self) -> Result<Vec<DetectedServer>, McpError> {
|
||||
if !self.installed {
|
||||
return Err(McpError::AgentNotInstalled(format!("{:?}", self.source)));
|
||||
}
|
||||
let servers = self.servers.lock().unwrap();
|
||||
Ok(servers.clone())
|
||||
}
|
||||
|
||||
async fn install_server(&self, name: &str, transport: &McpServerTransport) -> Result<(), McpError> {
|
||||
if !self.installed {
|
||||
return Err(McpError::AgentNotInstalled(format!("{:?}", self.source)));
|
||||
}
|
||||
let mut servers = self.servers.lock().unwrap();
|
||||
servers.retain(|s| s.name != name);
|
||||
servers.push(DetectedServer {
|
||||
name: name.to_owned(),
|
||||
transport: transport.clone(),
|
||||
importable: true,
|
||||
import_skip_reason: None,
|
||||
});
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn remove_server(&self, name: &str) -> Result<(), McpError> {
|
||||
if !self.installed {
|
||||
return Err(McpError::AgentNotInstalled(format!("{:?}", self.source)));
|
||||
}
|
||||
let mut servers = self.servers.lock().unwrap();
|
||||
servers.retain(|s| s.name != name);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mock_adapter_source() {
|
||||
let adapter = MockAdapter::new(McpSource::Claude, true);
|
||||
assert_eq!(adapter.source(), McpSource::Claude);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mock_adapter_is_installed() {
|
||||
let installed = MockAdapter::new(McpSource::Gemini, true);
|
||||
assert!(installed.is_installed().await.unwrap());
|
||||
|
||||
let not_installed = MockAdapter::new(McpSource::Gemini, false);
|
||||
assert!(!not_installed.is_installed().await.unwrap());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mock_adapter_install_and_detect() {
|
||||
let adapter = MockAdapter::new(McpSource::Claude, true);
|
||||
let transport = McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec!["-y".into(), "@test/server".into()],
|
||||
env: HashMap::new(),
|
||||
};
|
||||
|
||||
adapter.install_server("test-mcp", &transport).await.unwrap();
|
||||
|
||||
let detected = adapter.detect_existing().await.unwrap();
|
||||
assert_eq!(detected.len(), 1);
|
||||
assert_eq!(detected[0].name, "test-mcp");
|
||||
assert_eq!(detected[0].transport, transport);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mock_adapter_install_replaces_existing() {
|
||||
let adapter = MockAdapter::new(McpSource::Claude, true);
|
||||
let t1 = McpServerTransport::Stdio {
|
||||
command: "old".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
};
|
||||
let t2 = McpServerTransport::Stdio {
|
||||
command: "new".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
};
|
||||
|
||||
adapter.install_server("test-mcp", &t1).await.unwrap();
|
||||
adapter.install_server("test-mcp", &t2).await.unwrap();
|
||||
|
||||
let detected = adapter.detect_existing().await.unwrap();
|
||||
assert_eq!(detected.len(), 1);
|
||||
match &detected[0].transport {
|
||||
McpServerTransport::Stdio { command, .. } => assert_eq!(command, "new"),
|
||||
_ => panic!("expected Stdio"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mock_adapter_remove() {
|
||||
let adapter = MockAdapter::new(McpSource::Claude, true);
|
||||
let transport = McpServerTransport::Http {
|
||||
url: "http://x".into(),
|
||||
headers: HashMap::new(),
|
||||
};
|
||||
adapter.install_server("srv", &transport).await.unwrap();
|
||||
adapter.remove_server("srv").await.unwrap();
|
||||
|
||||
let detected = adapter.detect_existing().await.unwrap();
|
||||
assert!(detected.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mock_adapter_remove_nonexistent_is_idempotent() {
|
||||
let adapter = MockAdapter::new(McpSource::Claude, true);
|
||||
// Should not error
|
||||
adapter.remove_server("nonexistent").await.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn not_installed_detect_fails() {
|
||||
let adapter = MockAdapter::new(McpSource::Qwen, false);
|
||||
let result = adapter.detect_existing().await;
|
||||
assert!(matches!(result, Err(McpError::AgentNotInstalled(_))));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn not_installed_install_fails() {
|
||||
let adapter = MockAdapter::new(McpSource::Qwen, false);
|
||||
let transport = McpServerTransport::Stdio {
|
||||
command: "x".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
};
|
||||
let result = adapter.install_server("srv", &transport).await;
|
||||
assert!(matches!(result, Err(McpError::AgentNotInstalled(_))));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn trait_is_object_safe() {
|
||||
let adapter: Arc<dyn McpAgentAdapter> = Arc::new(MockAdapter::new(McpSource::Nomifun, true));
|
||||
assert_eq!(adapter.source(), McpSource::Nomifun);
|
||||
assert!(adapter.is_installed().await.unwrap());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,348 @@
|
||||
use std::collections::HashMap;
|
||||
|
||||
use nomifun_common::McpSource;
|
||||
|
||||
use crate::adapter::{DetectedServer, McpAgentAdapter};
|
||||
use crate::error::McpError;
|
||||
use crate::types::McpServerTransport;
|
||||
|
||||
use super::cli_helpers::{
|
||||
DETECT_TIMEOUT, MUTATE_TIMEOUT, is_cli_installed, normalize_detection_status, run_cli, strip_ansi,
|
||||
};
|
||||
|
||||
const CLI_NAME: &str = "claude";
|
||||
|
||||
/// Scopes to try when removing a server (user → local → project).
|
||||
const REMOVE_SCOPES: &[&str] = &["user", "local", "project"];
|
||||
|
||||
/// MCP Agent adapter for Claude CLI.
|
||||
///
|
||||
/// # CLI Commands
|
||||
///
|
||||
/// - **detect**: `claude mcp list`
|
||||
/// - **install (stdio)**: `claude mcp add-json -s user <name> <json>`
|
||||
/// - **install (http/sse)**: `claude mcp add -s user --transport <type> <name> <url> [--header ...]`
|
||||
/// - **remove**: `claude mcp remove -s <scope> <name>` (tries user → local → project)
|
||||
///
|
||||
/// Claude's list output uses a custom format:
|
||||
/// `name: command args - ✓ Connected` or `name: command args - ✗ Failed`
|
||||
pub struct ClaudeAdapter;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl McpAgentAdapter for ClaudeAdapter {
|
||||
fn source(&self) -> McpSource {
|
||||
McpSource::Claude
|
||||
}
|
||||
|
||||
async fn is_installed(&self) -> Result<bool, McpError> {
|
||||
is_cli_installed(CLI_NAME).await
|
||||
}
|
||||
|
||||
async fn detect_existing(&self) -> Result<Vec<DetectedServer>, McpError> {
|
||||
if !self.is_installed().await? {
|
||||
return Err(McpError::AgentNotInstalled(CLI_NAME.into()));
|
||||
}
|
||||
|
||||
let (stdout, _stderr) = run_cli(CLI_NAME, &["mcp", "list"], DETECT_TIMEOUT).await?;
|
||||
Ok(parse_claude_list_output(&stdout))
|
||||
}
|
||||
|
||||
async fn install_server(&self, name: &str, transport: &McpServerTransport) -> Result<(), McpError> {
|
||||
if !self.is_installed().await? {
|
||||
return Err(McpError::AgentNotInstalled(CLI_NAME.into()));
|
||||
}
|
||||
|
||||
match transport {
|
||||
McpServerTransport::Stdio { command, args, env } => {
|
||||
let config = build_stdio_json(command, args, env);
|
||||
let config_str =
|
||||
serde_json::to_string(&config).map_err(|e| McpError::AgentOperationFailed(e.to_string()))?;
|
||||
run_cli(
|
||||
CLI_NAME,
|
||||
&["mcp", "add-json", "-s", "user", name, &config_str],
|
||||
MUTATE_TIMEOUT,
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
McpServerTransport::Sse { url, headers } => {
|
||||
install_http_like(name, "sse", url, headers).await?;
|
||||
}
|
||||
McpServerTransport::Http { url, headers } => {
|
||||
install_http_like(name, "http", url, headers).await?;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn remove_server(&self, name: &str) -> Result<(), McpError> {
|
||||
if !self.is_installed().await? {
|
||||
return Err(McpError::AgentNotInstalled(CLI_NAME.into()));
|
||||
}
|
||||
|
||||
// Try each scope; stop on first success or "not found".
|
||||
for scope in REMOVE_SCOPES {
|
||||
let (stdout, _stderr) = run_cli(CLI_NAME, &["mcp", "remove", "-s", scope, name], MUTATE_TIMEOUT).await?;
|
||||
let lower = stdout.to_lowercase();
|
||||
if lower.contains("removed") || lower.contains("not found") {
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
|
||||
// If none of the scopes reported "removed" or "not found", treat as
|
||||
// idempotent success (server may simply not exist).
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// Install an HTTP-like (sse/http) server via `claude mcp add`.
|
||||
async fn install_http_like(
|
||||
name: &str,
|
||||
transport_type: &str,
|
||||
url: &str,
|
||||
headers: &HashMap<String, String>,
|
||||
) -> Result<(), McpError> {
|
||||
let mut args = vec![
|
||||
"mcp".to_owned(),
|
||||
"add".to_owned(),
|
||||
"-s".to_owned(),
|
||||
"user".to_owned(),
|
||||
"--transport".to_owned(),
|
||||
transport_type.to_owned(),
|
||||
name.to_owned(),
|
||||
url.to_owned(),
|
||||
];
|
||||
|
||||
for (key, value) in headers {
|
||||
args.push("--header".to_owned());
|
||||
args.push(format!("{key}: {value}"));
|
||||
}
|
||||
|
||||
let arg_refs: Vec<&str> = args.iter().map(|s| s.as_str()).collect();
|
||||
run_cli(CLI_NAME, &arg_refs, MUTATE_TIMEOUT).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Build the JSON config for `claude mcp add-json`.
|
||||
fn build_stdio_json(command: &str, args: &[String], env: &HashMap<String, String>) -> serde_json::Value {
|
||||
let mut config = serde_json::json!({
|
||||
"command": command,
|
||||
"args": args,
|
||||
});
|
||||
if !env.is_empty() {
|
||||
config["env"] = serde_json::json!(env);
|
||||
}
|
||||
config
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Output parsing
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Parse Claude CLI `mcp list` output.
|
||||
///
|
||||
/// Claude uses a custom format (not the standard Gemini/Qwen pattern):
|
||||
/// ```text
|
||||
/// name: command args - ✓ Connected
|
||||
/// name: command args - ✗ Failed to connect
|
||||
/// ```
|
||||
fn parse_claude_list_output(output: &str) -> Vec<DetectedServer> {
|
||||
let cleaned = strip_ansi(output);
|
||||
let mut servers = Vec::new();
|
||||
|
||||
for line in cleaned.lines() {
|
||||
let trimmed = line.trim();
|
||||
if let Some(server) = parse_claude_list_line(trimmed) {
|
||||
servers.push(server);
|
||||
}
|
||||
}
|
||||
|
||||
servers
|
||||
}
|
||||
|
||||
/// Parse a single line of Claude list output.
|
||||
///
|
||||
/// Pattern: `<name>: <command_or_url> - [✓|✗] <status>`
|
||||
fn parse_claude_list_line(line: &str) -> Option<DetectedServer> {
|
||||
// Split on " - " to separate "name: command" from status
|
||||
let dash_pos = line.rfind(" - ")?;
|
||||
let status = normalize_detection_status(&line[dash_pos + 3..]);
|
||||
|
||||
let name_cmd_part = &line[..dash_pos];
|
||||
|
||||
// Claude separates the name from command/URL with ": ". Names
|
||||
// themselves may contain ":" (for example plugin-scoped MCP entries).
|
||||
let separator_pos = name_cmd_part.find(": ")?;
|
||||
let name = name_cmd_part[..separator_pos].trim();
|
||||
if name.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let command_or_url = name_cmd_part[separator_pos + 2..].trim();
|
||||
if command_or_url.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let normalized_command_or_url = command_or_url
|
||||
.strip_suffix(" (HTTP)")
|
||||
.or_else(|| command_or_url.strip_suffix(" (SSE)"))
|
||||
.unwrap_or(command_or_url)
|
||||
.trim();
|
||||
|
||||
// Heuristic: if it looks like a URL, treat as HTTP; otherwise stdio.
|
||||
let transport =
|
||||
if normalized_command_or_url.starts_with("http://") || normalized_command_or_url.starts_with("https://") {
|
||||
// SSE heuristic: URL ending with /sse
|
||||
if normalized_command_or_url.ends_with("/sse") {
|
||||
McpServerTransport::Sse {
|
||||
url: normalized_command_or_url.to_owned(),
|
||||
headers: HashMap::new(),
|
||||
}
|
||||
} else {
|
||||
McpServerTransport::Http {
|
||||
url: normalized_command_or_url.to_owned(),
|
||||
headers: HashMap::new(),
|
||||
}
|
||||
}
|
||||
} else {
|
||||
McpServerTransport::Stdio {
|
||||
command: normalized_command_or_url.to_owned(),
|
||||
args: Vec::new(),
|
||||
env: HashMap::new(),
|
||||
}
|
||||
};
|
||||
|
||||
Some(DetectedServer {
|
||||
name: name.to_owned(),
|
||||
transport,
|
||||
importable: status.eq_ignore_ascii_case("Connected") && !name.starts_with("plugin:"),
|
||||
import_skip_reason: if name.starts_with("plugin:") {
|
||||
Some("Plugin-managed MCP".to_owned())
|
||||
} else if status.eq_ignore_ascii_case("Connected") {
|
||||
None
|
||||
} else {
|
||||
Some(status)
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn parse_claude_stdio_connected() {
|
||||
let output = "my-server: npx -y @test/server - ✓ Connected";
|
||||
let servers = parse_claude_list_output(output);
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert_eq!(servers[0].name, "my-server");
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Stdio { command, .. } => {
|
||||
assert_eq!(command, "npx -y @test/server");
|
||||
}
|
||||
_ => panic!("expected Stdio"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_claude_stdio_failed() {
|
||||
let output = "broken-srv: node index.js - ✗ Failed to connect";
|
||||
let servers = parse_claude_list_output(output);
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert!(!servers[0].importable);
|
||||
assert_eq!(servers[0].import_skip_reason.as_deref(), Some("Failed to connect"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_claude_http_server() {
|
||||
let output = "remote: https://example.com/mcp - ✓ Connected";
|
||||
let servers = parse_claude_list_output(output);
|
||||
assert_eq!(servers.len(), 1);
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Http { url, .. } => {
|
||||
assert_eq!(url, "https://example.com/mcp");
|
||||
}
|
||||
_ => panic!("expected Http"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_claude_sse_heuristic() {
|
||||
let output = "sse-srv: https://example.com/sse - ✓ Connected";
|
||||
let servers = parse_claude_list_output(output);
|
||||
assert_eq!(servers.len(), 1);
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Sse { url, .. } => {
|
||||
assert_eq!(url, "https://example.com/sse");
|
||||
}
|
||||
_ => panic!("expected Sse"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_claude_plugin_http_server_needing_auth() {
|
||||
let output = "plugin:slack:slack: https://mcp.slack.com/mcp (HTTP) - ! Needs authentication";
|
||||
let servers = parse_claude_list_output(output);
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert!(!servers[0].importable);
|
||||
assert_eq!(servers[0].import_skip_reason.as_deref(), Some("Plugin-managed MCP"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_claude_multiple_servers() {
|
||||
let output = "\
|
||||
my-mcp: npx -y @test/mcp - ✓ Connected
|
||||
broken: node bad.js - ✗ Failed to connect
|
||||
web: https://example.com/api - ✓ Connected";
|
||||
let servers = parse_claude_list_output(output);
|
||||
assert_eq!(servers.len(), 3);
|
||||
assert_eq!(servers[0].name, "my-mcp");
|
||||
assert_eq!(servers[1].name, "broken");
|
||||
assert!(!servers[1].importable);
|
||||
assert_eq!(servers[2].name, "web");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_claude_with_ansi() {
|
||||
let output = "\x1b[32m✓\x1b[0m test: npx srv - \x1b[32mConnected\x1b[0m";
|
||||
let servers = parse_claude_list_output(output);
|
||||
// After ANSI strip: "✓ test: npx srv - Connected"
|
||||
// The ✓ is at the beginning of the line, not in the "name: cmd" pattern
|
||||
// but it contains "Connected" so it should be parseable
|
||||
assert_eq!(servers.len(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_claude_no_servers() {
|
||||
let output = "No MCP servers configured.\nTry `claude mcp add` to get started.";
|
||||
let servers = parse_claude_list_output(output);
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_claude_empty_output() {
|
||||
let servers = parse_claude_list_output("");
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_stdio_json_without_env() {
|
||||
let json = build_stdio_json("npx", &["-y".into(), "srv".into()], &HashMap::new());
|
||||
assert_eq!(json["command"], "npx");
|
||||
assert_eq!(json["args"], serde_json::json!(["-y", "srv"]));
|
||||
assert!(json.get("env").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_stdio_json_with_env() {
|
||||
let mut env = HashMap::new();
|
||||
env.insert("KEY".into(), "VALUE".into());
|
||||
let json = build_stdio_json("node", &[], &env);
|
||||
assert_eq!(json["command"], "node");
|
||||
assert_eq!(json["env"]["KEY"], "VALUE");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,409 @@
|
||||
use std::collections::HashMap;
|
||||
use std::time::Duration;
|
||||
|
||||
use nomifun_runtime::Builder as CmdBuilder;
|
||||
use nomifun_runtime::resolve_command_path;
|
||||
|
||||
use crate::adapter::DetectedServer;
|
||||
use crate::error::McpError;
|
||||
use crate::types::McpServerTransport;
|
||||
|
||||
/// Timeout for detect/list operations (30 seconds).
|
||||
pub const DETECT_TIMEOUT: Duration = Duration::from_secs(30);
|
||||
|
||||
/// Timeout for install/remove operations (5 seconds).
|
||||
pub const MUTATE_TIMEOUT: Duration = Duration::from_secs(5);
|
||||
|
||||
/// Check whether a CLI binary is available on `$PATH`.
|
||||
///
|
||||
/// Uses `nomifun_runtime::resolve_command_path` so the lookup respects
|
||||
/// the bundled-bun shim and Windows `PATHEXT` rules. Previously this
|
||||
/// shelled out to `which`, which does not exist on Windows and made
|
||||
/// every MCP adapter report "not installed" there.
|
||||
pub async fn is_cli_installed(name: &str) -> Result<bool, McpError> {
|
||||
Ok(resolve_command_path(name).is_some())
|
||||
}
|
||||
|
||||
/// Run a CLI command with a timeout and clean environment variables.
|
||||
///
|
||||
/// Returns `(stdout, stderr)` on success. Returns an error if the command
|
||||
/// fails to start, times out, or exits with a non-zero status.
|
||||
pub async fn run_cli(program: &str, args: &[&str], timeout: Duration) -> Result<(String, String), McpError> {
|
||||
let mut builder = CmdBuilder::clean_cli(program);
|
||||
builder.args(args);
|
||||
let result = tokio::time::timeout(timeout, builder.output()).await;
|
||||
|
||||
let output = match result {
|
||||
Ok(Ok(output)) => output,
|
||||
Ok(Err(e)) => {
|
||||
return Err(McpError::AgentOperationFailed(format!(
|
||||
"`{program}` failed to start: {e}"
|
||||
)));
|
||||
}
|
||||
Err(_) => {
|
||||
return Err(McpError::AgentOperationFailed(format!(
|
||||
"`{program} {}` timed out after {}s",
|
||||
args.join(" "),
|
||||
timeout.as_secs()
|
||||
)));
|
||||
}
|
||||
};
|
||||
|
||||
let stdout = String::from_utf8_lossy(&output.stdout).to_string();
|
||||
let stderr = String::from_utf8_lossy(&output.stderr).to_string();
|
||||
|
||||
// Non-zero exit is not always fatal — callers inspect stdout/stderr.
|
||||
Ok((stdout, stderr))
|
||||
}
|
||||
|
||||
/// Run a CLI command and require zero exit status.
|
||||
pub async fn run_cli_strict(program: &str, args: &[&str], timeout: Duration) -> Result<String, McpError> {
|
||||
let mut builder = CmdBuilder::clean_cli(program);
|
||||
builder.args(args);
|
||||
let result = tokio::time::timeout(timeout, builder.output()).await;
|
||||
|
||||
let output = match result {
|
||||
Ok(Ok(output)) => output,
|
||||
Ok(Err(e)) => {
|
||||
return Err(McpError::AgentOperationFailed(format!(
|
||||
"`{program}` failed to start: {e}"
|
||||
)));
|
||||
}
|
||||
Err(_) => {
|
||||
return Err(McpError::AgentOperationFailed(format!(
|
||||
"`{program} {}` timed out after {}s",
|
||||
args.join(" "),
|
||||
timeout.as_secs()
|
||||
)));
|
||||
}
|
||||
};
|
||||
|
||||
let stdout = String::from_utf8_lossy(&output.stdout).to_string();
|
||||
|
||||
if !output.status.success() {
|
||||
let stderr = String::from_utf8_lossy(&output.stderr);
|
||||
return Err(McpError::AgentOperationFailed(format!(
|
||||
"`{program} {}` exited with {}: {}",
|
||||
args.join(" "),
|
||||
output.status,
|
||||
if stderr.is_empty() { &stdout } else { stderr.as_ref() }
|
||||
)));
|
||||
}
|
||||
|
||||
Ok(stdout)
|
||||
}
|
||||
|
||||
/// Strip ANSI escape codes from CLI output.
|
||||
pub fn strip_ansi(input: &str) -> String {
|
||||
// Matches: ESC[ ... m (SGR sequences) and other CSI sequences.
|
||||
let mut result = String::with_capacity(input.len());
|
||||
let mut chars = input.chars().peekable();
|
||||
|
||||
while let Some(ch) = chars.next() {
|
||||
if ch == '\x1b' {
|
||||
// Consume the '[' and everything until a letter in @ ..~ range.
|
||||
if chars.peek() == Some(&'[') {
|
||||
chars.next(); // consume '['
|
||||
while let Some(&c) = chars.peek() {
|
||||
chars.next();
|
||||
if c.is_ascii_alphabetic() || c == '~' || c == '@' {
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
result.push(ch);
|
||||
}
|
||||
}
|
||||
|
||||
result
|
||||
}
|
||||
|
||||
/// Normalize a CLI-reported MCP status string by stripping leading symbols
|
||||
/// such as `✓`, `✗`, `!`, bullets, and extra whitespace.
|
||||
pub fn normalize_detection_status(status: &str) -> String {
|
||||
status
|
||||
.trim()
|
||||
.trim_start_matches(|c: char| {
|
||||
matches!(c, '✓' | '✗' | '!' | '•' | '-' | '*' | '✔' | '✘' | ':' | '[' | ']') || c.is_whitespace()
|
||||
})
|
||||
.trim()
|
||||
.to_owned()
|
||||
}
|
||||
|
||||
/// Parse the "standard" `mcp list` text output shared by Gemini and Qwen.
|
||||
///
|
||||
/// Pattern: `[checkmark] name: command (transport_type) - Status`
|
||||
///
|
||||
/// Each matching line produces a `DetectedServer`.
|
||||
pub fn parse_standard_list_output(output: &str) -> Vec<DetectedServer> {
|
||||
let cleaned = strip_ansi(output);
|
||||
let mut servers = Vec::new();
|
||||
|
||||
for line in cleaned.lines() {
|
||||
let trimmed = line.trim();
|
||||
if let Some(server) = parse_standard_list_line(trimmed) {
|
||||
servers.push(server);
|
||||
}
|
||||
}
|
||||
|
||||
servers
|
||||
}
|
||||
|
||||
/// Parse a single line of standard list output.
|
||||
///
|
||||
/// Expected pattern:
|
||||
/// `[✓|✗] <name>: <command_or_url> (<transport_type>) - <Status>`
|
||||
fn parse_standard_list_line(line: &str) -> Option<DetectedServer> {
|
||||
// Must start with a check/cross mark (Unicode or ASCII fallback)
|
||||
let rest = if line.starts_with('\u{2713}') || line.starts_with('\u{2717}') {
|
||||
&line[3..] // UTF-8 multibyte ✓/✗
|
||||
} else if line.starts_with("✓") || line.starts_with("✗") {
|
||||
// Already handled above via char check
|
||||
return parse_standard_list_line_inner(line);
|
||||
} else {
|
||||
return None;
|
||||
};
|
||||
|
||||
parse_standard_list_line_inner_rest(rest.trim())
|
||||
}
|
||||
|
||||
fn parse_standard_list_line_inner(line: &str) -> Option<DetectedServer> {
|
||||
// Skip the leading mark character
|
||||
let rest = line.trim_start_matches(|c: char| !c.is_alphanumeric() && c != '_' && c != '-');
|
||||
parse_standard_list_line_inner_rest(rest)
|
||||
}
|
||||
|
||||
fn parse_standard_list_line_inner_rest(rest: &str) -> Option<DetectedServer> {
|
||||
// Find "name: command_or_url (type) - Status"
|
||||
let status_sep = rest.rfind(" - ")?;
|
||||
let status = normalize_detection_status(&rest[status_sep + 3..]);
|
||||
|
||||
let rest = &rest[..status_sep];
|
||||
|
||||
let colon_pos = rest.find(':')?;
|
||||
let name = rest[..colon_pos].trim();
|
||||
if name.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let after_colon = rest[colon_pos + 1..].trim();
|
||||
|
||||
// Find the transport type in parentheses
|
||||
let paren_open = after_colon.rfind('(')?;
|
||||
let paren_close = after_colon.rfind(')')?;
|
||||
if paren_close <= paren_open {
|
||||
return None;
|
||||
}
|
||||
|
||||
let transport_type = after_colon[paren_open + 1..paren_close].trim();
|
||||
let command_or_url = after_colon[..paren_open].trim();
|
||||
|
||||
let transport = match transport_type {
|
||||
"stdio" => McpServerTransport::Stdio {
|
||||
command: command_or_url.to_owned(),
|
||||
args: Vec::new(),
|
||||
env: HashMap::new(),
|
||||
},
|
||||
"sse" => McpServerTransport::Sse {
|
||||
url: command_or_url.to_owned(),
|
||||
headers: HashMap::new(),
|
||||
},
|
||||
"http" | "streamable_http" => McpServerTransport::Http {
|
||||
url: command_or_url.to_owned(),
|
||||
headers: HashMap::new(),
|
||||
},
|
||||
_ => return None,
|
||||
};
|
||||
|
||||
Some(DetectedServer {
|
||||
name: name.to_owned(),
|
||||
transport,
|
||||
importable: status.eq_ignore_ascii_case("connected"),
|
||||
import_skip_reason: if status.eq_ignore_ascii_case("connected") {
|
||||
None
|
||||
} else {
|
||||
Some(status)
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
/// Build `--env "KEY=VALUE"` argument pairs for CLI commands.
|
||||
pub fn build_env_args(env: &HashMap<String, String>, flag: &str) -> Vec<String> {
|
||||
env.iter()
|
||||
.flat_map(|(k, v)| [flag.to_owned(), format!("{k}={v}")])
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Build `--header "Key: Value"` or `-H "Key: Value"` argument pairs.
|
||||
pub fn build_header_args(headers: &HashMap<String, String>, flag: &str) -> Vec<String> {
|
||||
headers
|
||||
.iter()
|
||||
.flat_map(|(k, v)| [flag.to_owned(), format!("{k}: {v}")])
|
||||
.collect()
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn is_cli_installed_finds_known_binary() {
|
||||
// Both platforms ship a usable shell on PATH out of the box: `sh`
|
||||
// on Unix, `cmd` on Windows. resolve_command_path must locate it.
|
||||
#[cfg(unix)]
|
||||
let probe = "sh";
|
||||
#[cfg(windows)]
|
||||
let probe = "cmd";
|
||||
|
||||
assert!(is_cli_installed(probe).await.unwrap(), "expected `{probe}` on PATH");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn is_cli_installed_returns_false_for_missing_binary() {
|
||||
let result = is_cli_installed("nomifun-definitely-not-a-real-binary-xyz")
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(!result);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn strip_ansi_removes_color_codes() {
|
||||
let input = "\x1b[32m✓\x1b[0m my-server: npx (stdio) - \x1b[32mConnected\x1b[0m";
|
||||
let cleaned = strip_ansi(input);
|
||||
assert_eq!(cleaned, "✓ my-server: npx (stdio) - Connected");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn strip_ansi_preserves_plain_text() {
|
||||
let input = "hello world";
|
||||
assert_eq!(strip_ansi(input), "hello world");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn strip_ansi_handles_complex_sequences() {
|
||||
let input = "\x1b[1;34mBold Blue\x1b[0m normal \x1b[38;5;196mRed\x1b[0m";
|
||||
assert_eq!(strip_ansi(input), "Bold Blue normal Red");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalize_detection_status_strips_prefix_symbols() {
|
||||
assert_eq!(normalize_detection_status("✓ Connected"), "Connected");
|
||||
assert_eq!(normalize_detection_status("✗ Failed to connect"), "Failed to connect");
|
||||
assert_eq!(
|
||||
normalize_detection_status("! Needs authentication"),
|
||||
"Needs authentication"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_standard_list_stdio() {
|
||||
let output = "✓ my-server: npx -y @test/server (stdio) - Connected";
|
||||
let servers = parse_standard_list_output(output);
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert_eq!(servers[0].name, "my-server");
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Stdio { command, .. } => {
|
||||
assert_eq!(command, "npx -y @test/server");
|
||||
}
|
||||
_ => panic!("expected Stdio"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_standard_list_http() {
|
||||
let output = "✗ remote-srv: https://example.com/mcp (http) - Disconnected";
|
||||
let servers = parse_standard_list_output(output);
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert!(!servers[0].importable);
|
||||
assert_eq!(servers[0].import_skip_reason.as_deref(), Some("Disconnected"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_standard_list_sse() {
|
||||
let output = "✓ sse-srv: https://example.com/sse (sse) - Connected";
|
||||
let servers = parse_standard_list_output(output);
|
||||
assert_eq!(servers.len(), 1);
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Sse { url, .. } => {
|
||||
assert_eq!(url, "https://example.com/sse");
|
||||
}
|
||||
_ => panic!("expected Sse"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_standard_list_multiple_servers() {
|
||||
let output = "\
|
||||
Configured MCP servers:
|
||||
✓ server-a: npx -y @a/srv (stdio) - Connected
|
||||
✗ server-b: https://b.com/mcp (http) - Disconnected
|
||||
✓ server-c: https://c.com/sse (sse) - Connected
|
||||
Some footer text";
|
||||
let servers = parse_standard_list_output(output);
|
||||
assert_eq!(servers.len(), 3);
|
||||
assert_eq!(servers[0].name, "server-a");
|
||||
assert_eq!(servers[1].name, "server-b");
|
||||
assert!(!servers[1].importable);
|
||||
assert_eq!(servers[2].name, "server-c");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_standard_list_with_ansi() {
|
||||
let output = "\x1b[32m✓\x1b[0m my-mcp: npx -y @test/mcp (stdio) - \x1b[32mConnected\x1b[0m";
|
||||
let servers = parse_standard_list_output(output);
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert_eq!(servers[0].name, "my-mcp");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_standard_list_empty_output() {
|
||||
let servers = parse_standard_list_output("");
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_standard_list_no_matching_lines() {
|
||||
let output = "No MCP servers configured.\nTry `mcp add` to get started.";
|
||||
let servers = parse_standard_list_output(output);
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_standard_list_unknown_transport_skipped() {
|
||||
let output = "✓ srv: cmd (websocket) - Connected";
|
||||
let servers = parse_standard_list_output(output);
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_env_args_produces_pairs() {
|
||||
let mut env = HashMap::new();
|
||||
env.insert("K1".into(), "V1".into());
|
||||
let args = build_env_args(&env, "--env");
|
||||
assert_eq!(args.len(), 2);
|
||||
assert_eq!(args[0], "--env");
|
||||
assert_eq!(args[1], "K1=V1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_env_args_empty() {
|
||||
let env = HashMap::new();
|
||||
let args = build_env_args(&env, "--env");
|
||||
assert!(args.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_header_args_produces_pairs() {
|
||||
let mut headers = HashMap::new();
|
||||
headers.insert("Authorization".into(), "Bearer tok".into());
|
||||
let args = build_header_args(&headers, "--header");
|
||||
assert_eq!(args.len(), 2);
|
||||
assert_eq!(args[0], "--header");
|
||||
assert_eq!(args[1], "Authorization: Bearer tok");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,398 @@
|
||||
use std::collections::HashMap;
|
||||
|
||||
use nomifun_common::McpSource;
|
||||
|
||||
use crate::adapter::{DetectedServer, McpAgentAdapter};
|
||||
use crate::error::McpError;
|
||||
use crate::types::McpServerTransport;
|
||||
|
||||
use super::cli_helpers::{MUTATE_TIMEOUT, is_cli_installed, run_cli};
|
||||
|
||||
const CLI_NAME: &str = "codebuddy";
|
||||
|
||||
/// Scopes to try when removing a server.
|
||||
const REMOVE_SCOPES: &[&str] = &["user", "local", "project"];
|
||||
|
||||
/// MCP Agent adapter for CodeBuddy CLI.
|
||||
///
|
||||
/// # Detection
|
||||
///
|
||||
/// Detection reads `~/.codebuddy/mcp.json` directly (JSON) rather than
|
||||
/// parsing CLI text output. The file format:
|
||||
///
|
||||
/// ```json
|
||||
/// { "mcpServers": { "name": { "command": "...", "args": [...], ... } } }
|
||||
/// ```
|
||||
///
|
||||
/// # CLI Commands
|
||||
///
|
||||
/// - **install (stdio)**: `codebuddy mcp add -s user <name> <cmd> [-- args...] [-e K=V...]`
|
||||
/// - **install (http)**: `codebuddy mcp add-json -s user <name> <json>`
|
||||
/// - **remove**: `codebuddy mcp remove -s <scope> <name>` (tries user → local → project)
|
||||
pub struct CodeBuddyAdapter;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl McpAgentAdapter for CodeBuddyAdapter {
|
||||
fn source(&self) -> McpSource {
|
||||
McpSource::CodeBuddy
|
||||
}
|
||||
|
||||
async fn is_installed(&self) -> Result<bool, McpError> {
|
||||
is_cli_installed(CLI_NAME).await
|
||||
}
|
||||
|
||||
async fn detect_existing(&self) -> Result<Vec<DetectedServer>, McpError> {
|
||||
if !self.is_installed().await? {
|
||||
return Err(McpError::AgentNotInstalled(CLI_NAME.into()));
|
||||
}
|
||||
|
||||
let config_path = config_file_path()?;
|
||||
if !config_path.exists() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let content = tokio::fs::read_to_string(&config_path)
|
||||
.await
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("read codebuddy config: {e}")))?;
|
||||
|
||||
parse_codebuddy_config(&content)
|
||||
}
|
||||
|
||||
async fn install_server(&self, name: &str, transport: &McpServerTransport) -> Result<(), McpError> {
|
||||
if !self.is_installed().await? {
|
||||
return Err(McpError::AgentNotInstalled(CLI_NAME.into()));
|
||||
}
|
||||
|
||||
match transport {
|
||||
McpServerTransport::Stdio { command, args, env } => {
|
||||
let mut cli_args = vec![
|
||||
"mcp".to_owned(),
|
||||
"add".to_owned(),
|
||||
"-s".to_owned(),
|
||||
"user".to_owned(),
|
||||
name.to_owned(),
|
||||
command.clone(),
|
||||
];
|
||||
|
||||
// Separate args from env with --
|
||||
if !args.is_empty() {
|
||||
cli_args.push("--".to_owned());
|
||||
cli_args.extend(args.iter().cloned());
|
||||
}
|
||||
|
||||
// Env vars as -e KEY=VALUE
|
||||
for (k, v) in env {
|
||||
cli_args.push("-e".to_owned());
|
||||
cli_args.push(format!("{k}={v}"));
|
||||
}
|
||||
|
||||
let arg_refs: Vec<&str> = cli_args.iter().map(|s| s.as_str()).collect();
|
||||
run_cli(CLI_NAME, &arg_refs, MUTATE_TIMEOUT).await?;
|
||||
}
|
||||
McpServerTransport::Sse { url, headers } => {
|
||||
let config = build_http_json("sse", url, headers);
|
||||
install_via_add_json(name, &config).await?;
|
||||
}
|
||||
McpServerTransport::Http { url, headers } => {
|
||||
let config = build_http_json("streamable-http", url, headers);
|
||||
install_via_add_json(name, &config).await?;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn remove_server(&self, name: &str) -> Result<(), McpError> {
|
||||
if !self.is_installed().await? {
|
||||
return Err(McpError::AgentNotInstalled(CLI_NAME.into()));
|
||||
}
|
||||
|
||||
for scope in REMOVE_SCOPES {
|
||||
let (stdout, _stderr) = run_cli(CLI_NAME, &["mcp", "remove", "-s", scope, name], MUTATE_TIMEOUT).await?;
|
||||
let lower = stdout.to_lowercase();
|
||||
if lower.contains("removed") || lower.contains("not found") {
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// Install via `codebuddy mcp add-json -s user <name> <json>`.
|
||||
async fn install_via_add_json(name: &str, config: &serde_json::Value) -> Result<(), McpError> {
|
||||
let config_str = serde_json::to_string(config).map_err(|e| McpError::AgentOperationFailed(e.to_string()))?;
|
||||
run_cli(
|
||||
CLI_NAME,
|
||||
&["mcp", "add-json", "-s", "user", name, &config_str],
|
||||
MUTATE_TIMEOUT,
|
||||
)
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Build JSON config for HTTP-like servers.
|
||||
fn build_http_json(transport_type: &str, url: &str, headers: &HashMap<String, String>) -> serde_json::Value {
|
||||
let mut config = serde_json::json!({
|
||||
"url": url,
|
||||
"transportType": transport_type,
|
||||
});
|
||||
if !headers.is_empty() {
|
||||
config["headers"] = serde_json::json!(headers);
|
||||
}
|
||||
config
|
||||
}
|
||||
|
||||
/// Get the CodeBuddy config file path: `~/.codebuddy/mcp.json`.
|
||||
fn config_file_path() -> Result<std::path::PathBuf, McpError> {
|
||||
let home =
|
||||
dirs::home_dir().ok_or_else(|| McpError::AgentOperationFailed("cannot determine home directory".into()))?;
|
||||
Ok(home.join(".codebuddy").join("mcp.json"))
|
||||
}
|
||||
|
||||
/// Parse the CodeBuddy `mcp.json` config file.
|
||||
///
|
||||
/// Format:
|
||||
/// ```json
|
||||
/// {
|
||||
/// "mcpServers": {
|
||||
/// "name": {
|
||||
/// "command": "...",
|
||||
/// "args": [...],
|
||||
/// "env": { ... },
|
||||
/// "disabled": false,
|
||||
/// "url": "...",
|
||||
/// "transportType": "streamable-http",
|
||||
/// "headers": { ... }
|
||||
/// }
|
||||
/// }
|
||||
/// }
|
||||
/// ```
|
||||
fn parse_codebuddy_config(content: &str) -> Result<Vec<DetectedServer>, McpError> {
|
||||
let config: serde_json::Value = serde_json::from_str(content).map_err(McpError::from)?;
|
||||
|
||||
let servers_obj = match config.get("mcpServers").and_then(|v| v.as_object()) {
|
||||
Some(obj) => obj,
|
||||
None => return Ok(Vec::new()),
|
||||
};
|
||||
|
||||
let mut servers = Vec::new();
|
||||
|
||||
for (name, entry) in servers_obj {
|
||||
if let Some(transport) = parse_codebuddy_entry(entry) {
|
||||
let disabled = entry.get("disabled").and_then(|v| v.as_bool()).unwrap_or(false);
|
||||
servers.push(DetectedServer {
|
||||
name: name.clone(),
|
||||
transport,
|
||||
importable: !disabled,
|
||||
import_skip_reason: if disabled { Some("Disabled".into()) } else { None },
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
Ok(servers)
|
||||
}
|
||||
|
||||
/// Parse a single CodeBuddy config entry into a transport.
|
||||
fn parse_codebuddy_entry(entry: &serde_json::Value) -> Option<McpServerTransport> {
|
||||
let has_command = entry.get("command").and_then(|v| v.as_str()).is_some();
|
||||
let has_url = entry.get("url").and_then(|v| v.as_str()).is_some();
|
||||
|
||||
if has_command {
|
||||
let command = entry["command"].as_str()?.to_owned();
|
||||
let args = entry
|
||||
.get("args")
|
||||
.and_then(|v| v.as_array())
|
||||
.map(|arr| arr.iter().filter_map(|v| v.as_str().map(String::from)).collect())
|
||||
.unwrap_or_default();
|
||||
let env = parse_string_map(entry.get("env"));
|
||||
Some(McpServerTransport::Stdio { command, args, env })
|
||||
} else if has_url {
|
||||
let url = entry["url"].as_str()?.to_owned();
|
||||
let headers = parse_string_map(entry.get("headers"));
|
||||
let transport_type = entry
|
||||
.get("transportType")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("streamable-http");
|
||||
|
||||
// Normalize transport type
|
||||
match transport_type {
|
||||
"sse" => Some(McpServerTransport::Sse { url, headers }),
|
||||
_ => Some(McpServerTransport::Http { url, headers }),
|
||||
}
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
/// Parse a JSON object as `HashMap<String, String>`.
|
||||
fn parse_string_map(value: Option<&serde_json::Value>) -> HashMap<String, String> {
|
||||
value
|
||||
.and_then(|v| v.as_object())
|
||||
.map(|obj| {
|
||||
obj.iter()
|
||||
.filter_map(|(k, v)| v.as_str().map(|s| (k.clone(), s.to_owned())))
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn source_is_codebuddy() {
|
||||
assert_eq!(CodeBuddyAdapter.source(), McpSource::CodeBuddy);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_config_stdio_server() {
|
||||
let config = r#"{
|
||||
"mcpServers": {
|
||||
"test-server": {
|
||||
"command": "npx",
|
||||
"args": ["-y", "@test/server"],
|
||||
"env": { "NODE_ENV": "production" }
|
||||
}
|
||||
}
|
||||
}"#;
|
||||
let servers = parse_codebuddy_config(config).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert_eq!(servers[0].name, "test-server");
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Stdio { command, args, env } => {
|
||||
assert_eq!(command, "npx");
|
||||
assert_eq!(args, &["-y", "@test/server"]);
|
||||
assert_eq!(env.get("NODE_ENV").unwrap(), "production");
|
||||
}
|
||||
_ => panic!("expected Stdio"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_config_http_server() {
|
||||
let config = r#"{
|
||||
"mcpServers": {
|
||||
"remote": {
|
||||
"url": "https://example.com/mcp",
|
||||
"transportType": "streamable-http",
|
||||
"headers": { "Authorization": "Bearer tok" }
|
||||
}
|
||||
}
|
||||
}"#;
|
||||
let servers = parse_codebuddy_config(config).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert_eq!(servers[0].name, "remote");
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Http { url, headers } => {
|
||||
assert_eq!(url, "https://example.com/mcp");
|
||||
assert_eq!(headers.get("Authorization").unwrap(), "Bearer tok");
|
||||
}
|
||||
_ => panic!("expected Http"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_config_sse_server() {
|
||||
let config = r#"{
|
||||
"mcpServers": {
|
||||
"sse-srv": {
|
||||
"url": "https://example.com/sse",
|
||||
"transportType": "sse"
|
||||
}
|
||||
}
|
||||
}"#;
|
||||
let servers = parse_codebuddy_config(config).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Sse { url, .. } => {
|
||||
assert_eq!(url, "https://example.com/sse");
|
||||
}
|
||||
_ => panic!("expected Sse"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_config_skips_disabled() {
|
||||
let config = r#"{
|
||||
"mcpServers": {
|
||||
"active": { "command": "npx", "args": [] },
|
||||
"disabled": { "command": "npx", "args": [], "disabled": true }
|
||||
}
|
||||
}"#;
|
||||
let servers = parse_codebuddy_config(config).unwrap();
|
||||
assert_eq!(servers.len(), 2);
|
||||
assert_eq!(servers[0].name, "active");
|
||||
assert!(servers[0].importable);
|
||||
assert_eq!(servers[1].name, "disabled");
|
||||
assert!(!servers[1].importable);
|
||||
assert_eq!(servers[1].import_skip_reason.as_deref(), Some("Disabled"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_config_empty_mcp_servers() {
|
||||
let config = r#"{ "mcpServers": {} }"#;
|
||||
let servers = parse_codebuddy_config(config).unwrap();
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_config_no_mcp_servers_key() {
|
||||
let config = r#"{ "otherKey": 42 }"#;
|
||||
let servers = parse_codebuddy_config(config).unwrap();
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_config_multiple_servers() {
|
||||
let config = r#"{
|
||||
"mcpServers": {
|
||||
"stdio-srv": { "command": "node", "args": ["index.js"] },
|
||||
"http-srv": { "url": "https://a.com/mcp" },
|
||||
"sse-srv": { "url": "https://b.com/sse", "transportType": "sse" }
|
||||
}
|
||||
}"#;
|
||||
let servers = parse_codebuddy_config(config).unwrap();
|
||||
assert_eq!(servers.len(), 3);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_config_url_without_transport_type_defaults_to_http() {
|
||||
let config = r#"{
|
||||
"mcpServers": {
|
||||
"no-type": { "url": "https://example.com/api" }
|
||||
}
|
||||
}"#;
|
||||
let servers = parse_codebuddy_config(config).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert!(matches!(servers[0].transport, McpServerTransport::Http { .. }));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_http_json_without_headers() {
|
||||
let json = build_http_json("streamable-http", "https://example.com", &HashMap::new());
|
||||
assert_eq!(json["url"], "https://example.com");
|
||||
assert_eq!(json["transportType"], "streamable-http");
|
||||
assert!(json.get("headers").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_http_json_with_headers() {
|
||||
let mut headers = HashMap::new();
|
||||
headers.insert("Authorization".into(), "Bearer tok".into());
|
||||
let json = build_http_json("sse", "https://example.com/sse", &headers);
|
||||
assert_eq!(json["transportType"], "sse");
|
||||
assert_eq!(json["headers"]["Authorization"], "Bearer tok");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn trait_is_object_safe() {
|
||||
let adapter: Box<dyn McpAgentAdapter> = Box::new(CodeBuddyAdapter);
|
||||
assert_eq!(adapter.source(), McpSource::CodeBuddy);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,486 @@
|
||||
use std::collections::HashMap;
|
||||
|
||||
use nomifun_common::McpSource;
|
||||
|
||||
use crate::adapter::{DetectedServer, McpAgentAdapter};
|
||||
use crate::error::McpError;
|
||||
use crate::types::McpServerTransport;
|
||||
|
||||
use super::cli_helpers::{DETECT_TIMEOUT, MUTATE_TIMEOUT, is_cli_installed, run_cli_strict};
|
||||
|
||||
const CLI_NAME: &str = "codex";
|
||||
|
||||
/// MCP Agent adapter for Codex CLI.
|
||||
///
|
||||
/// # CLI Commands
|
||||
///
|
||||
/// - **detect**: `codex mcp list --json` (JSON output)
|
||||
/// - **install (stdio)**: `codex mcp add <name> [--env K=V]... -- <cmd> [args...]`
|
||||
/// - **install (http)**: `codex mcp add <name> --url <url>`
|
||||
/// - **remove**: `codex mcp remove <name>` (no scope parameter)
|
||||
///
|
||||
/// Codex outputs structured JSON for list, unlike the text-based agents.
|
||||
pub struct CodexAdapter;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl McpAgentAdapter for CodexAdapter {
|
||||
fn source(&self) -> McpSource {
|
||||
McpSource::Codex
|
||||
}
|
||||
|
||||
async fn is_installed(&self) -> Result<bool, McpError> {
|
||||
is_cli_installed(CLI_NAME).await
|
||||
}
|
||||
|
||||
async fn detect_existing(&self) -> Result<Vec<DetectedServer>, McpError> {
|
||||
if !self.is_installed().await? {
|
||||
return Err(McpError::AgentNotInstalled(CLI_NAME.into()));
|
||||
}
|
||||
|
||||
let stdout = run_cli_strict(CLI_NAME, &["mcp", "list", "--json"], DETECT_TIMEOUT).await?;
|
||||
|
||||
parse_codex_list_json(&stdout)
|
||||
}
|
||||
|
||||
async fn install_server(&self, name: &str, transport: &McpServerTransport) -> Result<(), McpError> {
|
||||
if !self.is_installed().await? {
|
||||
return Err(McpError::AgentNotInstalled(CLI_NAME.into()));
|
||||
}
|
||||
|
||||
match transport {
|
||||
McpServerTransport::Stdio { command, args, env } => {
|
||||
let mut cli_args = vec!["mcp".to_owned(), "add".to_owned(), name.to_owned()];
|
||||
|
||||
// Env vars come before --
|
||||
for (k, v) in env {
|
||||
cli_args.push("--env".to_owned());
|
||||
cli_args.push(format!("{k}={v}"));
|
||||
}
|
||||
|
||||
// Command and args come after --
|
||||
cli_args.push("--".to_owned());
|
||||
cli_args.push(command.clone());
|
||||
cli_args.extend(args.iter().cloned());
|
||||
|
||||
let arg_refs: Vec<&str> = cli_args.iter().map(|s| s.as_str()).collect();
|
||||
run_cli_strict(CLI_NAME, &arg_refs, MUTATE_TIMEOUT).await?;
|
||||
}
|
||||
McpServerTransport::Http { url, .. } | McpServerTransport::Sse { url, .. } => {
|
||||
// Codex only supports --url for HTTP, no headers via CLI
|
||||
run_cli_strict(CLI_NAME, &["mcp", "add", name, "--url", url], MUTATE_TIMEOUT).await?;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn remove_server(&self, name: &str) -> Result<(), McpError> {
|
||||
if !self.is_installed().await? {
|
||||
return Err(McpError::AgentNotInstalled(CLI_NAME.into()));
|
||||
}
|
||||
|
||||
// Codex has no scope parameter; remove is simple.
|
||||
let (stdout, _stderr) = super::cli_helpers::run_cli(CLI_NAME, &["mcp", "remove", name], MUTATE_TIMEOUT).await?;
|
||||
|
||||
// Idempotent: treat "not found" as success.
|
||||
let lower = stdout.to_lowercase();
|
||||
if lower.contains("not found") || lower.contains("removed") || lower.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// JSON parsing
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Parse the JSON output of `codex mcp list --json`.
|
||||
///
|
||||
/// Expected format: array of entries with transport details.
|
||||
///
|
||||
/// ```json
|
||||
/// [
|
||||
/// {
|
||||
/// "name": "...",
|
||||
/// "enabled": true,
|
||||
/// "transport": {
|
||||
/// "type": "stdio",
|
||||
/// "command": "...",
|
||||
/// "args": [...],
|
||||
/// "env": { ... },
|
||||
/// "env_vars": [{ "name": "...", "value": "..." }]
|
||||
/// }
|
||||
/// }
|
||||
/// ]
|
||||
/// ```
|
||||
fn parse_codex_list_json(json_str: &str) -> Result<Vec<DetectedServer>, McpError> {
|
||||
let trimmed = json_str.trim();
|
||||
if trimmed.is_empty() || trimmed == "[]" {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let entries: Vec<serde_json::Value> = serde_json::from_str(trimmed).map_err(McpError::from)?;
|
||||
|
||||
let mut servers = Vec::new();
|
||||
|
||||
for entry in &entries {
|
||||
if let Some(server) = parse_codex_entry(entry) {
|
||||
servers.push(server);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(servers)
|
||||
}
|
||||
|
||||
/// Parse a single Codex list entry.
|
||||
fn parse_codex_entry(entry: &serde_json::Value) -> Option<DetectedServer> {
|
||||
let name = entry.get("name")?.as_str()?.to_owned();
|
||||
let enabled = entry.get("enabled").and_then(|v| v.as_bool()).unwrap_or(true);
|
||||
|
||||
let transport_obj = entry.get("transport")?;
|
||||
let transport_type = transport_obj.get("type").and_then(|v| v.as_str()).unwrap_or("stdio");
|
||||
|
||||
let transport = match transport_type {
|
||||
"stdio" => {
|
||||
let command = transport_obj
|
||||
.get("command")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or_default()
|
||||
.to_owned();
|
||||
let args = transport_obj
|
||||
.get("args")
|
||||
.and_then(|v| v.as_array())
|
||||
.map(|arr| arr.iter().filter_map(|v| v.as_str().map(String::from)).collect())
|
||||
.unwrap_or_default();
|
||||
|
||||
// Codex supports both `env` (object) and `env_vars` (array of {name, value})
|
||||
let env = parse_codex_env(transport_obj);
|
||||
|
||||
McpServerTransport::Stdio { command, args, env }
|
||||
}
|
||||
"http" | "streamable_http" => {
|
||||
let url = transport_obj
|
||||
.get("url")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or_default()
|
||||
.to_owned();
|
||||
McpServerTransport::Http {
|
||||
url,
|
||||
headers: HashMap::new(),
|
||||
}
|
||||
}
|
||||
"sse" => {
|
||||
let url = transport_obj
|
||||
.get("url")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or_default()
|
||||
.to_owned();
|
||||
McpServerTransport::Sse {
|
||||
url,
|
||||
headers: HashMap::new(),
|
||||
}
|
||||
}
|
||||
_ => return None,
|
||||
};
|
||||
|
||||
Some(DetectedServer {
|
||||
name,
|
||||
transport,
|
||||
importable: enabled,
|
||||
import_skip_reason: if enabled { None } else { Some("Disabled".into()) },
|
||||
})
|
||||
}
|
||||
|
||||
/// Parse environment variables from Codex transport.
|
||||
///
|
||||
/// Handles both formats:
|
||||
/// - `"env": { "KEY": "VALUE" }` (object)
|
||||
/// - `"env_vars": [{ "name": "KEY", "value": "VALUE" }]` (array)
|
||||
fn parse_codex_env(transport: &serde_json::Value) -> HashMap<String, String> {
|
||||
// Try object format first
|
||||
if let Some(obj) = transport.get("env").and_then(|v| v.as_object()) {
|
||||
return obj
|
||||
.iter()
|
||||
.filter_map(|(k, v)| v.as_str().map(|s| (k.clone(), s.to_owned())))
|
||||
.collect();
|
||||
}
|
||||
|
||||
// Try array format
|
||||
if let Some(arr) = transport.get("env_vars").and_then(|v| v.as_array()) {
|
||||
return arr
|
||||
.iter()
|
||||
.filter_map(|entry| {
|
||||
let name = entry.get("name")?.as_str()?;
|
||||
let value = entry.get("value")?.as_str()?;
|
||||
Some((name.to_owned(), value.to_owned()))
|
||||
})
|
||||
.collect();
|
||||
}
|
||||
|
||||
HashMap::new()
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn source_is_codex() {
|
||||
assert_eq!(CodexAdapter.source(), McpSource::Codex);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_empty_json() {
|
||||
let servers = parse_codex_list_json("[]").unwrap();
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_empty_string() {
|
||||
let servers = parse_codex_list_json("").unwrap();
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_stdio_server_with_env_object() {
|
||||
let json = r#"[
|
||||
{
|
||||
"name": "test-mcp",
|
||||
"enabled": true,
|
||||
"transport": {
|
||||
"type": "stdio",
|
||||
"command": "npx",
|
||||
"args": ["-y", "@test/server"],
|
||||
"env": { "NODE_ENV": "production" }
|
||||
}
|
||||
}
|
||||
]"#;
|
||||
let servers = parse_codex_list_json(json).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert_eq!(servers[0].name, "test-mcp");
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Stdio { command, args, env } => {
|
||||
assert_eq!(command, "npx");
|
||||
assert_eq!(args, &["-y", "@test/server"]);
|
||||
assert_eq!(env.get("NODE_ENV").unwrap(), "production");
|
||||
}
|
||||
_ => panic!("expected Stdio"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_stdio_server_with_env_vars_array() {
|
||||
let json = r#"[
|
||||
{
|
||||
"name": "test-mcp",
|
||||
"enabled": true,
|
||||
"transport": {
|
||||
"type": "stdio",
|
||||
"command": "node",
|
||||
"args": ["index.js"],
|
||||
"env_vars": [
|
||||
{ "name": "KEY1", "value": "VAL1" },
|
||||
{ "name": "KEY2", "value": "VAL2" }
|
||||
]
|
||||
}
|
||||
}
|
||||
]"#;
|
||||
let servers = parse_codex_list_json(json).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Stdio { env, .. } => {
|
||||
assert_eq!(env.get("KEY1").unwrap(), "VAL1");
|
||||
assert_eq!(env.get("KEY2").unwrap(), "VAL2");
|
||||
}
|
||||
_ => panic!("expected Stdio"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_http_server() {
|
||||
let json = r#"[
|
||||
{
|
||||
"name": "remote",
|
||||
"enabled": true,
|
||||
"transport": {
|
||||
"type": "http",
|
||||
"url": "https://example.com/mcp"
|
||||
}
|
||||
}
|
||||
]"#;
|
||||
let servers = parse_codex_list_json(json).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Http { url, .. } => {
|
||||
assert_eq!(url, "https://example.com/mcp");
|
||||
}
|
||||
_ => panic!("expected Http"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_streamable_http_becomes_http() {
|
||||
let json = r#"[
|
||||
{
|
||||
"name": "streamable",
|
||||
"enabled": true,
|
||||
"transport": {
|
||||
"type": "streamable_http",
|
||||
"url": "https://example.com/api"
|
||||
}
|
||||
}
|
||||
]"#;
|
||||
let servers = parse_codex_list_json(json).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert!(matches!(servers[0].transport, McpServerTransport::Http { .. }));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_sse_server() {
|
||||
let json = r#"[
|
||||
{
|
||||
"name": "sse-srv",
|
||||
"enabled": true,
|
||||
"transport": {
|
||||
"type": "sse",
|
||||
"url": "https://example.com/sse"
|
||||
}
|
||||
}
|
||||
]"#;
|
||||
let servers = parse_codex_list_json(json).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Sse { url, .. } => {
|
||||
assert_eq!(url, "https://example.com/sse");
|
||||
}
|
||||
_ => panic!("expected Sse"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_multiple_servers() {
|
||||
let json = r#"[
|
||||
{
|
||||
"name": "stdio-srv",
|
||||
"enabled": true,
|
||||
"transport": { "type": "stdio", "command": "node", "args": [] }
|
||||
},
|
||||
{
|
||||
"name": "http-srv",
|
||||
"enabled": true,
|
||||
"transport": { "type": "http", "url": "https://a.com/mcp" }
|
||||
}
|
||||
]"#;
|
||||
let servers = parse_codex_list_json(json).unwrap();
|
||||
assert_eq!(servers.len(), 2);
|
||||
assert_eq!(servers[0].name, "stdio-srv");
|
||||
assert_eq!(servers[1].name, "http-srv");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_disabled_server_skipped_from_import_only() {
|
||||
let json = r#"[
|
||||
{
|
||||
"name": "disabled-srv",
|
||||
"enabled": false,
|
||||
"transport": { "type": "stdio", "command": "node", "args": [] }
|
||||
}
|
||||
]"#;
|
||||
let servers = parse_codex_list_json(json).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert!(!servers[0].importable);
|
||||
assert_eq!(servers[0].import_skip_reason.as_deref(), Some("Disabled"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_unknown_transport_skipped() {
|
||||
let json = r#"[
|
||||
{
|
||||
"name": "unknown",
|
||||
"enabled": true,
|
||||
"transport": { "type": "websocket", "url": "ws://localhost" }
|
||||
}
|
||||
]"#;
|
||||
let servers = parse_codex_list_json(json).unwrap();
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_missing_name_skipped() {
|
||||
let json = r#"[
|
||||
{
|
||||
"enabled": true,
|
||||
"transport": { "type": "stdio", "command": "node" }
|
||||
}
|
||||
]"#;
|
||||
let servers = parse_codex_list_json(json).unwrap();
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_default_type_is_stdio() {
|
||||
let json = r#"[
|
||||
{
|
||||
"name": "no-type",
|
||||
"enabled": true,
|
||||
"transport": { "command": "node", "args": ["srv.js"] }
|
||||
}
|
||||
]"#;
|
||||
let servers = parse_codex_list_json(json).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert!(matches!(servers[0].transport, McpServerTransport::Stdio { .. }));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_env_object_takes_precedence() {
|
||||
let json = r#"[
|
||||
{
|
||||
"name": "both-env",
|
||||
"enabled": true,
|
||||
"transport": {
|
||||
"type": "stdio",
|
||||
"command": "node",
|
||||
"env": { "FROM_OBJ": "yes" },
|
||||
"env_vars": [{ "name": "FROM_ARR", "value": "yes" }]
|
||||
}
|
||||
}
|
||||
]"#;
|
||||
let servers = parse_codex_list_json(json).unwrap();
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Stdio { env, .. } => {
|
||||
// Object format takes precedence
|
||||
assert_eq!(env.get("FROM_OBJ").unwrap(), "yes");
|
||||
assert!(env.get("FROM_ARR").is_none());
|
||||
}
|
||||
_ => panic!("expected Stdio"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_disabled_server_skipped() {
|
||||
let json = r#"[
|
||||
{
|
||||
"name": "disabled-mcp",
|
||||
"enabled": false,
|
||||
"transport": { "type": "stdio", "command": "node", "args": ["srv.js"] }
|
||||
}
|
||||
]"#;
|
||||
let servers = parse_codex_list_json(json).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert_eq!(servers[0].name, "disabled-mcp");
|
||||
assert!(!servers[0].importable);
|
||||
assert_eq!(servers[0].import_skip_reason.as_deref(), Some("Disabled"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn trait_is_object_safe() {
|
||||
let adapter: Box<dyn McpAgentAdapter> = Box::new(CodexAdapter);
|
||||
assert_eq!(adapter.source(), McpSource::Codex);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,145 @@
|
||||
use nomifun_common::McpSource;
|
||||
|
||||
use crate::adapter::{DetectedServer, McpAgentAdapter};
|
||||
use crate::error::McpError;
|
||||
use crate::types::McpServerTransport;
|
||||
|
||||
use super::cli_helpers::{DETECT_TIMEOUT, MUTATE_TIMEOUT, is_cli_installed, parse_standard_list_output, run_cli};
|
||||
|
||||
const CLI_NAME: &str = "gemini";
|
||||
|
||||
/// Scopes tried when removing (user first, then project).
|
||||
const REMOVE_SCOPES: &[&str] = &["user", "project"];
|
||||
|
||||
/// MCP Agent adapter for Gemini CLI.
|
||||
///
|
||||
/// # CLI Commands
|
||||
///
|
||||
/// - **detect**: `gemini mcp list`
|
||||
/// - **install (stdio)**: `gemini mcp add <name> <command> [args...] -s user`
|
||||
/// - **install (http/sse)**: `gemini mcp add <name> <url> --transport <type> -s user`
|
||||
/// - **remove**: `gemini mcp remove <name> -s user` (falls back to `-s project`)
|
||||
pub struct GeminiAdapter;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl McpAgentAdapter for GeminiAdapter {
|
||||
fn source(&self) -> McpSource {
|
||||
McpSource::Gemini
|
||||
}
|
||||
|
||||
async fn is_installed(&self) -> Result<bool, McpError> {
|
||||
is_cli_installed(CLI_NAME).await
|
||||
}
|
||||
|
||||
async fn detect_existing(&self) -> Result<Vec<DetectedServer>, McpError> {
|
||||
if !self.is_installed().await? {
|
||||
return Err(McpError::AgentNotInstalled(CLI_NAME.into()));
|
||||
}
|
||||
|
||||
let (stdout, _stderr) = run_cli(CLI_NAME, &["mcp", "list"], DETECT_TIMEOUT).await?;
|
||||
Ok(parse_standard_list_output(&stdout))
|
||||
}
|
||||
|
||||
async fn install_server(&self, name: &str, transport: &McpServerTransport) -> Result<(), McpError> {
|
||||
if !self.is_installed().await? {
|
||||
return Err(McpError::AgentNotInstalled(CLI_NAME.into()));
|
||||
}
|
||||
|
||||
match transport {
|
||||
McpServerTransport::Stdio { command, args, .. } => {
|
||||
let mut cli_args = vec!["mcp".to_owned(), "add".to_owned(), name.to_owned(), command.clone()];
|
||||
cli_args.extend(args.iter().cloned());
|
||||
cli_args.push("-s".to_owned());
|
||||
cli_args.push("user".to_owned());
|
||||
|
||||
let arg_refs: Vec<&str> = cli_args.iter().map(|s| s.as_str()).collect();
|
||||
run_cli(CLI_NAME, &arg_refs, MUTATE_TIMEOUT).await?;
|
||||
}
|
||||
McpServerTransport::Sse { url, .. } => {
|
||||
run_cli(
|
||||
CLI_NAME,
|
||||
&["mcp", "add", name, url, "--transport", "sse", "-s", "user"],
|
||||
MUTATE_TIMEOUT,
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
McpServerTransport::Http { url, .. } => {
|
||||
run_cli(
|
||||
CLI_NAME,
|
||||
&["mcp", "add", name, url, "--transport", "http", "-s", "user"],
|
||||
MUTATE_TIMEOUT,
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn remove_server(&self, name: &str) -> Result<(), McpError> {
|
||||
if !self.is_installed().await? {
|
||||
return Err(McpError::AgentNotInstalled(CLI_NAME.into()));
|
||||
}
|
||||
|
||||
for scope in REMOVE_SCOPES {
|
||||
let (stdout, _stderr) = run_cli(CLI_NAME, &["mcp", "remove", name, "-s", scope], MUTATE_TIMEOUT).await?;
|
||||
let lower = stdout.to_lowercase();
|
||||
if lower.contains("removed") || lower.contains("not found") {
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn source_is_gemini() {
|
||||
assert_eq!(GeminiAdapter.source(), McpSource::Gemini);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_gemini_list_output() {
|
||||
let output = "\
|
||||
Configured MCP servers:
|
||||
✓ my-server: npx -y @test/server (stdio) - Connected
|
||||
✗ broken: node bad.js (stdio) - Disconnected
|
||||
✓ remote: https://example.com/mcp (http) - Connected";
|
||||
|
||||
let servers = parse_standard_list_output(output);
|
||||
assert_eq!(servers.len(), 3);
|
||||
assert_eq!(servers[0].name, "my-server");
|
||||
assert_eq!(servers[1].name, "broken");
|
||||
assert_eq!(servers[2].name, "remote");
|
||||
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Stdio { command, .. } => {
|
||||
assert_eq!(command, "npx -y @test/server");
|
||||
}
|
||||
_ => panic!("expected Stdio"),
|
||||
}
|
||||
|
||||
match &servers[2].transport {
|
||||
McpServerTransport::Http { url, .. } => {
|
||||
assert_eq!(url, "https://example.com/mcp");
|
||||
}
|
||||
_ => panic!("expected Http"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_gemini_empty() {
|
||||
let servers = parse_standard_list_output("No MCP servers configured.");
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn trait_is_object_safe() {
|
||||
let adapter: Box<dyn McpAgentAdapter> = Box::new(GeminiAdapter);
|
||||
assert_eq!(adapter.source(), McpSource::Gemini);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
mod claude;
|
||||
mod cli_helpers;
|
||||
mod codebuddy;
|
||||
mod codex;
|
||||
mod gemini;
|
||||
mod nomi;
|
||||
mod nomifun;
|
||||
mod opencode;
|
||||
mod qwen;
|
||||
|
||||
pub use claude::ClaudeAdapter;
|
||||
pub use codebuddy::CodeBuddyAdapter;
|
||||
pub use codex::CodexAdapter;
|
||||
pub use gemini::GeminiAdapter;
|
||||
pub use nomi::NomiAdapter;
|
||||
pub use nomifun::NomifunAdapter;
|
||||
pub use opencode::OpencodeAdapter;
|
||||
pub use qwen::QwenAdapter;
|
||||
@@ -0,0 +1,603 @@
|
||||
use std::collections::HashMap;
|
||||
|
||||
use nomifun_common::McpSource;
|
||||
|
||||
use crate::adapter::{DetectedServer, McpAgentAdapter};
|
||||
use crate::error::McpError;
|
||||
use crate::types::McpServerTransport;
|
||||
|
||||
use super::cli_helpers::{DETECT_TIMEOUT, is_cli_installed, run_cli_strict};
|
||||
|
||||
const CLI_NAME: &str = "nomi";
|
||||
|
||||
/// MCP Agent adapter for Nomi.
|
||||
///
|
||||
/// Nomi stores MCP configuration in a TOML config file. The config path
|
||||
/// is obtained via `nomi --config-path`.
|
||||
///
|
||||
/// # Config Format (TOML)
|
||||
///
|
||||
/// ```toml
|
||||
/// [mcp.servers.server-name]
|
||||
/// transport = "stdio"
|
||||
/// command = "npx"
|
||||
/// args = ["-y", "@test/server"]
|
||||
///
|
||||
/// [mcp.servers.server-name.env]
|
||||
/// KEY = "VALUE"
|
||||
///
|
||||
/// [mcp.servers.remote-server]
|
||||
/// transport = "http"
|
||||
/// url = "https://example.com/mcp"
|
||||
///
|
||||
/// [mcp.servers.remote-server.headers]
|
||||
/// Authorization = "Bearer xxx"
|
||||
/// ```
|
||||
pub struct NomiAdapter;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl McpAgentAdapter for NomiAdapter {
|
||||
fn source(&self) -> McpSource {
|
||||
McpSource::Nomi
|
||||
}
|
||||
|
||||
async fn is_installed(&self) -> Result<bool, McpError> {
|
||||
is_cli_installed(CLI_NAME).await
|
||||
}
|
||||
|
||||
async fn detect_existing(&self) -> Result<Vec<DetectedServer>, McpError> {
|
||||
if !self.is_installed().await? {
|
||||
return Err(McpError::AgentNotInstalled(CLI_NAME.into()));
|
||||
}
|
||||
|
||||
let config_path = get_config_path().await?;
|
||||
let path = std::path::Path::new(&config_path);
|
||||
|
||||
if !path.exists() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let content = tokio::fs::read_to_string(path)
|
||||
.await
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("failed to read {config_path}: {e}")))?;
|
||||
|
||||
parse_toml_servers(&content)
|
||||
}
|
||||
|
||||
async fn install_server(&self, name: &str, transport: &McpServerTransport) -> Result<(), McpError> {
|
||||
if !self.is_installed().await? {
|
||||
return Err(McpError::AgentNotInstalled(CLI_NAME.into()));
|
||||
}
|
||||
|
||||
let config_path = get_config_path().await?;
|
||||
let path = std::path::Path::new(&config_path);
|
||||
|
||||
let mut doc = if path.exists() {
|
||||
let content = tokio::fs::read_to_string(path)
|
||||
.await
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("failed to read {config_path}: {e}")))?;
|
||||
content
|
||||
.parse::<toml::Value>()
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("failed to parse TOML: {e}")))?
|
||||
} else {
|
||||
// Ensure parent directory exists
|
||||
if let Some(parent) = path.parent() {
|
||||
tokio::fs::create_dir_all(parent)
|
||||
.await
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("failed to create dir: {e}")))?;
|
||||
}
|
||||
toml::Value::Table(toml::map::Map::new())
|
||||
};
|
||||
|
||||
// Ensure mcp.servers exists
|
||||
let root = doc
|
||||
.as_table_mut()
|
||||
.ok_or_else(|| McpError::AgentOperationFailed("TOML root is not a table".into()))?;
|
||||
|
||||
let mcp = root
|
||||
.entry("mcp")
|
||||
.or_insert_with(|| toml::Value::Table(toml::map::Map::new()));
|
||||
|
||||
let mcp_table = mcp
|
||||
.as_table_mut()
|
||||
.ok_or_else(|| McpError::AgentOperationFailed("mcp is not a table".into()))?;
|
||||
|
||||
let servers = mcp_table
|
||||
.entry("servers")
|
||||
.or_insert_with(|| toml::Value::Table(toml::map::Map::new()));
|
||||
|
||||
let servers_table = servers
|
||||
.as_table_mut()
|
||||
.ok_or_else(|| McpError::AgentOperationFailed("mcp.servers is not a table".into()))?;
|
||||
|
||||
servers_table.insert(name.to_owned(), transport_to_toml(transport));
|
||||
|
||||
let output = toml::to_string_pretty(&doc)
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("failed to serialize TOML: {e}")))?;
|
||||
|
||||
tokio::fs::write(path, output)
|
||||
.await
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("failed to write {config_path}: {e}")))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn remove_server(&self, name: &str) -> Result<(), McpError> {
|
||||
if !self.is_installed().await? {
|
||||
return Err(McpError::AgentNotInstalled(CLI_NAME.into()));
|
||||
}
|
||||
|
||||
let config_path = get_config_path().await?;
|
||||
let path = std::path::Path::new(&config_path);
|
||||
|
||||
if !path.exists() {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let content = tokio::fs::read_to_string(path)
|
||||
.await
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("failed to read {config_path}: {e}")))?;
|
||||
|
||||
let mut doc: toml::Value = content
|
||||
.parse()
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("failed to parse TOML: {e}")))?;
|
||||
|
||||
let removed = doc
|
||||
.as_table_mut()
|
||||
.and_then(|root| root.get_mut("mcp"))
|
||||
.and_then(|mcp| mcp.as_table_mut())
|
||||
.and_then(|mcp| mcp.get_mut("servers"))
|
||||
.and_then(|servers| servers.as_table_mut())
|
||||
.map(|servers| servers.remove(name).is_some())
|
||||
.unwrap_or(false);
|
||||
|
||||
if !removed {
|
||||
// Idempotent: not found is fine
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let output = toml::to_string_pretty(&doc)
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("failed to serialize TOML: {e}")))?;
|
||||
|
||||
tokio::fs::write(path, output)
|
||||
.await
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("failed to write {config_path}: {e}")))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Run `nomi --config-path` to get the TOML config file path.
|
||||
async fn get_config_path() -> Result<String, McpError> {
|
||||
let stdout = run_cli_strict(CLI_NAME, &["--config-path"], DETECT_TIMEOUT).await?;
|
||||
let path = stdout.trim().to_owned();
|
||||
if path.is_empty() {
|
||||
return Err(McpError::AgentOperationFailed(
|
||||
"nomi --config-path returned empty output".into(),
|
||||
));
|
||||
}
|
||||
Ok(path)
|
||||
}
|
||||
|
||||
/// Parse MCP servers from TOML config content.
|
||||
///
|
||||
/// Expects `[mcp.servers.<name>]` tables.
|
||||
fn parse_toml_servers(content: &str) -> Result<Vec<DetectedServer>, McpError> {
|
||||
let doc: toml::Value = content
|
||||
.parse()
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("failed to parse TOML: {e}")))?;
|
||||
|
||||
let servers_table = match doc
|
||||
.get("mcp")
|
||||
.and_then(|mcp| mcp.get("servers"))
|
||||
.and_then(|s| s.as_table())
|
||||
{
|
||||
Some(t) => t,
|
||||
None => return Ok(Vec::new()),
|
||||
};
|
||||
|
||||
let mut servers = Vec::new();
|
||||
for (name, config) in servers_table {
|
||||
if let Some(server) = parse_toml_server_entry(name, config) {
|
||||
servers.push(server);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(servers)
|
||||
}
|
||||
|
||||
/// Parse a single server entry from the TOML `[mcp.servers.*]` table.
|
||||
fn parse_toml_server_entry(name: &str, config: &toml::Value) -> Option<DetectedServer> {
|
||||
let table = config.as_table()?;
|
||||
|
||||
let transport_type = table
|
||||
.get("transport")
|
||||
.or_else(|| table.get("type"))
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("stdio");
|
||||
|
||||
let transport = match transport_type {
|
||||
"stdio" => {
|
||||
let command = table.get("command")?.as_str()?.to_owned();
|
||||
let args = table
|
||||
.get("args")
|
||||
.and_then(|v| v.as_array())
|
||||
.map(|arr| arr.iter().filter_map(|v| v.as_str().map(String::from)).collect())
|
||||
.unwrap_or_default();
|
||||
let env = table
|
||||
.get("env")
|
||||
.and_then(|v| v.as_table())
|
||||
.map(|t| {
|
||||
t.iter()
|
||||
.filter_map(|(k, v)| v.as_str().map(|s| (k.clone(), s.to_owned())))
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
McpServerTransport::Stdio { command, args, env }
|
||||
}
|
||||
"sse" => {
|
||||
let url = table.get("url")?.as_str()?.to_owned();
|
||||
let headers = parse_toml_headers(table);
|
||||
McpServerTransport::Sse { url, headers }
|
||||
}
|
||||
"http" | "streamable_http" => {
|
||||
let url = table.get("url")?.as_str()?.to_owned();
|
||||
let headers = parse_toml_headers(table);
|
||||
McpServerTransport::Http { url, headers }
|
||||
}
|
||||
_ => return None,
|
||||
};
|
||||
|
||||
Some(DetectedServer {
|
||||
name: name.to_owned(),
|
||||
transport,
|
||||
importable: true,
|
||||
import_skip_reason: None,
|
||||
})
|
||||
}
|
||||
|
||||
/// Extract headers from a TOML table's `headers` field.
|
||||
fn parse_toml_headers(table: &toml::map::Map<String, toml::Value>) -> HashMap<String, String> {
|
||||
table
|
||||
.get("headers")
|
||||
.and_then(|v| v.as_table())
|
||||
.map(|t| {
|
||||
t.iter()
|
||||
.filter_map(|(k, v)| v.as_str().map(|s| (k.clone(), s.to_owned())))
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
/// Convert a `McpServerTransport` to a TOML value for writing to config.
|
||||
fn transport_to_toml(transport: &McpServerTransport) -> toml::Value {
|
||||
let mut table = toml::map::Map::new();
|
||||
|
||||
match transport {
|
||||
McpServerTransport::Stdio { command, args, env } => {
|
||||
table.insert("transport".into(), toml::Value::String("stdio".into()));
|
||||
table.insert("command".into(), toml::Value::String(command.clone()));
|
||||
if !args.is_empty() {
|
||||
table.insert(
|
||||
"args".into(),
|
||||
toml::Value::Array(args.iter().map(|a| toml::Value::String(a.clone())).collect()),
|
||||
);
|
||||
}
|
||||
if !env.is_empty() {
|
||||
let env_table: toml::map::Map<String, toml::Value> = env
|
||||
.iter()
|
||||
.map(|(k, v)| (k.clone(), toml::Value::String(v.clone())))
|
||||
.collect();
|
||||
table.insert("env".into(), toml::Value::Table(env_table));
|
||||
}
|
||||
}
|
||||
McpServerTransport::Sse { url, headers } => {
|
||||
table.insert("transport".into(), toml::Value::String("sse".into()));
|
||||
table.insert("url".into(), toml::Value::String(url.clone()));
|
||||
insert_toml_headers(&mut table, headers);
|
||||
}
|
||||
McpServerTransport::Http { url, headers } => {
|
||||
table.insert("transport".into(), toml::Value::String("http".into()));
|
||||
table.insert("url".into(), toml::Value::String(url.clone()));
|
||||
insert_toml_headers(&mut table, headers);
|
||||
}
|
||||
}
|
||||
|
||||
toml::Value::Table(table)
|
||||
}
|
||||
|
||||
/// Insert headers into a TOML table if non-empty.
|
||||
fn insert_toml_headers(table: &mut toml::map::Map<String, toml::Value>, headers: &HashMap<String, String>) {
|
||||
if !headers.is_empty() {
|
||||
let headers_table: toml::map::Map<String, toml::Value> = headers
|
||||
.iter()
|
||||
.map(|(k, v)| (k.clone(), toml::Value::String(v.clone())))
|
||||
.collect();
|
||||
table.insert("headers".into(), toml::Value::Table(headers_table));
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn source_is_nomi() {
|
||||
assert_eq!(NomiAdapter.source(), McpSource::Nomi);
|
||||
}
|
||||
|
||||
// -- parse_toml_servers ---------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn parse_empty_config() {
|
||||
let servers = parse_toml_servers("").unwrap();
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_no_mcp_section() {
|
||||
let toml = r#"
|
||||
[some_other]
|
||||
key = "value"
|
||||
"#;
|
||||
let servers = parse_toml_servers(toml).unwrap();
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_empty_servers() {
|
||||
let toml = r#"
|
||||
[mcp]
|
||||
[mcp.servers]
|
||||
"#;
|
||||
let servers = parse_toml_servers(toml).unwrap();
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_stdio_server() {
|
||||
let toml = r#"
|
||||
[mcp.servers.test-mcp]
|
||||
type = "stdio"
|
||||
command = "npx"
|
||||
args = ["-y", "@test/server"]
|
||||
|
||||
[mcp.servers.test-mcp.env]
|
||||
KEY = "VALUE"
|
||||
NODE_ENV = "production"
|
||||
"#;
|
||||
let servers = parse_toml_servers(toml).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert_eq!(servers[0].name, "test-mcp");
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Stdio { command, args, env } => {
|
||||
assert_eq!(command, "npx");
|
||||
assert_eq!(args, &["-y", "@test/server"]);
|
||||
assert_eq!(env.get("KEY").unwrap(), "VALUE");
|
||||
assert_eq!(env.get("NODE_ENV").unwrap(), "production");
|
||||
}
|
||||
_ => panic!("expected Stdio"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_stdio_server_with_transport_key() {
|
||||
let toml = r#"
|
||||
[mcp.servers.test-mcp]
|
||||
transport = "stdio"
|
||||
command = "npx"
|
||||
args = ["-y", "@test/server"]
|
||||
|
||||
[mcp.servers.test-mcp.env]
|
||||
KEY = "VALUE"
|
||||
"#;
|
||||
let servers = parse_toml_servers(toml).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Stdio { command, args, env } => {
|
||||
assert_eq!(command, "npx");
|
||||
assert_eq!(args, &["-y", "@test/server"]);
|
||||
assert_eq!(env.get("KEY").unwrap(), "VALUE");
|
||||
}
|
||||
_ => panic!("expected Stdio"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_http_server() {
|
||||
let toml = r#"
|
||||
[mcp.servers.remote]
|
||||
type = "http"
|
||||
url = "https://example.com/mcp"
|
||||
|
||||
[mcp.servers.remote.headers]
|
||||
Authorization = "Bearer tok"
|
||||
"#;
|
||||
let servers = parse_toml_servers(toml).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert_eq!(servers[0].name, "remote");
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Http { url, headers } => {
|
||||
assert_eq!(url, "https://example.com/mcp");
|
||||
assert_eq!(headers.get("Authorization").unwrap(), "Bearer tok");
|
||||
}
|
||||
_ => panic!("expected Http"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_sse_server() {
|
||||
let toml = r#"
|
||||
[mcp.servers.sse-srv]
|
||||
type = "sse"
|
||||
url = "https://example.com/sse"
|
||||
"#;
|
||||
let servers = parse_toml_servers(toml).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Sse { url, .. } => {
|
||||
assert_eq!(url, "https://example.com/sse");
|
||||
}
|
||||
_ => panic!("expected Sse"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_streamable_http_becomes_http() {
|
||||
let toml = r#"
|
||||
[mcp.servers.sh]
|
||||
type = "streamable_http"
|
||||
url = "https://example.com/api"
|
||||
"#;
|
||||
let servers = parse_toml_servers(toml).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert!(matches!(servers[0].transport, McpServerTransport::Http { .. }));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_unknown_transport_skipped() {
|
||||
let toml = r#"
|
||||
[mcp.servers.ws]
|
||||
type = "websocket"
|
||||
url = "ws://localhost"
|
||||
"#;
|
||||
let servers = parse_toml_servers(toml).unwrap();
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_stdio_missing_command_skipped() {
|
||||
let toml = r#"
|
||||
[mcp.servers.bad]
|
||||
type = "stdio"
|
||||
args = []
|
||||
"#;
|
||||
let servers = parse_toml_servers(toml).unwrap();
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_multiple_servers() {
|
||||
let toml = r#"
|
||||
[mcp.servers.srv-a]
|
||||
type = "stdio"
|
||||
command = "node"
|
||||
|
||||
[mcp.servers.srv-b]
|
||||
type = "http"
|
||||
url = "https://b.com/mcp"
|
||||
|
||||
[mcp.servers.srv-c]
|
||||
type = "sse"
|
||||
url = "https://c.com/sse"
|
||||
"#;
|
||||
let servers = parse_toml_servers(toml).unwrap();
|
||||
assert_eq!(servers.len(), 3);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_default_type_is_stdio() {
|
||||
let toml = r#"
|
||||
[mcp.servers.no-type]
|
||||
command = "node"
|
||||
args = ["srv.js"]
|
||||
"#;
|
||||
let servers = parse_toml_servers(toml).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert!(matches!(servers[0].transport, McpServerTransport::Stdio { .. }));
|
||||
}
|
||||
|
||||
// -- transport_to_toml roundtrip ------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn stdio_to_toml_roundtrip() {
|
||||
let transport = McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec!["-y".into(), "@test/srv".into()],
|
||||
env: HashMap::from([("K".into(), "V".into())]),
|
||||
};
|
||||
let toml_val = transport_to_toml(&transport);
|
||||
let table = toml_val.as_table().unwrap();
|
||||
assert_eq!(table.get("transport").unwrap().as_str().unwrap(), "stdio");
|
||||
assert!(table.get("type").is_none());
|
||||
assert_eq!(
|
||||
table
|
||||
.get("env")
|
||||
.and_then(|v| v.as_table())
|
||||
.and_then(|env| env.get("K"))
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap(),
|
||||
"V"
|
||||
);
|
||||
let server = parse_toml_server_entry("test", &toml_val).unwrap();
|
||||
assert_eq!(server.transport, transport);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn http_to_toml_roundtrip() {
|
||||
let transport = McpServerTransport::Http {
|
||||
url: "https://example.com/mcp".into(),
|
||||
headers: HashMap::from([("Authorization".into(), "Bearer tok".into())]),
|
||||
};
|
||||
let toml_val = transport_to_toml(&transport);
|
||||
let server = parse_toml_server_entry("test", &toml_val).unwrap();
|
||||
assert_eq!(server.transport, transport);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sse_to_toml_roundtrip() {
|
||||
let transport = McpServerTransport::Sse {
|
||||
url: "https://example.com/sse".into(),
|
||||
headers: HashMap::new(),
|
||||
};
|
||||
let toml_val = transport_to_toml(&transport);
|
||||
let server = parse_toml_server_entry("test", &toml_val).unwrap();
|
||||
assert_eq!(server.transport, transport);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stdio_to_toml_omits_empty_args_and_env() {
|
||||
let transport = McpServerTransport::Stdio {
|
||||
command: "node".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
};
|
||||
let toml_val = transport_to_toml(&transport);
|
||||
let table = toml_val.as_table().unwrap();
|
||||
assert!(table.get("args").is_none());
|
||||
assert!(table.get("env").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn http_to_toml_omits_empty_headers() {
|
||||
let transport = McpServerTransport::Http {
|
||||
url: "https://x.com".into(),
|
||||
headers: HashMap::new(),
|
||||
};
|
||||
let toml_val = transport_to_toml(&transport);
|
||||
let table = toml_val.as_table().unwrap();
|
||||
assert!(table.get("headers").is_none());
|
||||
}
|
||||
|
||||
// -- invalid TOML ---------------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn parse_invalid_toml_fails() {
|
||||
let result = parse_toml_servers("not valid toml [[[");
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn trait_is_object_safe() {
|
||||
let adapter: Box<dyn McpAgentAdapter> = Box::new(NomiAdapter);
|
||||
assert_eq!(adapter.source(), McpSource::Nomi);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,240 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use nomifun_common::McpSource;
|
||||
use nomifun_db::IMcpServerRepository;
|
||||
|
||||
use crate::adapter::{DetectedServer, McpAgentAdapter};
|
||||
use crate::error::McpError;
|
||||
use crate::types::{McpServer, McpServerTransport};
|
||||
|
||||
/// MCP Agent adapter for Nomi itself.
|
||||
///
|
||||
/// Unlike CLI-based adapters, this adapter reads/writes directly to the
|
||||
/// local database. It is always "installed" since Nomi is the host
|
||||
/// application.
|
||||
///
|
||||
/// # Behavior
|
||||
///
|
||||
/// - `is_installed()` → always `true`
|
||||
/// - `detect_existing()` → reads all MCP servers from the DB
|
||||
/// - `install_server()` → no-op (DB writes are handled by `McpConfigService`)
|
||||
/// - `remove_server()` → no-op (configuration is managed via the frontend)
|
||||
pub struct NomifunAdapter {
|
||||
repo: Arc<dyn IMcpServerRepository>,
|
||||
}
|
||||
|
||||
impl NomifunAdapter {
|
||||
pub fn new(repo: Arc<dyn IMcpServerRepository>) -> Self {
|
||||
Self { repo }
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl McpAgentAdapter for NomifunAdapter {
|
||||
fn source(&self) -> McpSource {
|
||||
McpSource::Nomifun
|
||||
}
|
||||
|
||||
async fn is_installed(&self) -> Result<bool, McpError> {
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn detect_existing(&self) -> Result<Vec<DetectedServer>, McpError> {
|
||||
let rows = self.repo.list().await?;
|
||||
|
||||
let mut servers = Vec::new();
|
||||
for row in rows {
|
||||
let server = McpServer::from_row(row)?;
|
||||
servers.push(DetectedServer {
|
||||
name: server.name,
|
||||
transport: server.transport,
|
||||
importable: true,
|
||||
import_skip_reason: None,
|
||||
});
|
||||
}
|
||||
|
||||
Ok(servers)
|
||||
}
|
||||
|
||||
async fn install_server(&self, _name: &str, _transport: &McpServerTransport) -> Result<(), McpError> {
|
||||
// No-op: DB writes are handled by McpConfigService.
|
||||
// The sync service calls install_server on all adapters, but for
|
||||
// Nomi the server is already in the DB.
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn remove_server(&self, _name: &str) -> Result<(), McpError> {
|
||||
// No-op: configuration is managed via the frontend/REST API.
|
||||
// Removing from the DB is done through McpConfigService.delete_server().
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::collections::HashMap;
|
||||
|
||||
use super::*;
|
||||
use crate::types::McpServerTransport;
|
||||
use nomifun_db::models::McpServerRow;
|
||||
|
||||
/// In-memory mock repository for testing.
|
||||
struct MockRepo {
|
||||
servers: Vec<McpServerRow>,
|
||||
}
|
||||
|
||||
impl MockRepo {
|
||||
fn new(servers: Vec<McpServerRow>) -> Self {
|
||||
Self { servers }
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl IMcpServerRepository for MockRepo {
|
||||
async fn list(&self) -> Result<Vec<McpServerRow>, nomifun_db::DbError> {
|
||||
Ok(self.servers.clone())
|
||||
}
|
||||
|
||||
async fn find_by_id(&self, id: i64) -> Result<Option<McpServerRow>, nomifun_db::DbError> {
|
||||
Ok(self.servers.iter().find(|s| s.id == id).cloned())
|
||||
}
|
||||
|
||||
async fn find_by_name(&self, name: &str) -> Result<Option<McpServerRow>, nomifun_db::DbError> {
|
||||
Ok(self.servers.iter().find(|s| s.name == name).cloned())
|
||||
}
|
||||
|
||||
async fn create(
|
||||
&self,
|
||||
_params: nomifun_db::CreateMcpServerParams<'_>,
|
||||
) -> Result<McpServerRow, nomifun_db::DbError> {
|
||||
unimplemented!("not needed for adapter tests")
|
||||
}
|
||||
|
||||
async fn update(
|
||||
&self,
|
||||
_id: i64,
|
||||
_params: nomifun_db::UpdateMcpServerParams<'_>,
|
||||
) -> Result<McpServerRow, nomifun_db::DbError> {
|
||||
unimplemented!("not needed for adapter tests")
|
||||
}
|
||||
|
||||
async fn delete(&self, _id: i64) -> Result<(), nomifun_db::DbError> {
|
||||
unimplemented!("not needed for adapter tests")
|
||||
}
|
||||
|
||||
async fn batch_upsert(
|
||||
&self,
|
||||
_servers: &[nomifun_db::CreateMcpServerParams<'_>],
|
||||
) -> Result<Vec<McpServerRow>, nomifun_db::DbError> {
|
||||
unimplemented!("not needed for adapter tests")
|
||||
}
|
||||
|
||||
async fn update_status(
|
||||
&self,
|
||||
_id: i64,
|
||||
_status: &str,
|
||||
_last_connected: Option<nomifun_common::TimestampMs>,
|
||||
) -> Result<(), nomifun_db::DbError> {
|
||||
unimplemented!("not needed for adapter tests")
|
||||
}
|
||||
|
||||
async fn update_tools(&self, _id: i64, _tools: Option<&str>) -> Result<(), nomifun_db::DbError> {
|
||||
unimplemented!("not needed for adapter tests")
|
||||
}
|
||||
}
|
||||
|
||||
fn make_row(name: &str, transport_type: &str, transport_config: &str) -> McpServerRow {
|
||||
McpServerRow {
|
||||
// Host-local integer PK; never compared in adapter tests (detection keys on name).
|
||||
id: name.bytes().map(i64::from).sum::<i64>().max(1),
|
||||
name: name.to_owned(),
|
||||
description: None,
|
||||
enabled: true,
|
||||
transport_type: transport_type.into(),
|
||||
transport_config: transport_config.into(),
|
||||
tools: None,
|
||||
last_test_status: "disconnected".into(),
|
||||
last_connected: None,
|
||||
original_json: None,
|
||||
builtin: false,
|
||||
deleted_at: None,
|
||||
created_at: 1000,
|
||||
updated_at: 1000,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn source_is_nomifun() {
|
||||
let repo = Arc::new(MockRepo::new(vec![]));
|
||||
let adapter = NomifunAdapter::new(repo);
|
||||
assert_eq!(adapter.source(), McpSource::Nomifun);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn is_always_installed() {
|
||||
let repo = Arc::new(MockRepo::new(vec![]));
|
||||
let adapter = NomifunAdapter::new(repo);
|
||||
assert!(adapter.is_installed().await.unwrap());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn detect_existing_returns_db_servers() {
|
||||
let rows = vec![
|
||||
make_row("srv-a", "stdio", r#"{"command":"npx","args":[]}"#),
|
||||
make_row("srv-b", "http", r#"{"url":"https://b.com/mcp","headers":{}}"#),
|
||||
];
|
||||
let repo = Arc::new(MockRepo::new(rows));
|
||||
let adapter = NomifunAdapter::new(repo);
|
||||
|
||||
let servers = adapter.detect_existing().await.unwrap();
|
||||
assert_eq!(servers.len(), 2);
|
||||
assert_eq!(servers[0].name, "srv-a");
|
||||
assert_eq!(servers[1].name, "srv-b");
|
||||
assert!(matches!(servers[0].transport, McpServerTransport::Stdio { .. }));
|
||||
assert!(matches!(servers[1].transport, McpServerTransport::Http { .. }));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn detect_existing_empty_db() {
|
||||
let repo = Arc::new(MockRepo::new(vec![]));
|
||||
let adapter = NomifunAdapter::new(repo);
|
||||
|
||||
let servers = adapter.detect_existing().await.unwrap();
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn install_server_is_noop() {
|
||||
let repo = Arc::new(MockRepo::new(vec![]));
|
||||
let adapter = NomifunAdapter::new(repo);
|
||||
|
||||
let transport = McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
};
|
||||
// Should succeed without side effects
|
||||
adapter.install_server("test", &transport).await.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn remove_server_is_noop() {
|
||||
let repo = Arc::new(MockRepo::new(vec![]));
|
||||
let adapter = NomifunAdapter::new(repo);
|
||||
|
||||
// Should succeed without side effects
|
||||
adapter.remove_server("test").await.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn trait_is_object_safe() {
|
||||
let repo = Arc::new(MockRepo::new(vec![]));
|
||||
let adapter: Arc<dyn McpAgentAdapter> = Arc::new(NomifunAdapter::new(repo));
|
||||
assert_eq!(adapter.source(), McpSource::Nomifun);
|
||||
assert!(adapter.is_installed().await.unwrap());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,637 @@
|
||||
use std::collections::HashMap;
|
||||
use std::path::PathBuf;
|
||||
|
||||
use nomifun_common::McpSource;
|
||||
|
||||
use crate::adapter::{DetectedServer, McpAgentAdapter};
|
||||
use crate::error::McpError;
|
||||
use crate::types::McpServerTransport;
|
||||
|
||||
/// MCP Agent adapter for Opencode.
|
||||
///
|
||||
/// Opencode stores configuration in `~/.config/opencode/opencode.json`.
|
||||
/// The `mcp` field is a map of server names to transport configs.
|
||||
///
|
||||
/// # Config Format (JSONC)
|
||||
///
|
||||
/// ```jsonc
|
||||
/// {
|
||||
/// // other opencode config...
|
||||
/// "mcp": {
|
||||
/// "server-name": {
|
||||
/// "type": "stdio",
|
||||
/// "command": "npx",
|
||||
/// "args": ["-y", "@test/server"],
|
||||
/// "env": { "KEY": "VALUE" }
|
||||
/// },
|
||||
/// "remote-server": {
|
||||
/// "type": "http",
|
||||
/// "url": "https://example.com/mcp",
|
||||
/// "headers": { "Authorization": "Bearer xxx" }
|
||||
/// }
|
||||
/// }
|
||||
/// }
|
||||
/// ```
|
||||
///
|
||||
/// Opencode config files may contain JSON comments (JSONC), so we
|
||||
/// strip comments before parsing and preserve the original structure
|
||||
/// when writing back.
|
||||
pub struct OpencodeAdapter;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl McpAgentAdapter for OpencodeAdapter {
|
||||
fn source(&self) -> McpSource {
|
||||
McpSource::OpenCode
|
||||
}
|
||||
|
||||
async fn is_installed(&self) -> Result<bool, McpError> {
|
||||
Ok(config_dir().is_some_and(|d| d.exists()))
|
||||
}
|
||||
|
||||
async fn detect_existing(&self) -> Result<Vec<DetectedServer>, McpError> {
|
||||
let path = config_file_path().ok_or_else(|| McpError::AgentNotInstalled("opencode".into()))?;
|
||||
|
||||
if !path.exists() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let content = tokio::fs::read_to_string(&path)
|
||||
.await
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("failed to read {}: {e}", path.display())))?;
|
||||
|
||||
let root = parse_jsonc(&content)?;
|
||||
parse_mcp_field(&root)
|
||||
}
|
||||
|
||||
async fn install_server(&self, name: &str, transport: &McpServerTransport) -> Result<(), McpError> {
|
||||
let path = config_file_path().ok_or_else(|| McpError::AgentNotInstalled("opencode".into()))?;
|
||||
|
||||
let mut root = if path.exists() {
|
||||
let content = tokio::fs::read_to_string(&path)
|
||||
.await
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("failed to read {}: {e}", path.display())))?;
|
||||
parse_jsonc(&content)?
|
||||
} else {
|
||||
// Ensure directory exists
|
||||
if let Some(parent) = path.parent() {
|
||||
tokio::fs::create_dir_all(parent)
|
||||
.await
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("failed to create dir: {e}")))?;
|
||||
}
|
||||
serde_json::json!({})
|
||||
};
|
||||
|
||||
let mcp = root
|
||||
.as_object_mut()
|
||||
.ok_or_else(|| McpError::AgentOperationFailed("config root is not an object".into()))?
|
||||
.entry("mcp")
|
||||
.or_insert_with(|| serde_json::json!({}));
|
||||
|
||||
let mcp_obj = mcp
|
||||
.as_object_mut()
|
||||
.ok_or_else(|| McpError::AgentOperationFailed("mcp field is not an object".into()))?;
|
||||
|
||||
mcp_obj.insert(name.to_owned(), transport_to_json(transport));
|
||||
|
||||
let output = serde_json::to_string_pretty(&root)
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("failed to serialize config: {e}")))?;
|
||||
|
||||
tokio::fs::write(&path, output)
|
||||
.await
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("failed to write {}: {e}", path.display())))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn remove_server(&self, name: &str) -> Result<(), McpError> {
|
||||
let path = config_file_path().ok_or_else(|| McpError::AgentNotInstalled("opencode".into()))?;
|
||||
|
||||
if !path.exists() {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let content = tokio::fs::read_to_string(&path)
|
||||
.await
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("failed to read {}: {e}", path.display())))?;
|
||||
|
||||
let mut root = parse_jsonc(&content)?;
|
||||
|
||||
let removed = root
|
||||
.as_object_mut()
|
||||
.and_then(|obj| obj.get_mut("mcp"))
|
||||
.and_then(|mcp| mcp.as_object_mut())
|
||||
.map(|mcp_obj| mcp_obj.remove(name).is_some())
|
||||
.unwrap_or(false);
|
||||
|
||||
if !removed {
|
||||
// Idempotent: not found is fine
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let output = serde_json::to_string_pretty(&root)
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("failed to serialize config: {e}")))?;
|
||||
|
||||
tokio::fs::write(&path, output)
|
||||
.await
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("failed to write {}: {e}", path.display())))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Returns `~/.config/opencode/` if HOME is available.
|
||||
fn config_dir() -> Option<PathBuf> {
|
||||
dirs::config_dir().map(|d| d.join("opencode"))
|
||||
}
|
||||
|
||||
/// Returns `~/.config/opencode/opencode.json` if HOME is available.
|
||||
fn config_file_path() -> Option<PathBuf> {
|
||||
config_dir().map(|d| d.join("opencode.json"))
|
||||
}
|
||||
|
||||
/// Strip single-line (`//`) and multi-line (`/* ... */`) JSON comments.
|
||||
///
|
||||
/// Preserves string contents (comments inside strings are left alone).
|
||||
fn strip_json_comments(input: &str) -> String {
|
||||
let mut result = String::with_capacity(input.len());
|
||||
let bytes = input.as_bytes();
|
||||
let len = bytes.len();
|
||||
let mut i = 0;
|
||||
|
||||
while i < len {
|
||||
// Check for string literal
|
||||
if bytes[i] == b'"' {
|
||||
result.push('"');
|
||||
i += 1;
|
||||
// Consume until closing quote, respecting escapes
|
||||
while i < len {
|
||||
if bytes[i] == b'\\' && i + 1 < len {
|
||||
result.push(bytes[i] as char);
|
||||
result.push(bytes[i + 1] as char);
|
||||
i += 2;
|
||||
} else if bytes[i] == b'"' {
|
||||
result.push('"');
|
||||
i += 1;
|
||||
break;
|
||||
} else {
|
||||
result.push(bytes[i] as char);
|
||||
i += 1;
|
||||
}
|
||||
}
|
||||
} else if bytes[i] == b'/' && i + 1 < len {
|
||||
if bytes[i + 1] == b'/' {
|
||||
// Single-line comment: skip until newline
|
||||
i += 2;
|
||||
while i < len && bytes[i] != b'\n' {
|
||||
i += 1;
|
||||
}
|
||||
} else if bytes[i + 1] == b'*' {
|
||||
// Multi-line comment: skip until */
|
||||
i += 2;
|
||||
while i + 1 < len {
|
||||
if bytes[i] == b'*' && bytes[i + 1] == b'/' {
|
||||
i += 2;
|
||||
break;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
// Handle unterminated block comment
|
||||
if i >= len {
|
||||
break;
|
||||
}
|
||||
} else {
|
||||
result.push(bytes[i] as char);
|
||||
i += 1;
|
||||
}
|
||||
} else {
|
||||
result.push(bytes[i] as char);
|
||||
i += 1;
|
||||
}
|
||||
}
|
||||
|
||||
result
|
||||
}
|
||||
|
||||
/// Parse JSONC (JSON with comments) into a `serde_json::Value`.
|
||||
fn parse_jsonc(input: &str) -> Result<serde_json::Value, McpError> {
|
||||
let stripped = strip_json_comments(input);
|
||||
serde_json::from_str(&stripped).map_err(McpError::from)
|
||||
}
|
||||
|
||||
/// Extract MCP servers from the parsed config root.
|
||||
fn parse_mcp_field(root: &serde_json::Value) -> Result<Vec<DetectedServer>, McpError> {
|
||||
let mcp = match root.get("mcp") {
|
||||
Some(v) => v,
|
||||
None => return Ok(Vec::new()),
|
||||
};
|
||||
|
||||
let mcp_obj = mcp
|
||||
.as_object()
|
||||
.ok_or_else(|| McpError::AgentOperationFailed("mcp field is not an object".into()))?;
|
||||
|
||||
let mut servers = Vec::new();
|
||||
|
||||
for (name, config) in mcp_obj {
|
||||
if let Some(server) = parse_server_entry(name, config) {
|
||||
servers.push(server);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(servers)
|
||||
}
|
||||
|
||||
/// Parse a single server entry from the `mcp` object.
|
||||
fn parse_server_entry(name: &str, config: &serde_json::Value) -> Option<DetectedServer> {
|
||||
let transport_type = config.get("type").and_then(|v| v.as_str()).unwrap_or("stdio");
|
||||
|
||||
let transport = match transport_type {
|
||||
"stdio" => {
|
||||
let command = config.get("command")?.as_str()?.to_owned();
|
||||
let args = config
|
||||
.get("args")
|
||||
.and_then(|v| v.as_array())
|
||||
.map(|arr| arr.iter().filter_map(|v| v.as_str().map(String::from)).collect())
|
||||
.unwrap_or_default();
|
||||
let env = config
|
||||
.get("env")
|
||||
.and_then(|v| v.as_object())
|
||||
.map(|obj| {
|
||||
obj.iter()
|
||||
.filter_map(|(k, v)| v.as_str().map(|s| (k.clone(), s.to_owned())))
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
McpServerTransport::Stdio { command, args, env }
|
||||
}
|
||||
"sse" => {
|
||||
let url = config.get("url")?.as_str()?.to_owned();
|
||||
let headers = parse_headers(config);
|
||||
McpServerTransport::Sse { url, headers }
|
||||
}
|
||||
"http" | "streamable_http" => {
|
||||
let url = config.get("url")?.as_str()?.to_owned();
|
||||
let headers = parse_headers(config);
|
||||
McpServerTransport::Http { url, headers }
|
||||
}
|
||||
_ => return None,
|
||||
};
|
||||
|
||||
Some(DetectedServer {
|
||||
name: name.to_owned(),
|
||||
transport,
|
||||
importable: true,
|
||||
import_skip_reason: None,
|
||||
})
|
||||
}
|
||||
|
||||
/// Extract headers from a config object's `headers` field.
|
||||
fn parse_headers(config: &serde_json::Value) -> HashMap<String, String> {
|
||||
config
|
||||
.get("headers")
|
||||
.and_then(|v| v.as_object())
|
||||
.map(|obj| {
|
||||
obj.iter()
|
||||
.filter_map(|(k, v)| v.as_str().map(|s| (k.clone(), s.to_owned())))
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
/// Convert a `McpServerTransport` to a JSON value for writing to config.
|
||||
fn transport_to_json(transport: &McpServerTransport) -> serde_json::Value {
|
||||
match transport {
|
||||
McpServerTransport::Stdio { command, args, env } => {
|
||||
let mut obj = serde_json::json!({
|
||||
"type": "stdio",
|
||||
"command": command,
|
||||
"args": args,
|
||||
});
|
||||
if !env.is_empty() {
|
||||
obj["env"] = serde_json::json!(env);
|
||||
}
|
||||
obj
|
||||
}
|
||||
McpServerTransport::Sse { url, headers } => {
|
||||
let mut obj = serde_json::json!({
|
||||
"type": "sse",
|
||||
"url": url,
|
||||
});
|
||||
if !headers.is_empty() {
|
||||
obj["headers"] = serde_json::json!(headers);
|
||||
}
|
||||
obj
|
||||
}
|
||||
McpServerTransport::Http { url, headers } => {
|
||||
let mut obj = serde_json::json!({
|
||||
"type": "http",
|
||||
"url": url,
|
||||
});
|
||||
if !headers.is_empty() {
|
||||
obj["headers"] = serde_json::json!(headers);
|
||||
}
|
||||
obj
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn source_is_opencode() {
|
||||
assert_eq!(OpencodeAdapter.source(), McpSource::OpenCode);
|
||||
}
|
||||
|
||||
// -- strip_json_comments --------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn strip_single_line_comments() {
|
||||
let input = r#"{
|
||||
// This is a comment
|
||||
"key": "value" // inline comment
|
||||
}"#;
|
||||
let stripped = strip_json_comments(input);
|
||||
let parsed: serde_json::Value = serde_json::from_str(&stripped).unwrap();
|
||||
assert_eq!(parsed["key"], "value");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn strip_multi_line_comments() {
|
||||
let input = r#"{
|
||||
/* multi-line
|
||||
comment */
|
||||
"key": "value"
|
||||
}"#;
|
||||
let stripped = strip_json_comments(input);
|
||||
let parsed: serde_json::Value = serde_json::from_str(&stripped).unwrap();
|
||||
assert_eq!(parsed["key"], "value");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn preserve_comments_inside_strings() {
|
||||
let input = r#"{
|
||||
"key": "value with // comment inside",
|
||||
"key2": "value with /* block */ inside"
|
||||
}"#;
|
||||
let stripped = strip_json_comments(input);
|
||||
let parsed: serde_json::Value = serde_json::from_str(&stripped).unwrap();
|
||||
assert_eq!(parsed["key"], "value with // comment inside");
|
||||
assert_eq!(parsed["key2"], "value with /* block */ inside");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn strip_comments_preserves_escaped_quotes() {
|
||||
let input = r#"{"key": "val\"ue // not a comment"}"#;
|
||||
let stripped = strip_json_comments(input);
|
||||
let parsed: serde_json::Value = serde_json::from_str(&stripped).unwrap();
|
||||
assert_eq!(parsed["key"], "val\"ue // not a comment");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn strip_no_comments() {
|
||||
let input = r#"{"key": "value"}"#;
|
||||
assert_eq!(strip_json_comments(input), input);
|
||||
}
|
||||
|
||||
// -- parse_mcp_field ------------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn parse_empty_mcp() {
|
||||
let root = serde_json::json!({ "mcp": {} });
|
||||
let servers = parse_mcp_field(&root).unwrap();
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_no_mcp_field() {
|
||||
let root = serde_json::json!({ "other": "stuff" });
|
||||
let servers = parse_mcp_field(&root).unwrap();
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_stdio_server() {
|
||||
let root = serde_json::json!({
|
||||
"mcp": {
|
||||
"test-mcp": {
|
||||
"type": "stdio",
|
||||
"command": "npx",
|
||||
"args": ["-y", "@test/server"],
|
||||
"env": { "KEY": "VALUE" }
|
||||
}
|
||||
}
|
||||
});
|
||||
let servers = parse_mcp_field(&root).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert_eq!(servers[0].name, "test-mcp");
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Stdio { command, args, env } => {
|
||||
assert_eq!(command, "npx");
|
||||
assert_eq!(args, &["-y", "@test/server"]);
|
||||
assert_eq!(env.get("KEY").unwrap(), "VALUE");
|
||||
}
|
||||
_ => panic!("expected Stdio"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_http_server() {
|
||||
let root = serde_json::json!({
|
||||
"mcp": {
|
||||
"remote": {
|
||||
"type": "http",
|
||||
"url": "https://example.com/mcp",
|
||||
"headers": { "Authorization": "Bearer tok" }
|
||||
}
|
||||
}
|
||||
});
|
||||
let servers = parse_mcp_field(&root).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Http { url, headers } => {
|
||||
assert_eq!(url, "https://example.com/mcp");
|
||||
assert_eq!(headers.get("Authorization").unwrap(), "Bearer tok");
|
||||
}
|
||||
_ => panic!("expected Http"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_sse_server() {
|
||||
let root = serde_json::json!({
|
||||
"mcp": {
|
||||
"sse-srv": {
|
||||
"type": "sse",
|
||||
"url": "https://example.com/sse"
|
||||
}
|
||||
}
|
||||
});
|
||||
let servers = parse_mcp_field(&root).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Sse { url, .. } => {
|
||||
assert_eq!(url, "https://example.com/sse");
|
||||
}
|
||||
_ => panic!("expected Sse"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_streamable_http_becomes_http() {
|
||||
let root = serde_json::json!({
|
||||
"mcp": {
|
||||
"sh": {
|
||||
"type": "streamable_http",
|
||||
"url": "https://example.com/api"
|
||||
}
|
||||
}
|
||||
});
|
||||
let servers = parse_mcp_field(&root).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert!(matches!(servers[0].transport, McpServerTransport::Http { .. }));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_unknown_transport_skipped() {
|
||||
let root = serde_json::json!({
|
||||
"mcp": {
|
||||
"ws": { "type": "websocket", "url": "ws://localhost" }
|
||||
}
|
||||
});
|
||||
let servers = parse_mcp_field(&root).unwrap();
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_stdio_missing_command_skipped() {
|
||||
let root = serde_json::json!({
|
||||
"mcp": {
|
||||
"bad": { "type": "stdio", "args": [] }
|
||||
}
|
||||
});
|
||||
let servers = parse_mcp_field(&root).unwrap();
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_multiple_servers() {
|
||||
let root = serde_json::json!({
|
||||
"mcp": {
|
||||
"srv-a": { "type": "stdio", "command": "node" },
|
||||
"srv-b": { "type": "http", "url": "https://b.com/mcp" }
|
||||
}
|
||||
});
|
||||
let servers = parse_mcp_field(&root).unwrap();
|
||||
assert_eq!(servers.len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_default_type_is_stdio() {
|
||||
let root = serde_json::json!({
|
||||
"mcp": {
|
||||
"no-type": { "command": "node", "args": ["srv.js"] }
|
||||
}
|
||||
});
|
||||
let servers = parse_mcp_field(&root).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert!(matches!(servers[0].transport, McpServerTransport::Stdio { .. }));
|
||||
}
|
||||
|
||||
// -- transport_to_json ----------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn stdio_to_json_roundtrip() {
|
||||
let transport = McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec!["-y".into(), "@test/srv".into()],
|
||||
env: HashMap::from([("K".into(), "V".into())]),
|
||||
};
|
||||
let json = transport_to_json(&transport);
|
||||
let server = parse_server_entry("test", &json).unwrap();
|
||||
assert_eq!(server.transport, transport);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn http_to_json_roundtrip() {
|
||||
let transport = McpServerTransport::Http {
|
||||
url: "https://example.com/mcp".into(),
|
||||
headers: HashMap::from([("Authorization".into(), "Bearer tok".into())]),
|
||||
};
|
||||
let json = transport_to_json(&transport);
|
||||
let server = parse_server_entry("test", &json).unwrap();
|
||||
assert_eq!(server.transport, transport);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sse_to_json_roundtrip() {
|
||||
let transport = McpServerTransport::Sse {
|
||||
url: "https://example.com/sse".into(),
|
||||
headers: HashMap::new(),
|
||||
};
|
||||
let json = transport_to_json(&transport);
|
||||
let server = parse_server_entry("test", &json).unwrap();
|
||||
assert_eq!(server.transport, transport);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stdio_to_json_omits_empty_env() {
|
||||
let transport = McpServerTransport::Stdio {
|
||||
command: "node".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
};
|
||||
let json = transport_to_json(&transport);
|
||||
assert!(json.get("env").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn http_to_json_omits_empty_headers() {
|
||||
let transport = McpServerTransport::Http {
|
||||
url: "https://x.com".into(),
|
||||
headers: HashMap::new(),
|
||||
};
|
||||
let json = transport_to_json(&transport);
|
||||
assert!(json.get("headers").is_none());
|
||||
}
|
||||
|
||||
// -- parse_jsonc ----------------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn parse_jsonc_with_comments() {
|
||||
let input = r#"{
|
||||
// comment
|
||||
"mcp": {
|
||||
/* block comment */
|
||||
"srv": {
|
||||
"type": "stdio",
|
||||
"command": "npx"
|
||||
}
|
||||
}
|
||||
}"#;
|
||||
let root = parse_jsonc(input).unwrap();
|
||||
let servers = parse_mcp_field(&root).unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
assert_eq!(servers[0].name, "srv");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_jsonc_invalid_json_fails() {
|
||||
let result = parse_jsonc("not json at all");
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn trait_is_object_safe() {
|
||||
let adapter: Box<dyn McpAgentAdapter> = Box::new(OpencodeAdapter);
|
||||
assert_eq!(adapter.source(), McpSource::OpenCode);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,198 @@
|
||||
use std::collections::HashMap;
|
||||
|
||||
use nomifun_common::McpSource;
|
||||
|
||||
use crate::adapter::{DetectedServer, McpAgentAdapter};
|
||||
use crate::error::McpError;
|
||||
use crate::types::McpServerTransport;
|
||||
|
||||
use super::cli_helpers::{
|
||||
DETECT_TIMEOUT, MUTATE_TIMEOUT, build_env_args, build_header_args, is_cli_installed, parse_standard_list_output,
|
||||
run_cli,
|
||||
};
|
||||
|
||||
const CLI_NAME: &str = "qwen";
|
||||
|
||||
/// Scopes tried when removing (user first, then project).
|
||||
const REMOVE_SCOPES: &[&str] = &["user", "project"];
|
||||
|
||||
/// MCP Agent adapter for Qwen CLI.
|
||||
///
|
||||
/// # CLI Commands
|
||||
///
|
||||
/// - **detect**: `qwen mcp list`
|
||||
/// - **install (stdio)**: `qwen mcp add <name> <command> [args...] [--env K=V]... -s user`
|
||||
/// - **install (http/sse)**: `qwen mcp add <name> <url> --transport <type> [--header K: V]... -s user`
|
||||
/// - **remove**: `qwen mcp remove <name> -s user` → `-s project` → file fallback
|
||||
///
|
||||
/// If CLI remove fails, falls back to editing `~/.qwen/client_config.json`.
|
||||
pub struct QwenAdapter;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl McpAgentAdapter for QwenAdapter {
|
||||
fn source(&self) -> McpSource {
|
||||
McpSource::Qwen
|
||||
}
|
||||
|
||||
async fn is_installed(&self) -> Result<bool, McpError> {
|
||||
is_cli_installed(CLI_NAME).await
|
||||
}
|
||||
|
||||
async fn detect_existing(&self) -> Result<Vec<DetectedServer>, McpError> {
|
||||
if !self.is_installed().await? {
|
||||
return Err(McpError::AgentNotInstalled(CLI_NAME.into()));
|
||||
}
|
||||
|
||||
let (stdout, _stderr) = run_cli(CLI_NAME, &["mcp", "list"], DETECT_TIMEOUT).await?;
|
||||
Ok(parse_standard_list_output(&stdout))
|
||||
}
|
||||
|
||||
async fn install_server(&self, name: &str, transport: &McpServerTransport) -> Result<(), McpError> {
|
||||
if !self.is_installed().await? {
|
||||
return Err(McpError::AgentNotInstalled(CLI_NAME.into()));
|
||||
}
|
||||
|
||||
match transport {
|
||||
McpServerTransport::Stdio { command, args, env } => {
|
||||
let mut cli_args = vec!["mcp".to_owned(), "add".to_owned(), name.to_owned(), command.clone()];
|
||||
cli_args.extend(args.iter().cloned());
|
||||
cli_args.extend(build_env_args(env, "--env"));
|
||||
cli_args.push("-s".to_owned());
|
||||
cli_args.push("user".to_owned());
|
||||
|
||||
let arg_refs: Vec<&str> = cli_args.iter().map(|s| s.as_str()).collect();
|
||||
run_cli(CLI_NAME, &arg_refs, MUTATE_TIMEOUT).await?;
|
||||
}
|
||||
McpServerTransport::Sse { url, headers } => {
|
||||
install_http_like(name, "sse", url, headers).await?;
|
||||
}
|
||||
McpServerTransport::Http { url, headers } => {
|
||||
install_http_like(name, "http", url, headers).await?;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn remove_server(&self, name: &str) -> Result<(), McpError> {
|
||||
if !self.is_installed().await? {
|
||||
return Err(McpError::AgentNotInstalled(CLI_NAME.into()));
|
||||
}
|
||||
|
||||
// Try CLI removal with each scope.
|
||||
for scope in REMOVE_SCOPES {
|
||||
let (stdout, _stderr) = run_cli(CLI_NAME, &["mcp", "remove", name, "-s", scope], MUTATE_TIMEOUT).await?;
|
||||
let lower = stdout.to_lowercase();
|
||||
if lower.contains("removed") {
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
|
||||
// Fallback: directly edit ~/.qwen/client_config.json
|
||||
remove_from_config_file(name).await
|
||||
}
|
||||
}
|
||||
|
||||
/// Install an HTTP-like (sse/http) server via `qwen mcp add`.
|
||||
async fn install_http_like(
|
||||
name: &str,
|
||||
transport_type: &str,
|
||||
url: &str,
|
||||
headers: &HashMap<String, String>,
|
||||
) -> Result<(), McpError> {
|
||||
let mut cli_args = vec![
|
||||
"mcp".to_owned(),
|
||||
"add".to_owned(),
|
||||
name.to_owned(),
|
||||
url.to_owned(),
|
||||
"--transport".to_owned(),
|
||||
transport_type.to_owned(),
|
||||
];
|
||||
cli_args.extend(build_header_args(headers, "--header"));
|
||||
cli_args.push("-s".to_owned());
|
||||
cli_args.push("user".to_owned());
|
||||
|
||||
let arg_refs: Vec<&str> = cli_args.iter().map(|s| s.as_str()).collect();
|
||||
run_cli(CLI_NAME, &arg_refs, MUTATE_TIMEOUT).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Fallback: remove server from `~/.qwen/client_config.json` directly.
|
||||
///
|
||||
/// Reads the file, deletes the key from `mcpServers`, writes back.
|
||||
/// Silently succeeds if the file doesn't exist or the key is absent.
|
||||
async fn remove_from_config_file(name: &str) -> Result<(), McpError> {
|
||||
let home = home_dir()?;
|
||||
let config_path = home.join(".qwen").join("client_config.json");
|
||||
|
||||
if !config_path.exists() {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let content = tokio::fs::read_to_string(&config_path)
|
||||
.await
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("read qwen config: {e}")))?;
|
||||
|
||||
let mut config: serde_json::Value = serde_json::from_str(&content).map_err(McpError::from)?;
|
||||
|
||||
let removed = config
|
||||
.get_mut("mcpServers")
|
||||
.and_then(|servers| servers.as_object_mut())
|
||||
.map(|servers| servers.remove(name).is_some())
|
||||
.unwrap_or(false);
|
||||
|
||||
if removed {
|
||||
let new_content = serde_json::to_string_pretty(&config).map_err(McpError::from)?;
|
||||
tokio::fs::write(&config_path, new_content)
|
||||
.await
|
||||
.map_err(|e| McpError::AgentOperationFailed(format!("write qwen config: {e}")))?;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Get the user's home directory.
|
||||
fn home_dir() -> Result<std::path::PathBuf, McpError> {
|
||||
dirs::home_dir().ok_or_else(|| McpError::AgentOperationFailed("cannot determine home directory".into()))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn source_is_qwen() {
|
||||
assert_eq!(QwenAdapter.source(), McpSource::Qwen);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_qwen_list_output() {
|
||||
let output = "\
|
||||
✓ my-server: npx -y @test/server (stdio) - Connected
|
||||
✗ broken: node bad.js (stdio) - Disconnected";
|
||||
|
||||
let servers = parse_standard_list_output(output);
|
||||
assert_eq!(servers.len(), 2);
|
||||
assert_eq!(servers[0].name, "my-server");
|
||||
assert_eq!(servers[1].name, "broken");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_qwen_http_server() {
|
||||
let output = "✓ remote: https://example.com/mcp (http) - Connected";
|
||||
let servers = parse_standard_list_output(output);
|
||||
assert_eq!(servers.len(), 1);
|
||||
match &servers[0].transport {
|
||||
McpServerTransport::Http { url, .. } => {
|
||||
assert_eq!(url, "https://example.com/mcp");
|
||||
}
|
||||
_ => panic!("expected Http"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn trait_is_object_safe() {
|
||||
let adapter: Box<dyn McpAgentAdapter> = Box::new(QwenAdapter);
|
||||
assert_eq!(adapter.source(), McpSource::Qwen);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,479 @@
|
||||
mod protocol;
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::ffi::OsString;
|
||||
use std::time::Duration;
|
||||
|
||||
use nomifun_api_types::{McpConnectionTestErrorCode, McpConnectionTestResult};
|
||||
use nomifun_runtime::{Builder as CmdBuilder, kill_process_tree, resolve_command_path};
|
||||
use serde::Serialize;
|
||||
use tokio::sync::mpsc;
|
||||
use tracing::{debug, warn};
|
||||
|
||||
use crate::types::McpServerTransport;
|
||||
use protocol::{
|
||||
JsonRpcRequest, JsonRpcResponse, SseEvent, build_http_headers, build_initialize_request,
|
||||
build_initialized_notification, build_tools_list_request, error_result, read_sse_events, rpc_error_result,
|
||||
run_stdio_protocol, spawn_error_result, success_result, timeout_result, wait_for_endpoint,
|
||||
wait_for_jsonrpc_response,
|
||||
};
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Constants
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
const CONNECTION_TIMEOUT: Duration = Duration::from_secs(30);
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// McpConnectionTestService
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Service for testing MCP server connectivity.
|
||||
///
|
||||
/// Creates a temporary MCP client, performs the protocol handshake
|
||||
/// (initialize -> initialized -> tools/list), and returns the tool list
|
||||
/// or an error. Supports stdio, HTTP (Streamable HTTP), and SSE transports.
|
||||
#[derive(Clone)]
|
||||
pub struct McpConnectionTestService {
|
||||
http_client: reqwest::Client,
|
||||
timeout: Duration,
|
||||
}
|
||||
|
||||
impl McpConnectionTestService {
|
||||
pub fn new(http_client: reqwest::Client) -> Self {
|
||||
Self {
|
||||
http_client,
|
||||
timeout: CONNECTION_TIMEOUT,
|
||||
}
|
||||
}
|
||||
|
||||
/// Override the connection test timeout (default: 30s).
|
||||
pub fn with_timeout(self, timeout: Duration) -> Self {
|
||||
Self { timeout, ..self }
|
||||
}
|
||||
|
||||
/// Test connectivity to an MCP server.
|
||||
///
|
||||
/// Dispatches to the appropriate transport handler. Always returns
|
||||
/// a result (never errors) -- failures are encoded in the struct.
|
||||
pub async fn test_connection(&self, name: &str, transport: &McpServerTransport) -> McpConnectionTestResult {
|
||||
debug!(name, ?transport, "starting MCP connection test");
|
||||
match transport {
|
||||
McpServerTransport::Stdio { command, args, env } => self.test_stdio(command, args, env).await,
|
||||
McpServerTransport::Http { url, headers } => self.test_http(url, headers).await,
|
||||
McpServerTransport::Sse { url, headers } => self.test_sse(url, headers).await,
|
||||
}
|
||||
}
|
||||
|
||||
// -- Stdio transport --------------------------------------------------
|
||||
|
||||
async fn test_stdio(
|
||||
&self,
|
||||
command: &str,
|
||||
args: &[String],
|
||||
env: &HashMap<String, String>,
|
||||
) -> McpConnectionTestResult {
|
||||
self.test_stdio_inner(command, args, env).await
|
||||
}
|
||||
|
||||
async fn test_stdio_inner(
|
||||
&self,
|
||||
command: &str,
|
||||
args: &[String],
|
||||
env: &HashMap<String, String>,
|
||||
) -> McpConnectionTestResult {
|
||||
let program = resolve_stdio_command(command);
|
||||
let mut cmd = CmdBuilder::new(&program);
|
||||
cmd.args(args)
|
||||
.envs(env.iter())
|
||||
.stdin(std::process::Stdio::piped())
|
||||
.stdout(std::process::Stdio::piped())
|
||||
.stderr(std::process::Stdio::null());
|
||||
|
||||
let mut child = match cmd.spawn() {
|
||||
Ok(c) => c,
|
||||
Err(e) => return spawn_error_result(command, &e),
|
||||
};
|
||||
|
||||
let stdin = child.stdin.take().expect("stdin was piped");
|
||||
let stdout = child.stdout.take().expect("stdout was piped");
|
||||
let result = match tokio::time::timeout(self.timeout, run_stdio_protocol(stdin, stdout)).await {
|
||||
Ok(r) => r,
|
||||
Err(_) => timeout_result(self.timeout),
|
||||
};
|
||||
if let Err(error) = kill_process_tree(&mut child).await {
|
||||
warn!(%error, "failed to clean up MCP stdio connection test process tree");
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
// -- HTTP (Streamable HTTP) transport ---------------------------------
|
||||
|
||||
async fn test_http(&self, url: &str, headers: &HashMap<String, String>) -> McpConnectionTestResult {
|
||||
match tokio::time::timeout(self.timeout, self.test_http_inner(url, headers)).await {
|
||||
Ok(r) => r,
|
||||
Err(_) => timeout_result(self.timeout),
|
||||
}
|
||||
}
|
||||
|
||||
async fn test_http_inner(&self, url: &str, headers: &HashMap<String, String>) -> McpConnectionTestResult {
|
||||
let mut req_headers = build_http_headers(headers);
|
||||
req_headers.insert(
|
||||
reqwest::header::CONTENT_TYPE,
|
||||
"application/json".parse().expect("valid header"),
|
||||
);
|
||||
req_headers.insert(
|
||||
reqwest::header::ACCEPT,
|
||||
"application/json, text/event-stream".parse().expect("valid header"),
|
||||
);
|
||||
|
||||
// 1. initialize
|
||||
let init_resp = match self
|
||||
.http_post_mcp(url, &req_headers, &build_initialize_request(1))
|
||||
.await
|
||||
{
|
||||
Ok(r) => r,
|
||||
Err(result) => return result,
|
||||
};
|
||||
if let Some(err) = init_resp.rpc.error {
|
||||
return rpc_error_result("initialize", &err);
|
||||
}
|
||||
|
||||
// Extract session ID for subsequent requests
|
||||
if let Some(sid) = init_resp.session_id
|
||||
&& let Ok(val) = reqwest::header::HeaderValue::from_str(&sid)
|
||||
{
|
||||
req_headers.insert("mcp-session-id", val);
|
||||
}
|
||||
|
||||
// 2. initialized notification (fire-and-forget)
|
||||
let _ = self
|
||||
.http_client
|
||||
.post(url)
|
||||
.headers(req_headers.clone())
|
||||
.json(&build_initialized_notification())
|
||||
.send()
|
||||
.await;
|
||||
|
||||
// 3. tools/list
|
||||
let tools_resp = match self
|
||||
.http_post_mcp(url, &req_headers, &build_tools_list_request(2))
|
||||
.await
|
||||
{
|
||||
Ok(r) => r,
|
||||
Err(result) => return result,
|
||||
};
|
||||
if let Some(err) = tools_resp.rpc.error {
|
||||
return rpc_error_result("tools/list", &err);
|
||||
}
|
||||
|
||||
success_result(tools_resp.rpc.result)
|
||||
}
|
||||
|
||||
/// POST a JSON-RPC message and parse the response.
|
||||
///
|
||||
/// Returns `Err(McpConnectionTestResult)` for HTTP-level failures
|
||||
/// (connection error, 401, non-success status).
|
||||
async fn http_post_mcp(
|
||||
&self,
|
||||
url: &str,
|
||||
headers: &reqwest::header::HeaderMap,
|
||||
body: &JsonRpcRequest,
|
||||
) -> Result<HttpMcpResponse, McpConnectionTestResult> {
|
||||
let resp = self
|
||||
.http_client
|
||||
.post(url)
|
||||
.headers(headers.clone())
|
||||
.json(body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| {
|
||||
error_result(
|
||||
McpConnectionTestErrorCode::ConnectionFailed,
|
||||
format!("Connection failed: {e}"),
|
||||
Some(serde_json::json!({ "transport": "http" })),
|
||||
)
|
||||
})?;
|
||||
|
||||
if resp.status() == reqwest::StatusCode::UNAUTHORIZED {
|
||||
return Err(protocol::auth_result(resp.headers()));
|
||||
}
|
||||
if !resp.status().is_success() {
|
||||
return Err(error_result(
|
||||
McpConnectionTestErrorCode::HttpError,
|
||||
format!("HTTP {} from server", resp.status()),
|
||||
Some(serde_json::json!({ "status": resp.status().as_u16() })),
|
||||
));
|
||||
}
|
||||
|
||||
let session_id = resp
|
||||
.headers()
|
||||
.get("mcp-session-id")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.map(String::from);
|
||||
|
||||
let rpc = protocol::parse_http_response(resp).await.map_err(|error| {
|
||||
error_result(
|
||||
McpConnectionTestErrorCode::ProtocolError,
|
||||
error,
|
||||
Some(serde_json::json!({ "transport": "http" })),
|
||||
)
|
||||
})?;
|
||||
Ok(HttpMcpResponse { rpc, session_id })
|
||||
}
|
||||
|
||||
// -- SSE transport ----------------------------------------------------
|
||||
|
||||
async fn test_sse(&self, url: &str, headers: &HashMap<String, String>) -> McpConnectionTestResult {
|
||||
match tokio::time::timeout(self.timeout, self.test_sse_inner(url, headers)).await {
|
||||
Ok(r) => r,
|
||||
Err(_) => timeout_result(self.timeout),
|
||||
}
|
||||
}
|
||||
|
||||
async fn test_sse_inner(&self, url: &str, headers: &HashMap<String, String>) -> McpConnectionTestResult {
|
||||
let mut req_headers = build_http_headers(headers);
|
||||
|
||||
// 1. Open SSE connection
|
||||
let resp = match self
|
||||
.http_client
|
||||
.get(url)
|
||||
.headers(req_headers.clone())
|
||||
.header(reqwest::header::ACCEPT, "text/event-stream")
|
||||
.send()
|
||||
.await
|
||||
{
|
||||
Ok(r) => r,
|
||||
Err(e) => {
|
||||
return error_result(
|
||||
McpConnectionTestErrorCode::ConnectionFailed,
|
||||
format!("Connection failed: {e}"),
|
||||
Some(serde_json::json!({ "transport": "sse" })),
|
||||
);
|
||||
}
|
||||
};
|
||||
if resp.status() == reqwest::StatusCode::UNAUTHORIZED {
|
||||
return protocol::auth_result(resp.headers());
|
||||
}
|
||||
if !resp.status().is_success() {
|
||||
return error_result(
|
||||
McpConnectionTestErrorCode::HttpError,
|
||||
format!("HTTP {} from server", resp.status()),
|
||||
Some(serde_json::json!({ "status": resp.status().as_u16() })),
|
||||
);
|
||||
}
|
||||
|
||||
// 2. Start SSE reader task
|
||||
let (event_tx, mut event_rx) = mpsc::channel::<SseEvent>(16);
|
||||
let reader_handle = tokio::spawn(read_sse_events(resp, event_tx));
|
||||
|
||||
req_headers.insert(
|
||||
reqwest::header::CONTENT_TYPE,
|
||||
"application/json".parse().expect("valid header"),
|
||||
);
|
||||
|
||||
let result = self.run_sse_protocol(url, &req_headers, &mut event_rx).await;
|
||||
reader_handle.abort();
|
||||
result
|
||||
}
|
||||
|
||||
async fn run_sse_protocol(
|
||||
&self,
|
||||
base_url: &str,
|
||||
headers: &reqwest::header::HeaderMap,
|
||||
event_rx: &mut mpsc::Receiver<SseEvent>,
|
||||
) -> McpConnectionTestResult {
|
||||
// 3. Wait for endpoint event
|
||||
let endpoint = match wait_for_endpoint(event_rx, base_url).await {
|
||||
Ok(ep) => ep,
|
||||
Err(e) => {
|
||||
return error_result(
|
||||
McpConnectionTestErrorCode::ProtocolError,
|
||||
e,
|
||||
Some(serde_json::json!({ "transport": "sse", "stage": "endpoint" })),
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
// 4. initialize
|
||||
if let Err(e) = self.sse_post(&endpoint, headers, &build_initialize_request(1)).await {
|
||||
return error_result(
|
||||
McpConnectionTestErrorCode::ProtocolError,
|
||||
format!("Failed to send initialize: {e}"),
|
||||
Some(serde_json::json!({ "transport": "sse", "stage": "initialize_send" })),
|
||||
);
|
||||
}
|
||||
let init_resp = match wait_for_jsonrpc_response(event_rx).await {
|
||||
Ok(r) => r,
|
||||
Err(e) => {
|
||||
return error_result(
|
||||
McpConnectionTestErrorCode::ProtocolError,
|
||||
format!("initialize response: {e}"),
|
||||
Some(serde_json::json!({ "transport": "sse", "stage": "initialize_response" })),
|
||||
);
|
||||
}
|
||||
};
|
||||
if let Some(err) = init_resp.error {
|
||||
return rpc_error_result("initialize", &err);
|
||||
}
|
||||
|
||||
// 5. initialized notification
|
||||
let _ = self
|
||||
.sse_post(&endpoint, headers, &build_initialized_notification())
|
||||
.await;
|
||||
|
||||
// 6. tools/list
|
||||
if let Err(e) = self.sse_post(&endpoint, headers, &build_tools_list_request(2)).await {
|
||||
return error_result(
|
||||
McpConnectionTestErrorCode::ProtocolError,
|
||||
format!("Failed to send tools/list: {e}"),
|
||||
Some(serde_json::json!({ "transport": "sse", "stage": "tools_list_send" })),
|
||||
);
|
||||
}
|
||||
let tools_resp = match wait_for_jsonrpc_response(event_rx).await {
|
||||
Ok(r) => r,
|
||||
Err(e) => {
|
||||
return error_result(
|
||||
McpConnectionTestErrorCode::ProtocolError,
|
||||
format!("tools/list response: {e}"),
|
||||
Some(serde_json::json!({ "transport": "sse", "stage": "tools_list_response" })),
|
||||
);
|
||||
}
|
||||
};
|
||||
if let Some(err) = tools_resp.error {
|
||||
return rpc_error_result("tools/list", &err);
|
||||
}
|
||||
|
||||
success_result(tools_resp.result)
|
||||
}
|
||||
|
||||
/// POST a JSON-RPC message to an SSE endpoint (fire-and-forget semantics).
|
||||
async fn sse_post<T: Serialize>(
|
||||
&self,
|
||||
endpoint: &str,
|
||||
headers: &reqwest::header::HeaderMap,
|
||||
body: &T,
|
||||
) -> Result<(), String> {
|
||||
self.http_client
|
||||
.post(endpoint)
|
||||
.headers(headers.clone())
|
||||
.json(body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
fn resolve_stdio_command(command: &str) -> OsString {
|
||||
if !command.is_empty()
|
||||
&& !command.contains('/')
|
||||
&& !command.contains('\\')
|
||||
&& let Some(path) = resolve_command_path(command)
|
||||
{
|
||||
return path.into_os_string();
|
||||
}
|
||||
|
||||
OsString::from(command)
|
||||
}
|
||||
|
||||
/// Intermediate struct for HTTP transport response parsing.
|
||||
struct HttpMcpResponse {
|
||||
rpc: JsonRpcResponse,
|
||||
session_id: Option<String>,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn service_clone() {
|
||||
let svc = McpConnectionTestService::new(reqwest::Client::new());
|
||||
let _cloned = svc.clone();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn service_with_timeout() {
|
||||
let svc = McpConnectionTestService::new(reqwest::Client::new()).with_timeout(Duration::from_secs(5));
|
||||
assert_eq!(svc.timeout, Duration::from_secs(5));
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
#[tokio::test]
|
||||
async fn stdio_timeout_cleans_up_process_group() {
|
||||
let marker_path = std::env::temp_dir().join(format!(
|
||||
"nomifun-mcp-timeout-pid-{}-{}",
|
||||
std::process::id(),
|
||||
std::time::SystemTime::now()
|
||||
.duration_since(std::time::UNIX_EPOCH)
|
||||
.unwrap()
|
||||
.as_nanos()
|
||||
));
|
||||
let transport = McpServerTransport::Stdio {
|
||||
command: "sh".into(),
|
||||
args: vec![
|
||||
"-c".into(),
|
||||
"printf '%s\n' \"$$\" > \"$1\"; sleep 30".into(),
|
||||
"mcp-timeout-child".into(),
|
||||
marker_path.to_string_lossy().into_owned(),
|
||||
],
|
||||
env: HashMap::new(),
|
||||
};
|
||||
let svc = McpConnectionTestService::new(reqwest::Client::new()).with_timeout(Duration::from_millis(100));
|
||||
|
||||
let result = svc.test_connection("timeout-cleanup", &transport).await;
|
||||
assert!(!result.success);
|
||||
assert!(
|
||||
result.error.as_deref().unwrap_or_default().contains("timed out"),
|
||||
"expected timeout result, got {result:?}"
|
||||
);
|
||||
|
||||
let pid: i32 = std::fs::read_to_string(&marker_path)
|
||||
.expect("stdio child should write its pid")
|
||||
.trim()
|
||||
.parse()
|
||||
.expect("pid marker should be numeric");
|
||||
|
||||
let group_alive = wait_for_process_group_exit(pid, Duration::from_secs(1)).await;
|
||||
if group_alive {
|
||||
let _ = kill_process_group(pid, libc_sigkill());
|
||||
}
|
||||
let _ = std::fs::remove_file(marker_path);
|
||||
|
||||
assert!(
|
||||
!group_alive,
|
||||
"stdio timeout should terminate the spawned process group for pid={pid}"
|
||||
);
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
async fn wait_for_process_group_exit(pid: i32, timeout: Duration) -> bool {
|
||||
let deadline = tokio::time::Instant::now() + timeout;
|
||||
while tokio::time::Instant::now() < deadline {
|
||||
if !is_process_group_alive(pid) {
|
||||
return false;
|
||||
}
|
||||
tokio::time::sleep(Duration::from_millis(50)).await;
|
||||
}
|
||||
is_process_group_alive(pid)
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
fn is_process_group_alive(pid: i32) -> bool {
|
||||
kill_process_group(pid, 0)
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
fn kill_process_group(pid: i32, signal: i32) -> bool {
|
||||
unsafe extern "C" {
|
||||
fn kill(pid: i32, sig: i32) -> i32;
|
||||
}
|
||||
unsafe { kill(-pid, signal) == 0 }
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
fn libc_sigkill() -> i32 {
|
||||
9
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,884 @@
|
||||
// Protocol types, helpers, and message builders for MCP connection testing.
|
||||
//
|
||||
// This module implements the minimal JSON-RPC 2.0 subset needed for the
|
||||
// MCP handshake: initialize → initialized → tools/list.
|
||||
|
||||
use std::collections::HashMap;
|
||||
|
||||
use nomifun_api_types::{McpAuthMethod, McpConnectionTestErrorCode, McpConnectionTestResult, McpToolResponse};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
use std::time::Duration;
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Constants
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
const PROTOCOL_VERSION: &str = "2024-11-05";
|
||||
const CLIENT_NAME: &str = "nomifun-mcp-test";
|
||||
const CLIENT_VERSION: &str = "1.0.0";
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// JSON-RPC message types
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
pub(super) struct JsonRpcRequest {
|
||||
pub jsonrpc: &'static str,
|
||||
pub id: u64,
|
||||
pub method: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub params: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
pub(super) struct JsonRpcNotification {
|
||||
pub jsonrpc: &'static str,
|
||||
pub method: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub(super) struct JsonRpcResponse {
|
||||
#[allow(dead_code)]
|
||||
pub jsonrpc: String,
|
||||
pub id: Option<u64>,
|
||||
pub result: Option<serde_json::Value>,
|
||||
pub error: Option<JsonRpcError>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub(super) struct JsonRpcError {
|
||||
pub code: i64,
|
||||
pub message: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct ToolsListResult {
|
||||
tools: Vec<McpToolInfo>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct McpToolInfo {
|
||||
name: String,
|
||||
description: Option<String>,
|
||||
#[serde(rename = "inputSchema")]
|
||||
input_schema: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
/// A single SSE event parsed from the stream.
|
||||
pub(super) struct SseEvent {
|
||||
pub event_type: String,
|
||||
pub data: String,
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Stdio protocol helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Run the MCP protocol handshake over stdio (newline-delimited JSON-RPC).
|
||||
pub(super) async fn run_stdio_protocol(
|
||||
mut stdin: tokio::process::ChildStdin,
|
||||
stdout: tokio::process::ChildStdout,
|
||||
) -> McpConnectionTestResult {
|
||||
let mut reader = BufReader::new(stdout);
|
||||
|
||||
// 1. initialize
|
||||
if let Err(e) = write_jsonrpc_line(&mut stdin, &build_initialize_request(1)).await {
|
||||
return error_result(
|
||||
McpConnectionTestErrorCode::ProtocolError,
|
||||
format!("Failed to send initialize: {e}"),
|
||||
Some(serde_json::json!({ "transport": "stdio", "stage": "initialize_send" })),
|
||||
);
|
||||
}
|
||||
let init_resp = match read_jsonrpc_response(&mut reader).await {
|
||||
Ok(r) => r,
|
||||
Err(e) => {
|
||||
return error_result(
|
||||
McpConnectionTestErrorCode::ProtocolError,
|
||||
format!("initialize response: {e}"),
|
||||
Some(serde_json::json!({ "transport": "stdio", "stage": "initialize_response" })),
|
||||
);
|
||||
}
|
||||
};
|
||||
if let Some(err) = init_resp.error {
|
||||
return rpc_error_result("initialize", &err);
|
||||
}
|
||||
|
||||
// 2. initialized notification
|
||||
if let Err(e) = write_jsonrpc_line(&mut stdin, &build_initialized_notification()).await {
|
||||
return error_result(
|
||||
McpConnectionTestErrorCode::ProtocolError,
|
||||
format!("Failed to send initialized: {e}"),
|
||||
Some(serde_json::json!({ "transport": "stdio", "stage": "initialized_send" })),
|
||||
);
|
||||
}
|
||||
|
||||
// 3. tools/list
|
||||
if let Err(e) = write_jsonrpc_line(&mut stdin, &build_tools_list_request(2)).await {
|
||||
return error_result(
|
||||
McpConnectionTestErrorCode::ProtocolError,
|
||||
format!("Failed to send tools/list: {e}"),
|
||||
Some(serde_json::json!({ "transport": "stdio", "stage": "tools_list_send" })),
|
||||
);
|
||||
}
|
||||
let tools_resp = match read_jsonrpc_response(&mut reader).await {
|
||||
Ok(r) => r,
|
||||
Err(e) => {
|
||||
return error_result(
|
||||
McpConnectionTestErrorCode::ProtocolError,
|
||||
format!("tools/list response: {e}"),
|
||||
Some(serde_json::json!({ "transport": "stdio", "stage": "tools_list_response" })),
|
||||
);
|
||||
}
|
||||
};
|
||||
if let Some(err) = tools_resp.error {
|
||||
return rpc_error_result("tools/list", &err);
|
||||
}
|
||||
|
||||
success_result(tools_resp.result)
|
||||
}
|
||||
|
||||
/// Write a JSON-RPC message as a newline-delimited line to stdin.
|
||||
async fn write_jsonrpc_line<T: Serialize>(stdin: &mut tokio::process::ChildStdin, msg: &T) -> std::io::Result<()> {
|
||||
let json = serde_json::to_string(msg).map_err(std::io::Error::other)?;
|
||||
stdin.write_all(json.as_bytes()).await?;
|
||||
stdin.write_all(b"\n").await?;
|
||||
stdin.flush().await
|
||||
}
|
||||
|
||||
/// Read the next JSON-RPC response from stdout.
|
||||
///
|
||||
/// Skips server notifications (messages without an `id` field) and
|
||||
/// non-JSON lines (e.g. logging output).
|
||||
async fn read_jsonrpc_response(reader: &mut BufReader<tokio::process::ChildStdout>) -> Result<JsonRpcResponse, String> {
|
||||
let mut line = String::new();
|
||||
loop {
|
||||
line.clear();
|
||||
let n = reader
|
||||
.read_line(&mut line)
|
||||
.await
|
||||
.map_err(|e| format!("I/O error: {e}"))?;
|
||||
if n == 0 {
|
||||
return Err("Server closed stdout before responding".into());
|
||||
}
|
||||
let trimmed = line.trim();
|
||||
if trimmed.is_empty() {
|
||||
continue;
|
||||
}
|
||||
if let Ok(resp) = serde_json::from_str::<JsonRpcResponse>(trimmed)
|
||||
&& resp.id.is_some()
|
||||
{
|
||||
return Ok(resp);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// SSE helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Read SSE events from a streaming HTTP response and forward via channel.
|
||||
pub(super) async fn read_sse_events(mut resp: reqwest::Response, tx: mpsc::Sender<SseEvent>) {
|
||||
let mut buffer = String::new();
|
||||
loop {
|
||||
match resp.chunk().await {
|
||||
Ok(Some(chunk)) => {
|
||||
// Normalize line endings for consistent parsing
|
||||
let text = String::from_utf8_lossy(&chunk);
|
||||
buffer.push_str(&text.replace("\r\n", "\n"));
|
||||
while let Some(event) = parse_next_sse_event(&mut buffer) {
|
||||
if tx.send(event).await.is_err() {
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(None) | Err(_) => return,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Parse the next complete SSE event from the buffer.
|
||||
///
|
||||
/// An SSE event is terminated by a blank line (`\n\n`).
|
||||
/// Returns `None` if no complete event is available yet.
|
||||
fn parse_next_sse_event(buffer: &mut String) -> Option<SseEvent> {
|
||||
let end = buffer.find("\n\n")?;
|
||||
let event_text: String = buffer.drain(..end + 2).collect();
|
||||
|
||||
let mut event_type = String::new();
|
||||
let mut data_parts: Vec<&str> = Vec::new();
|
||||
|
||||
for line in event_text.lines() {
|
||||
if let Some(rest) = line.strip_prefix("event:") {
|
||||
event_type = rest.trim().to_owned();
|
||||
} else if let Some(rest) = line.strip_prefix("data:") {
|
||||
// SSE spec: strip one leading space if present
|
||||
data_parts.push(rest.strip_prefix(' ').unwrap_or(rest));
|
||||
}
|
||||
}
|
||||
|
||||
Some(SseEvent {
|
||||
event_type,
|
||||
data: data_parts.join("\n"),
|
||||
})
|
||||
}
|
||||
|
||||
/// Wait for an `endpoint` event from the SSE stream and resolve the URL.
|
||||
pub(super) async fn wait_for_endpoint(
|
||||
event_rx: &mut mpsc::Receiver<SseEvent>,
|
||||
base_url: &str,
|
||||
) -> Result<String, String> {
|
||||
loop {
|
||||
match event_rx.recv().await {
|
||||
Some(event) if event.event_type == "endpoint" => {
|
||||
return resolve_endpoint_url(base_url, &event.data);
|
||||
}
|
||||
Some(_) => continue,
|
||||
None => return Err("SSE stream closed before endpoint event".into()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Wait for the next JSON-RPC response from the SSE stream.
|
||||
pub(super) async fn wait_for_jsonrpc_response(
|
||||
event_rx: &mut mpsc::Receiver<SseEvent>,
|
||||
) -> Result<JsonRpcResponse, String> {
|
||||
loop {
|
||||
match event_rx.recv().await {
|
||||
Some(event) if event.event_type == "message" => {
|
||||
let resp: JsonRpcResponse =
|
||||
serde_json::from_str(&event.data).map_err(|e| format!("Invalid JSON-RPC in SSE: {e}"))?;
|
||||
if resp.id.is_some() {
|
||||
return Ok(resp);
|
||||
}
|
||||
}
|
||||
Some(_) => continue,
|
||||
None => return Err("SSE stream closed before response".into()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Resolve a potentially relative endpoint URL against a base URL.
|
||||
fn resolve_endpoint_url(base_url: &str, endpoint: &str) -> Result<String, String> {
|
||||
let base = reqwest::Url::parse(base_url).map_err(|e| format!("Invalid base URL: {e}"))?;
|
||||
base.join(endpoint)
|
||||
.map(|u| u.to_string())
|
||||
.map_err(|e| format!("Invalid endpoint URL: {e}"))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// HTTP helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Build an HTTP header map from a string-to-string map.
|
||||
pub(super) fn build_http_headers(headers: &HashMap<String, String>) -> reqwest::header::HeaderMap {
|
||||
let mut map = reqwest::header::HeaderMap::new();
|
||||
for (k, v) in headers {
|
||||
if let (Ok(name), Ok(val)) = (
|
||||
reqwest::header::HeaderName::from_bytes(k.as_bytes()),
|
||||
reqwest::header::HeaderValue::from_str(v),
|
||||
) {
|
||||
map.insert(name, val);
|
||||
}
|
||||
}
|
||||
map
|
||||
}
|
||||
|
||||
/// Parse a JSON-RPC response from an HTTP response body.
|
||||
///
|
||||
/// Handles both `application/json` and `text/event-stream` content types
|
||||
/// (Streamable HTTP servers may respond with either).
|
||||
pub(super) async fn parse_http_response(resp: reqwest::Response) -> Result<JsonRpcResponse, String> {
|
||||
let is_sse = resp
|
||||
.headers()
|
||||
.get(reqwest::header::CONTENT_TYPE)
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.is_some_and(|ct| ct.contains("text/event-stream"));
|
||||
|
||||
let body = resp.text().await.map_err(|e| format!("Failed to read response: {e}"))?;
|
||||
|
||||
if is_sse {
|
||||
extract_jsonrpc_from_sse(&body)
|
||||
} else {
|
||||
serde_json::from_str(&body).map_err(|e| format!("Invalid JSON-RPC response: {e}"))
|
||||
}
|
||||
}
|
||||
|
||||
/// Extract the first JSON-RPC response from SSE event data.
|
||||
fn extract_jsonrpc_from_sse(body: &str) -> Result<JsonRpcResponse, String> {
|
||||
for line in body.lines() {
|
||||
if let Some(data) = line.strip_prefix("data:") {
|
||||
let data = data.strip_prefix(' ').unwrap_or(data);
|
||||
if let Ok(resp) = serde_json::from_str::<JsonRpcResponse>(data)
|
||||
&& resp.id.is_some()
|
||||
{
|
||||
return Ok(resp);
|
||||
}
|
||||
}
|
||||
}
|
||||
Err("No JSON-RPC response found in SSE data".into())
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// JSON-RPC message builders
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
pub(super) fn build_initialize_request(id: u64) -> JsonRpcRequest {
|
||||
JsonRpcRequest {
|
||||
jsonrpc: "2.0",
|
||||
id,
|
||||
method: "initialize".into(),
|
||||
params: Some(serde_json::json!({
|
||||
"protocolVersion": PROTOCOL_VERSION,
|
||||
"capabilities": {},
|
||||
"clientInfo": {
|
||||
"name": CLIENT_NAME,
|
||||
"version": CLIENT_VERSION
|
||||
}
|
||||
})),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn build_initialized_notification() -> JsonRpcNotification {
|
||||
JsonRpcNotification {
|
||||
jsonrpc: "2.0",
|
||||
method: "notifications/initialized".into(),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn build_tools_list_request(id: u64) -> JsonRpcRequest {
|
||||
JsonRpcRequest {
|
||||
jsonrpc: "2.0",
|
||||
id,
|
||||
method: "tools/list".into(),
|
||||
params: None,
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Result builders
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
pub(super) fn success_result(tools_value: Option<serde_json::Value>) -> McpConnectionTestResult {
|
||||
let tools = tools_value
|
||||
.and_then(|v| serde_json::from_value::<ToolsListResult>(v).ok())
|
||||
.map(|r| {
|
||||
r.tools
|
||||
.into_iter()
|
||||
.map(|t| McpToolResponse {
|
||||
name: t.name,
|
||||
description: t.description,
|
||||
input_schema: t.input_schema,
|
||||
})
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
|
||||
McpConnectionTestResult {
|
||||
success: true,
|
||||
tools: Some(tools),
|
||||
error: None,
|
||||
code: None,
|
||||
details: None,
|
||||
needs_auth: None,
|
||||
auth_method: None,
|
||||
www_authenticate: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn error_result(
|
||||
code: McpConnectionTestErrorCode,
|
||||
msg: String,
|
||||
details: Option<serde_json::Value>,
|
||||
) -> McpConnectionTestResult {
|
||||
McpConnectionTestResult {
|
||||
success: false,
|
||||
tools: None,
|
||||
error: Some(msg),
|
||||
code: Some(code),
|
||||
details,
|
||||
needs_auth: None,
|
||||
auth_method: None,
|
||||
www_authenticate: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn timeout_result(duration: Duration) -> McpConnectionTestResult {
|
||||
error_result(
|
||||
McpConnectionTestErrorCode::Timeout,
|
||||
format!("Connection test timed out after {}s", duration.as_secs()),
|
||||
Some(serde_json::json!({ "timeout_seconds": duration.as_secs() })),
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) fn spawn_error_result(command: &str, error: &std::io::Error) -> McpConnectionTestResult {
|
||||
match error.kind() {
|
||||
std::io::ErrorKind::NotFound => {
|
||||
let runtime = missing_command_runtime(command);
|
||||
error_result(
|
||||
McpConnectionTestErrorCode::CommandNotFound,
|
||||
command_not_found_message(command),
|
||||
Some(serde_json::json!({
|
||||
"command": command,
|
||||
"runtime": runtime,
|
||||
})),
|
||||
)
|
||||
}
|
||||
std::io::ErrorKind::PermissionDenied => error_result(
|
||||
McpConnectionTestErrorCode::CommandPermissionDenied,
|
||||
format!("Permission denied: {command}"),
|
||||
Some(serde_json::json!({ "command": command })),
|
||||
),
|
||||
_ => error_result(
|
||||
McpConnectionTestErrorCode::CommandStartFailed,
|
||||
format!("Failed to start '{command}': {error}"),
|
||||
Some(serde_json::json!({
|
||||
"command": command,
|
||||
"io_error": error.to_string(),
|
||||
})),
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
fn command_basename(command: &str) -> String {
|
||||
let mut command_name = command
|
||||
.rsplit(['/', '\\'])
|
||||
.next()
|
||||
.unwrap_or(command)
|
||||
.to_ascii_lowercase();
|
||||
for suffix in [".exe", ".cmd", ".bat"] {
|
||||
if let Some(stripped) = command_name.strip_suffix(suffix) {
|
||||
command_name = stripped.to_owned();
|
||||
break;
|
||||
}
|
||||
}
|
||||
command_name
|
||||
}
|
||||
|
||||
fn missing_command_runtime(command: &str) -> &'static str {
|
||||
let command_name = command_basename(command);
|
||||
match command_name.as_str() {
|
||||
"npx" | "npm" | "node" | "pnpx" => "node",
|
||||
"bun" | "bunx" => "bun",
|
||||
"uv" | "uvx" => "uv",
|
||||
"python" | "python3" => "python",
|
||||
"deno" => "deno",
|
||||
_ => "generic",
|
||||
}
|
||||
}
|
||||
|
||||
fn command_not_found_message(command: &str) -> String {
|
||||
match missing_command_runtime(command) {
|
||||
"node" => format!(
|
||||
"Command not found: {command}. Install Node.js (which includes npm/npx), then restart Nomi or configure this MCP server to use an absolute command path."
|
||||
),
|
||||
"bun" => format!(
|
||||
"Command not found: {command}. Install Bun (which includes bun/bunx), then restart Nomi or configure this MCP server to use an absolute command path."
|
||||
),
|
||||
"uv" => format!(
|
||||
"Command not found: {command}. Install uv, then restart Nomi or configure this MCP server to use an absolute command path."
|
||||
),
|
||||
"python" => format!(
|
||||
"Command not found: {command}. Install Python, then restart Nomi or configure this MCP server to use an absolute command path."
|
||||
),
|
||||
"deno" => format!(
|
||||
"Command not found: {command}. Install Deno, then restart Nomi or configure this MCP server to use an absolute command path."
|
||||
),
|
||||
_ => format!(
|
||||
"Command not found: {command}. Install the command or configure this MCP server to use an absolute command path."
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn rpc_error_result(method: &str, err: &JsonRpcError) -> McpConnectionTestResult {
|
||||
error_result(
|
||||
McpConnectionTestErrorCode::RpcError,
|
||||
format!("{method} error: {} (code {})", err.message, err.code),
|
||||
Some(serde_json::json!({
|
||||
"method": method,
|
||||
"rpc_code": err.code,
|
||||
})),
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) fn auth_result(headers: &reqwest::header::HeaderMap) -> McpConnectionTestResult {
|
||||
let www_authenticate = headers
|
||||
.get("www-authenticate")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.map(String::from);
|
||||
|
||||
let auth_method = www_authenticate.as_deref().map(detect_auth_method);
|
||||
|
||||
McpConnectionTestResult {
|
||||
success: false,
|
||||
tools: None,
|
||||
error: None,
|
||||
code: None,
|
||||
details: None,
|
||||
needs_auth: Some(true),
|
||||
auth_method,
|
||||
www_authenticate,
|
||||
}
|
||||
}
|
||||
|
||||
fn detect_auth_method(www_authenticate: &str) -> McpAuthMethod {
|
||||
let lower = www_authenticate.to_lowercase();
|
||||
if lower.contains("bearer") || lower.contains("oauth") {
|
||||
McpAuthMethod::Oauth
|
||||
} else {
|
||||
McpAuthMethod::Basic
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
// -- SSE event parsing ------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn parse_sse_event_basic() {
|
||||
let mut buf = "event: endpoint\ndata: /messages\n\n".to_string();
|
||||
let event = parse_next_sse_event(&mut buf).unwrap();
|
||||
assert_eq!(event.event_type, "endpoint");
|
||||
assert_eq!(event.data, "/messages");
|
||||
assert!(buf.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_sse_event_with_json_data() {
|
||||
let mut buf = "event: message\ndata: {\"jsonrpc\":\"2.0\",\"id\":1,\"result\":{}}\n\n".to_string();
|
||||
let event = parse_next_sse_event(&mut buf).unwrap();
|
||||
assert_eq!(event.event_type, "message");
|
||||
assert!(event.data.contains("jsonrpc"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_sse_event_multiline_data() {
|
||||
let mut buf = "event: message\ndata: line1\ndata: line2\n\n".to_string();
|
||||
let event = parse_next_sse_event(&mut buf).unwrap();
|
||||
assert_eq!(event.data, "line1\nline2");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_sse_event_no_leading_space() {
|
||||
let mut buf = "event: test\ndata:no-space\n\n".to_string();
|
||||
let event = parse_next_sse_event(&mut buf).unwrap();
|
||||
assert_eq!(event.data, "no-space");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_sse_event_incomplete_returns_none() {
|
||||
let mut buf = "event: endpoint\ndata: /msg".to_string();
|
||||
assert!(parse_next_sse_event(&mut buf).is_none());
|
||||
assert_eq!(buf, "event: endpoint\ndata: /msg");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_sse_event_empty_buffer() {
|
||||
let mut buf = String::new();
|
||||
assert!(parse_next_sse_event(&mut buf).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_sse_event_multiple_in_buffer() {
|
||||
let mut buf = "event: a\ndata: 1\n\nevent: b\ndata: 2\n\n".to_string();
|
||||
let first = parse_next_sse_event(&mut buf).unwrap();
|
||||
assert_eq!(first.event_type, "a");
|
||||
assert_eq!(first.data, "1");
|
||||
let second = parse_next_sse_event(&mut buf).unwrap();
|
||||
assert_eq!(second.event_type, "b");
|
||||
assert_eq!(second.data, "2");
|
||||
}
|
||||
|
||||
// -- URL resolution ---------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn resolve_absolute_endpoint() {
|
||||
let result = resolve_endpoint_url("https://example.com/sse", "https://other.com/messages");
|
||||
assert_eq!(result.unwrap(), "https://other.com/messages");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolve_relative_endpoint() {
|
||||
let result = resolve_endpoint_url("https://example.com/sse", "/messages?s=123");
|
||||
assert_eq!(result.unwrap(), "https://example.com/messages?s=123");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolve_relative_path_endpoint() {
|
||||
let result = resolve_endpoint_url("https://example.com/mcp/sse", "messages");
|
||||
assert_eq!(result.unwrap(), "https://example.com/mcp/messages");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolve_invalid_base_url() {
|
||||
let result = resolve_endpoint_url("not-a-url", "/messages");
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
// -- Auth detection ---------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn detect_bearer_as_oauth() {
|
||||
assert!(matches!(
|
||||
detect_auth_method("Bearer realm=\"mcp\""),
|
||||
McpAuthMethod::Oauth
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn detect_oauth_keyword() {
|
||||
assert!(matches!(
|
||||
detect_auth_method("OAuth realm=\"mcp\""),
|
||||
McpAuthMethod::Oauth
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn detect_basic_auth() {
|
||||
assert!(matches!(
|
||||
detect_auth_method("Basic realm=\"mcp\""),
|
||||
McpAuthMethod::Basic
|
||||
));
|
||||
}
|
||||
|
||||
// -- Result builders --------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn success_result_with_tools() {
|
||||
let tools_json = serde_json::json!({
|
||||
"tools": [
|
||||
{ "name": "read_file", "description": "Read a file" },
|
||||
{ "name": "write_file" }
|
||||
]
|
||||
});
|
||||
let result = success_result(Some(tools_json));
|
||||
assert!(result.success);
|
||||
let tools = result.tools.unwrap();
|
||||
assert_eq!(tools.len(), 2);
|
||||
assert_eq!(tools[0].name, "read_file");
|
||||
assert_eq!(tools[0].description.as_deref(), Some("Read a file"));
|
||||
assert!(tools[1].description.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn success_result_empty_tools() {
|
||||
let tools_json = serde_json::json!({ "tools": [] });
|
||||
let result = success_result(Some(tools_json));
|
||||
assert!(result.success);
|
||||
assert!(result.tools.unwrap().is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn success_result_none_gives_empty_tools() {
|
||||
let result = success_result(None);
|
||||
assert!(result.success);
|
||||
assert!(result.tools.unwrap().is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn success_result_malformed_gives_empty_tools() {
|
||||
let result = success_result(Some(serde_json::json!("not an object")));
|
||||
assert!(result.success);
|
||||
assert!(result.tools.unwrap().is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn error_result_fields() {
|
||||
let result = error_result(
|
||||
McpConnectionTestErrorCode::ProtocolError,
|
||||
"something broke".into(),
|
||||
Some(serde_json::json!({ "stage": "initialize" })),
|
||||
);
|
||||
assert!(!result.success);
|
||||
assert_eq!(result.error.as_deref(), Some("something broke"));
|
||||
assert_eq!(result.code, Some(McpConnectionTestErrorCode::ProtocolError));
|
||||
assert_eq!(result.details.unwrap()["stage"], "initialize");
|
||||
assert!(result.tools.is_none());
|
||||
assert!(result.needs_auth.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn timeout_result_message() {
|
||||
let result = timeout_result(Duration::from_secs(30));
|
||||
assert!(!result.success);
|
||||
assert!(result.error.as_deref().unwrap().contains("30s"));
|
||||
assert_eq!(result.code, Some(McpConnectionTestErrorCode::Timeout));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn spawn_error_not_found() {
|
||||
let err = std::io::Error::new(std::io::ErrorKind::NotFound, "not found");
|
||||
let result = spawn_error_result("npx", &err);
|
||||
let error = result.error.as_deref().unwrap();
|
||||
assert!(error.contains("Command not found: npx"));
|
||||
assert!(error.contains("Install Node.js"));
|
||||
assert!(error.contains("absolute command path"));
|
||||
assert_eq!(result.code, Some(McpConnectionTestErrorCode::CommandNotFound));
|
||||
assert_eq!(result.details.as_ref().unwrap()["runtime"], "node");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn spawn_error_not_found_generic_command() {
|
||||
let err = std::io::Error::new(std::io::ErrorKind::NotFound, "not found");
|
||||
let result = spawn_error_result("missing-mcp", &err);
|
||||
let error = result.error.as_deref().unwrap();
|
||||
assert!(error.contains("Command not found: missing-mcp"));
|
||||
assert!(error.contains("Install the command"));
|
||||
assert!(error.contains("absolute command path"));
|
||||
assert_eq!(result.details.as_ref().unwrap()["runtime"], "generic");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn spawn_error_not_found_bun_command() {
|
||||
let err = std::io::Error::new(std::io::ErrorKind::NotFound, "not found");
|
||||
let result = spawn_error_result("bunx", &err);
|
||||
let error = result.error.as_deref().unwrap();
|
||||
assert!(error.contains("Command not found: bunx"));
|
||||
assert!(error.contains("Install Bun"));
|
||||
assert_eq!(result.details.as_ref().unwrap()["runtime"], "bun");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn spawn_error_not_found_uv_command() {
|
||||
let err = std::io::Error::new(std::io::ErrorKind::NotFound, "not found");
|
||||
let result = spawn_error_result("uvx", &err);
|
||||
let error = result.error.as_deref().unwrap();
|
||||
assert!(error.contains("Command not found: uvx"));
|
||||
assert!(error.contains("Install uv"));
|
||||
assert_eq!(result.details.as_ref().unwrap()["runtime"], "uv");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn spawn_error_not_found_python_command() {
|
||||
let err = std::io::Error::new(std::io::ErrorKind::NotFound, "not found");
|
||||
let result = spawn_error_result("python3", &err);
|
||||
let error = result.error.as_deref().unwrap();
|
||||
assert!(error.contains("Command not found: python3"));
|
||||
assert!(error.contains("Install Python"));
|
||||
assert_eq!(result.details.as_ref().unwrap()["runtime"], "python");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn spawn_error_not_found_deno_command() {
|
||||
let err = std::io::Error::new(std::io::ErrorKind::NotFound, "not found");
|
||||
let result = spawn_error_result("deno", &err);
|
||||
let error = result.error.as_deref().unwrap();
|
||||
assert!(error.contains("Command not found: deno"));
|
||||
assert!(error.contains("Install Deno"));
|
||||
assert_eq!(result.details.as_ref().unwrap()["runtime"], "deno");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn spawn_error_permission_denied() {
|
||||
let err = std::io::Error::new(std::io::ErrorKind::PermissionDenied, "denied");
|
||||
let result = spawn_error_result("./script.sh", &err);
|
||||
assert!(result.error.as_deref().unwrap().contains("Permission denied"));
|
||||
assert_eq!(result.code, Some(McpConnectionTestErrorCode::CommandPermissionDenied));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn spawn_error_other() {
|
||||
let err = std::io::Error::other("broken pipe");
|
||||
let result = spawn_error_result("cmd", &err);
|
||||
assert!(result.error.as_deref().unwrap().contains("Failed to start"));
|
||||
assert_eq!(result.code, Some(McpConnectionTestErrorCode::CommandStartFailed));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auth_result_with_bearer() {
|
||||
let mut headers = reqwest::header::HeaderMap::new();
|
||||
headers.insert("www-authenticate", "Bearer realm=\"mcp\"".parse().unwrap());
|
||||
let result = auth_result(&headers);
|
||||
assert!(!result.success);
|
||||
assert_eq!(result.needs_auth, Some(true));
|
||||
assert!(matches!(result.auth_method, Some(McpAuthMethod::Oauth)));
|
||||
assert!(result.www_authenticate.is_some());
|
||||
assert!(result.code.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auth_result_without_www_authenticate() {
|
||||
let headers = reqwest::header::HeaderMap::new();
|
||||
let result = auth_result(&headers);
|
||||
assert_eq!(result.needs_auth, Some(true));
|
||||
assert!(result.auth_method.is_none());
|
||||
assert!(result.www_authenticate.is_none());
|
||||
}
|
||||
|
||||
// -- JSON-RPC builders ------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn initialize_request_structure() {
|
||||
let req = build_initialize_request(1);
|
||||
assert_eq!(req.jsonrpc, "2.0");
|
||||
assert_eq!(req.id, 1);
|
||||
assert_eq!(req.method, "initialize");
|
||||
let params = req.params.unwrap();
|
||||
assert_eq!(params["protocolVersion"], PROTOCOL_VERSION);
|
||||
assert_eq!(params["clientInfo"]["name"], CLIENT_NAME);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn initialized_notification_structure() {
|
||||
let n = build_initialized_notification();
|
||||
assert_eq!(n.jsonrpc, "2.0");
|
||||
assert_eq!(n.method, "notifications/initialized");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tools_list_request_structure() {
|
||||
let req = build_tools_list_request(2);
|
||||
assert_eq!(req.id, 2);
|
||||
assert_eq!(req.method, "tools/list");
|
||||
assert!(req.params.is_none());
|
||||
}
|
||||
|
||||
// -- HTTP header builder ----------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn build_headers_from_map() {
|
||||
let mut map = HashMap::new();
|
||||
map.insert("Authorization".into(), "Bearer tok".into());
|
||||
map.insert("X-Custom".into(), "val".into());
|
||||
let headers = build_http_headers(&map);
|
||||
assert_eq!(headers.get("authorization").unwrap().to_str().unwrap(), "Bearer tok");
|
||||
assert_eq!(headers.get("x-custom").unwrap().to_str().unwrap(), "val");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_headers_empty() {
|
||||
let headers = build_http_headers(&HashMap::new());
|
||||
assert!(headers.is_empty());
|
||||
}
|
||||
|
||||
// -- extract_jsonrpc_from_sse -----------------------------------------
|
||||
|
||||
#[test]
|
||||
fn extract_jsonrpc_from_sse_basic() {
|
||||
let body = "event: message\ndata: {\"jsonrpc\":\"2.0\",\"id\":1,\"result\":{}}\n";
|
||||
let resp = extract_jsonrpc_from_sse(body).unwrap();
|
||||
assert_eq!(resp.id, Some(1));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn extract_jsonrpc_skips_notifications() {
|
||||
let body =
|
||||
"data: {\"jsonrpc\":\"2.0\",\"method\":\"log\"}\ndata: {\"jsonrpc\":\"2.0\",\"id\":1,\"result\":{}}\n";
|
||||
let resp = extract_jsonrpc_from_sse(body).unwrap();
|
||||
assert_eq!(resp.id, Some(1));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn extract_jsonrpc_from_sse_no_response() {
|
||||
let body = "data: not json\ndata: {\"jsonrpc\":\"2.0\",\"method\":\"log\"}\n";
|
||||
assert!(extract_jsonrpc_from_sse(body).is_err());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,131 @@
|
||||
use nomifun_common::AppError;
|
||||
|
||||
/// MCP crate-level errors.
|
||||
///
|
||||
/// Uses `thiserror` (library crate convention).
|
||||
/// Converts to `AppError` for HTTP response mapping.
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum McpError {
|
||||
#[error("MCP server not found: {0}")]
|
||||
NotFound(String),
|
||||
|
||||
#[error("MCP server name conflict: {0}")]
|
||||
Conflict(String),
|
||||
|
||||
#[error("Invalid MCP server edit: {0}")]
|
||||
InvalidEdit(String),
|
||||
|
||||
#[error("Invalid transport configuration: {0}")]
|
||||
InvalidTransport(String),
|
||||
|
||||
#[error("Agent CLI not installed: {0}")]
|
||||
AgentNotInstalled(String),
|
||||
|
||||
#[error("Agent operation failed: {0}")]
|
||||
AgentOperationFailed(String),
|
||||
|
||||
#[error("Connection test failed: {0}")]
|
||||
ConnectionFailed(String),
|
||||
|
||||
#[error("OAuth error: {0}")]
|
||||
OAuth(String),
|
||||
|
||||
#[error("{0}")]
|
||||
Database(#[from] nomifun_db::DbError),
|
||||
|
||||
#[error("JSON error: {0}")]
|
||||
Json(#[from] serde_json::Error),
|
||||
}
|
||||
|
||||
impl From<McpError> for AppError {
|
||||
fn from(err: McpError) -> Self {
|
||||
match err {
|
||||
McpError::NotFound(msg) => AppError::NotFound(msg),
|
||||
McpError::Conflict(msg) => AppError::Conflict(msg),
|
||||
McpError::InvalidEdit(msg) => AppError::BadRequest(msg),
|
||||
McpError::InvalidTransport(msg) => AppError::BadRequest(msg),
|
||||
McpError::AgentNotInstalled(msg) => AppError::BadRequest(msg),
|
||||
McpError::AgentOperationFailed(msg) => AppError::Internal(msg),
|
||||
McpError::ConnectionFailed(msg) => AppError::BadGateway(msg),
|
||||
McpError::OAuth(msg) => AppError::Internal(format!("OAuth error: {msg}")),
|
||||
McpError::Database(db_err) => AppError::from(db_err),
|
||||
McpError::Json(e) => AppError::Internal(format!("JSON error: {e}")),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn not_found_maps_to_app_not_found() {
|
||||
let err: AppError = McpError::NotFound("mcp_123".into()).into();
|
||||
assert!(matches!(err, AppError::NotFound(msg) if msg == "mcp_123"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn conflict_maps_to_app_conflict() {
|
||||
let err: AppError = McpError::Conflict("test-server".into()).into();
|
||||
assert!(matches!(err, AppError::Conflict(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_transport_maps_to_bad_request() {
|
||||
let err: AppError = McpError::InvalidTransport("missing command".into()).into();
|
||||
assert!(matches!(err, AppError::BadRequest(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_edit_maps_to_bad_request() {
|
||||
let err: AppError = McpError::InvalidEdit("rename forbidden".into()).into();
|
||||
assert!(matches!(err, AppError::BadRequest(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn agent_not_installed_maps_to_bad_request() {
|
||||
let err: AppError = McpError::AgentNotInstalled("claude".into()).into();
|
||||
assert!(matches!(err, AppError::BadRequest(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn agent_operation_failed_maps_to_internal() {
|
||||
let err: AppError = McpError::AgentOperationFailed("exit code 1".into()).into();
|
||||
assert!(matches!(err, AppError::Internal(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn connection_failed_maps_to_bad_gateway() {
|
||||
let err: AppError = McpError::ConnectionFailed("timeout".into()).into();
|
||||
assert!(matches!(err, AppError::BadGateway(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn oauth_maps_to_internal() {
|
||||
let err: AppError = McpError::OAuth("discovery failed".into()).into();
|
||||
assert!(matches!(err, AppError::Internal(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn json_error_maps_to_internal() {
|
||||
let json_err = serde_json::from_str::<serde_json::Value>("invalid").unwrap_err();
|
||||
let err: AppError = McpError::Json(json_err).into();
|
||||
assert!(matches!(err, AppError::Internal(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn display_messages() {
|
||||
assert_eq!(
|
||||
McpError::NotFound("mcp_1".into()).to_string(),
|
||||
"MCP server not found: mcp_1"
|
||||
);
|
||||
assert_eq!(
|
||||
McpError::InvalidTransport("bad".into()).to_string(),
|
||||
"Invalid transport configuration: bad"
|
||||
);
|
||||
assert_eq!(
|
||||
McpError::InvalidEdit("rename forbidden".into()).to_string(),
|
||||
"Invalid MCP server edit: rename forbidden"
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
//! MCP server configuration, multi-agent sync adapters, OAuth, and connection testing.
|
||||
pub mod adapter;
|
||||
pub mod adapters;
|
||||
pub mod connection_test;
|
||||
pub mod error;
|
||||
pub mod oauth_service;
|
||||
pub mod routes;
|
||||
pub mod service;
|
||||
pub mod session_injection;
|
||||
pub mod sync_service;
|
||||
pub mod types;
|
||||
|
||||
pub use adapter::{DetectedServer, McpAgentAdapter};
|
||||
pub use adapters::{
|
||||
ClaudeAdapter, CodeBuddyAdapter, CodexAdapter, GeminiAdapter, NomiAdapter, NomifunAdapter, OpencodeAdapter,
|
||||
QwenAdapter,
|
||||
};
|
||||
pub use connection_test::McpConnectionTestService;
|
||||
pub use error::McpError;
|
||||
pub use oauth_service::McpOAuthService;
|
||||
pub use routes::{McpRouterState, mcp_routes};
|
||||
pub use service::McpConfigService;
|
||||
pub use session_injection::{
|
||||
AcpMcpCapabilities, AcpSessionMcpServer, ImageGenConfig, NameValuePair, build_builtin_image_gen_server,
|
||||
build_session_mcp_servers, parse_acp_mcp_capabilities,
|
||||
};
|
||||
pub use sync_service::McpSyncService;
|
||||
pub use types::{McpServer, McpServerTransport, McpTool};
|
||||
@@ -0,0 +1,875 @@
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use nomifun_api_types::{OAuthLoginResponse, OAuthStatusResponse};
|
||||
use nomifun_common::{TimestampMs, now_ms};
|
||||
use nomifun_db::{IOAuthTokenRepository, UpsertOAuthTokenParams};
|
||||
use oauth2::basic::BasicClient;
|
||||
use oauth2::{
|
||||
AuthUrl, AuthorizationCode, ClientId, CsrfToken, PkceCodeChallenge, PkceCodeVerifier, RedirectUrl, RefreshToken,
|
||||
TokenResponse, TokenUrl,
|
||||
};
|
||||
use serde::Deserialize;
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tokio::net::TcpListener;
|
||||
use tokio::sync::Mutex;
|
||||
use tracing::{debug, warn};
|
||||
|
||||
use crate::error::McpError;
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Constants
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Default timeout for the OAuth callback server waiting for the redirect.
|
||||
const CALLBACK_TIMEOUT: Duration = Duration::from_secs(120);
|
||||
|
||||
/// Default OAuth client ID for MCP servers (public client, no secret).
|
||||
const DEFAULT_CLIENT_ID: &str = "nomifun";
|
||||
|
||||
/// Token expiry safety margin (refresh 5 minutes before expiration).
|
||||
const EXPIRY_MARGIN_MS: i64 = 5 * 60 * 1000;
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Discovery response
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// OAuth Authorization Server Metadata (RFC 8414) — subset of fields we need.
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct OAuthServerMetadata {
|
||||
authorization_endpoint: String,
|
||||
token_endpoint: String,
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Pending login state
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// State held while waiting for the OAuth callback redirect.
|
||||
///
|
||||
/// Stores endpoint URLs rather than the typed `BasicClient` to avoid
|
||||
/// complex generic type parameters from the `oauth2` crate.
|
||||
struct PendingLogin {
|
||||
csrf_token: CsrfToken,
|
||||
pkce_verifier: PkceCodeVerifier,
|
||||
auth_url: String,
|
||||
token_url: String,
|
||||
redirect_url: String,
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// McpOAuthService
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Service for MCP server OAuth 2.0 PKCE authentication.
|
||||
///
|
||||
/// Manages the full lifecycle: discovery → authorize → callback → token
|
||||
/// exchange → storage → refresh → logout.
|
||||
#[derive(Clone)]
|
||||
pub struct McpOAuthService {
|
||||
token_repo: Arc<dyn IOAuthTokenRepository>,
|
||||
http_client: reqwest::Client,
|
||||
/// Mutex protecting the pending login state (only one login at a time).
|
||||
pending: Arc<Mutex<Option<PendingLogin>>>,
|
||||
}
|
||||
|
||||
impl McpOAuthService {
|
||||
pub fn new(token_repo: Arc<dyn IOAuthTokenRepository>, http_client: reqwest::Client) -> Self {
|
||||
Self {
|
||||
token_repo,
|
||||
http_client,
|
||||
pending: Arc::new(Mutex::new(None)),
|
||||
}
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Public API
|
||||
// -----------------------------------------------------------------------
|
||||
|
||||
/// Check whether the given server URL has a valid (non-expired) OAuth token.
|
||||
pub async fn check_oauth_status(&self, server_url: &str) -> Result<OAuthStatusResponse, McpError> {
|
||||
let authenticated = self.has_valid_token(server_url).await?;
|
||||
Ok(OAuthStatusResponse { authenticated })
|
||||
}
|
||||
|
||||
/// Start the OAuth PKCE login flow for the given MCP server URL.
|
||||
///
|
||||
/// 1. Discover authorization/token endpoints
|
||||
/// 2. Generate PKCE challenge
|
||||
/// 3. Start local callback server on a random port
|
||||
/// 4. Build authorization URL and open it in the system browser
|
||||
/// 5. Wait for the redirect with the authorization code
|
||||
/// 6. Exchange code for tokens and persist them
|
||||
pub async fn login(&self, server_url: &str) -> Result<OAuthLoginResponse, McpError> {
|
||||
let (authorize_url, listener) = self.prepare_login_flow(server_url).await?;
|
||||
|
||||
// Open browser.
|
||||
debug!(url = %authorize_url, "Opening browser for OAuth authorization");
|
||||
if let Err(e) = open::that(&authorize_url) {
|
||||
warn!("Failed to open browser: {e}");
|
||||
}
|
||||
|
||||
// Wait for callback.
|
||||
let code = match self.wait_for_callback(listener).await {
|
||||
Ok(code) => code,
|
||||
Err(e) => {
|
||||
self.clear_pending().await;
|
||||
return Ok(OAuthLoginResponse {
|
||||
success: false,
|
||||
error: Some(e.to_string()),
|
||||
});
|
||||
}
|
||||
};
|
||||
|
||||
// Exchange code for tokens.
|
||||
match self.exchange_code(server_url, code).await {
|
||||
Ok(()) => Ok(OAuthLoginResponse {
|
||||
success: true,
|
||||
error: None,
|
||||
}),
|
||||
Err(e) => {
|
||||
self.clear_pending().await;
|
||||
Ok(OAuthLoginResponse {
|
||||
success: false,
|
||||
error: Some(e.to_string()),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Logout from the given MCP server URL (delete stored token).
|
||||
///
|
||||
/// Idempotent: returns Ok even if no token was stored.
|
||||
pub async fn logout(&self, server_url: &str) -> Result<(), McpError> {
|
||||
match self.token_repo.delete(server_url).await {
|
||||
Ok(()) => {
|
||||
debug!(server_url, "OAuth token deleted");
|
||||
Ok(())
|
||||
}
|
||||
Err(nomifun_db::DbError::NotFound(_)) => {
|
||||
debug!(server_url, "No OAuth token to delete (idempotent)");
|
||||
Ok(())
|
||||
}
|
||||
Err(e) => Err(McpError::Database(e)),
|
||||
}
|
||||
}
|
||||
|
||||
/// Return the list of server URLs that have stored OAuth tokens.
|
||||
pub async fn get_authenticated_servers(&self) -> Result<Vec<String>, McpError> {
|
||||
let urls = self.token_repo.list_authenticated_urls().await?;
|
||||
Ok(urls)
|
||||
}
|
||||
|
||||
/// Get a valid access token for the given server URL.
|
||||
///
|
||||
/// If the stored token is expired and a refresh token is available,
|
||||
/// automatically refreshes before returning.
|
||||
/// Returns `None` if no token is stored for this URL.
|
||||
pub async fn get_token(&self, server_url: &str) -> Result<Option<String>, McpError> {
|
||||
let row = match self.token_repo.get_by_url(server_url).await? {
|
||||
Some(row) => row,
|
||||
None => return Ok(None),
|
||||
};
|
||||
|
||||
// Check if token is expired (with safety margin).
|
||||
if let Some(expires_at) = row.expires_at {
|
||||
let now = now_ms();
|
||||
if now >= expires_at - EXPIRY_MARGIN_MS
|
||||
&& let Some(ref refresh_token) = row.refresh_token
|
||||
{
|
||||
match self.refresh_token(server_url, refresh_token).await {
|
||||
Ok(new_token) => return Ok(Some(new_token)),
|
||||
Err(e) => {
|
||||
warn!(
|
||||
server_url,
|
||||
error = %e,
|
||||
"Token refresh failed, returning expired token"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(Some(row.access_token))
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Internal helpers
|
||||
// -----------------------------------------------------------------------
|
||||
|
||||
/// Discover endpoints, build OAuth client, generate PKCE, bind callback
|
||||
/// server, store pending state, and return the authorization URL + listener.
|
||||
async fn prepare_login_flow(&self, server_url: &str) -> Result<(String, TcpListener), McpError> {
|
||||
let metadata = self.discover_endpoints(server_url).await?;
|
||||
|
||||
let auth_url_str = metadata.authorization_endpoint.clone();
|
||||
let token_url_str = metadata.token_endpoint.clone();
|
||||
|
||||
let auth_url = AuthUrl::new(metadata.authorization_endpoint)
|
||||
.map_err(|e| McpError::OAuth(format!("Invalid auth URL: {e}")))?;
|
||||
let token_url =
|
||||
TokenUrl::new(metadata.token_endpoint).map_err(|e| McpError::OAuth(format!("Invalid token URL: {e}")))?;
|
||||
|
||||
let listener = TcpListener::bind("127.0.0.1:0")
|
||||
.await
|
||||
.map_err(|e| McpError::OAuth(format!("Failed to bind callback server: {e}")))?;
|
||||
let callback_port = listener
|
||||
.local_addr()
|
||||
.map_err(|e| McpError::OAuth(format!("Failed to get callback port: {e}")))?
|
||||
.port();
|
||||
|
||||
let redirect_url_str = format!("http://127.0.0.1:{callback_port}/callback");
|
||||
let redirect = RedirectUrl::new(redirect_url_str.clone())
|
||||
.map_err(|e| McpError::OAuth(format!("Invalid redirect URL: {e}")))?;
|
||||
|
||||
let client = BasicClient::new(ClientId::new(DEFAULT_CLIENT_ID.to_string()))
|
||||
.set_auth_uri(auth_url)
|
||||
.set_token_uri(token_url)
|
||||
.set_redirect_uri(redirect);
|
||||
|
||||
let (pkce_challenge, pkce_verifier) = PkceCodeChallenge::new_random_sha256();
|
||||
|
||||
let (authorize_url, csrf_token) = client
|
||||
.authorize_url(CsrfToken::new_random)
|
||||
.set_pkce_challenge(pkce_challenge)
|
||||
.url();
|
||||
|
||||
{
|
||||
let mut pending = self.pending.lock().await;
|
||||
*pending = Some(PendingLogin {
|
||||
csrf_token,
|
||||
pkce_verifier,
|
||||
auth_url: auth_url_str,
|
||||
token_url: token_url_str,
|
||||
redirect_url: redirect_url_str,
|
||||
});
|
||||
}
|
||||
|
||||
Ok((authorize_url.to_string(), listener))
|
||||
}
|
||||
|
||||
/// Check if a valid (non-expired) token exists for the URL.
|
||||
async fn has_valid_token(&self, server_url: &str) -> Result<bool, McpError> {
|
||||
let row = match self.token_repo.get_by_url(server_url).await? {
|
||||
Some(row) => row,
|
||||
None => return Ok(false),
|
||||
};
|
||||
|
||||
if let Some(expires_at) = row.expires_at
|
||||
&& now_ms() >= expires_at
|
||||
{
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
/// Discover OAuth authorization server metadata.
|
||||
///
|
||||
/// Tries `.well-known/oauth-authorization-server` first,
|
||||
/// falls back to `.well-known/openid-configuration`.
|
||||
async fn discover_endpoints(&self, server_url: &str) -> Result<OAuthServerMetadata, McpError> {
|
||||
let base = server_url.trim_end_matches('/');
|
||||
|
||||
let well_known_url = format!("{base}/.well-known/oauth-authorization-server");
|
||||
if let Ok(metadata) = self.fetch_metadata(&well_known_url).await {
|
||||
debug!(server_url, "Discovered OAuth metadata via RFC 8414");
|
||||
return Ok(metadata);
|
||||
}
|
||||
|
||||
let oidc_url = format!("{base}/.well-known/openid-configuration");
|
||||
if let Ok(metadata) = self.fetch_metadata(&oidc_url).await {
|
||||
debug!(server_url, "Discovered OAuth metadata via OIDC");
|
||||
return Ok(metadata);
|
||||
}
|
||||
|
||||
Err(McpError::OAuth(format!(
|
||||
"Failed to discover OAuth endpoints for '{server_url}': \
|
||||
no .well-known/oauth-authorization-server or \
|
||||
.well-known/openid-configuration found"
|
||||
)))
|
||||
}
|
||||
|
||||
/// Fetch and parse OAuth server metadata from a URL.
|
||||
async fn fetch_metadata(&self, url: &str) -> Result<OAuthServerMetadata, McpError> {
|
||||
let resp = self
|
||||
.http_client
|
||||
.get(url)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| McpError::OAuth(format!("HTTP request failed: {e}")))?;
|
||||
|
||||
if !resp.status().is_success() {
|
||||
return Err(McpError::OAuth(format!("Metadata endpoint returned {}", resp.status())));
|
||||
}
|
||||
|
||||
resp.json()
|
||||
.await
|
||||
.map_err(|e| McpError::OAuth(format!("Failed to parse metadata: {e}")))
|
||||
}
|
||||
|
||||
/// Wait for the OAuth callback redirect on the given listener.
|
||||
async fn wait_for_callback(&self, listener: TcpListener) -> Result<String, McpError> {
|
||||
let (code_tx, code_rx) = tokio::sync::oneshot::channel::<Result<String, McpError>>();
|
||||
let pending = self.pending.clone();
|
||||
|
||||
tokio::spawn(async move {
|
||||
let result = Self::handle_callback_connection(listener, pending).await;
|
||||
let _ = code_tx.send(result);
|
||||
});
|
||||
|
||||
match tokio::time::timeout(CALLBACK_TIMEOUT, code_rx).await {
|
||||
Ok(Ok(result)) => result,
|
||||
Ok(Err(_)) => Err(McpError::OAuth("Callback channel closed unexpectedly".to_string())),
|
||||
Err(_) => Err(McpError::OAuth(
|
||||
"OAuth callback timed out — no redirect received within 120s".to_string(),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
/// Handle a single HTTP connection on the callback server.
|
||||
async fn handle_callback_connection(
|
||||
listener: TcpListener,
|
||||
pending: Arc<Mutex<Option<PendingLogin>>>,
|
||||
) -> Result<String, McpError> {
|
||||
let (mut stream, _) = listener
|
||||
.accept()
|
||||
.await
|
||||
.map_err(|e| McpError::OAuth(format!("Failed to accept connection: {e}")))?;
|
||||
|
||||
let mut buf = vec![0u8; 4096];
|
||||
let n = stream
|
||||
.read(&mut buf)
|
||||
.await
|
||||
.map_err(|e| McpError::OAuth(format!("Failed to read request: {e}")))?;
|
||||
|
||||
let request = String::from_utf8_lossy(&buf[..n]);
|
||||
let (code, state) = parse_callback_query(&request)?;
|
||||
|
||||
// Validate CSRF state.
|
||||
let guard = pending.lock().await;
|
||||
let pending_login = guard
|
||||
.as_ref()
|
||||
.ok_or_else(|| McpError::OAuth("No pending login state".to_string()))?;
|
||||
|
||||
if state != *pending_login.csrf_token.secret() {
|
||||
return Err(McpError::OAuth("CSRF state mismatch".to_string()));
|
||||
}
|
||||
|
||||
// Send a success response to the browser.
|
||||
let response = "HTTP/1.1 200 OK\r\n\
|
||||
Content-Type: text/html; charset=utf-8\r\n\
|
||||
Connection: close\r\n\r\n\
|
||||
<html><body><h1>Authorization successful!</h1>\
|
||||
<p>You can close this window and return to Nomi.</p>\
|
||||
</body></html>";
|
||||
|
||||
let _ = stream.write_all(response.as_bytes()).await;
|
||||
|
||||
Ok(code)
|
||||
}
|
||||
|
||||
/// Build a no-redirect reqwest client for OAuth token exchange.
|
||||
fn build_no_redirect_client() -> Result<reqwest::Client, McpError> {
|
||||
nomifun_net::proxy::apply_detected_proxy(reqwest::ClientBuilder::new())
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.map_err(|e| McpError::OAuth(format!("Failed to build HTTP client: {e}")))
|
||||
}
|
||||
|
||||
/// Exchange the authorization code for tokens and persist them.
|
||||
async fn exchange_code(&self, server_url: &str, code: String) -> Result<(), McpError> {
|
||||
let (auth_url_str, token_url_str, redirect_url_str, pkce_verifier) = {
|
||||
let mut guard = self.pending.lock().await;
|
||||
let pending = guard
|
||||
.take()
|
||||
.ok_or_else(|| McpError::OAuth("No pending login state".to_string()))?;
|
||||
(
|
||||
pending.auth_url,
|
||||
pending.token_url,
|
||||
pending.redirect_url,
|
||||
pending.pkce_verifier,
|
||||
)
|
||||
};
|
||||
|
||||
let auth_url = AuthUrl::new(auth_url_str).map_err(|e| McpError::OAuth(format!("Invalid auth URL: {e}")))?;
|
||||
let token_url = TokenUrl::new(token_url_str).map_err(|e| McpError::OAuth(format!("Invalid token URL: {e}")))?;
|
||||
let redirect =
|
||||
RedirectUrl::new(redirect_url_str).map_err(|e| McpError::OAuth(format!("Invalid redirect URL: {e}")))?;
|
||||
|
||||
let client = BasicClient::new(ClientId::new(DEFAULT_CLIENT_ID.to_string()))
|
||||
.set_auth_uri(auth_url)
|
||||
.set_token_uri(token_url)
|
||||
.set_redirect_uri(redirect);
|
||||
|
||||
let http_client = Self::build_no_redirect_client()?;
|
||||
|
||||
let token_result = client
|
||||
.exchange_code(AuthorizationCode::new(code))
|
||||
.set_pkce_verifier(pkce_verifier)
|
||||
.request_async(&http_client)
|
||||
.await
|
||||
.map_err(|e| McpError::OAuth(format!("Token exchange failed: {e}")))?;
|
||||
|
||||
self.persist_token(server_url, &token_result).await?;
|
||||
debug!(server_url, "OAuth tokens stored successfully");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Refresh an expired access token using the refresh token.
|
||||
async fn refresh_token(&self, server_url: &str, refresh_token_value: &str) -> Result<String, McpError> {
|
||||
let metadata = self.discover_endpoints(server_url).await?;
|
||||
let token_url =
|
||||
TokenUrl::new(metadata.token_endpoint).map_err(|e| McpError::OAuth(format!("Invalid token URL: {e}")))?;
|
||||
|
||||
let client = BasicClient::new(ClientId::new(DEFAULT_CLIENT_ID.to_string())).set_token_uri(token_url);
|
||||
|
||||
let http_client = Self::build_no_redirect_client()?;
|
||||
|
||||
let refresh_token = RefreshToken::new(refresh_token_value.to_string());
|
||||
let token_result = client
|
||||
.exchange_refresh_token(&refresh_token)
|
||||
.request_async(&http_client)
|
||||
.await
|
||||
.map_err(|e| McpError::OAuth(format!("Token refresh failed: {e}")))?;
|
||||
|
||||
let new_access_token = token_result.access_token().secret().clone();
|
||||
|
||||
let expires_at: Option<TimestampMs> = token_result.expires_in().map(|d| now_ms() + d.as_millis() as i64);
|
||||
|
||||
// Prefer new refresh_token if provided, otherwise keep the old one.
|
||||
let new_refresh = token_result
|
||||
.refresh_token()
|
||||
.map(|t| t.secret().as_str())
|
||||
.unwrap_or(refresh_token_value);
|
||||
|
||||
self.token_repo
|
||||
.upsert(UpsertOAuthTokenParams {
|
||||
server_url,
|
||||
access_token: &new_access_token,
|
||||
refresh_token: Some(new_refresh),
|
||||
token_type: "bearer",
|
||||
expires_at,
|
||||
})
|
||||
.await?;
|
||||
|
||||
debug!(server_url, "OAuth token refreshed successfully");
|
||||
Ok(new_access_token)
|
||||
}
|
||||
|
||||
/// Persist token response to DB.
|
||||
async fn persist_token<TR: TokenResponse>(&self, server_url: &str, token_result: &TR) -> Result<(), McpError> {
|
||||
let expires_at: Option<TimestampMs> = token_result.expires_in().map(|d| now_ms() + d.as_millis() as i64);
|
||||
|
||||
self.token_repo
|
||||
.upsert(UpsertOAuthTokenParams {
|
||||
server_url,
|
||||
access_token: token_result.access_token().secret(),
|
||||
refresh_token: token_result.refresh_token().map(|t| t.secret().as_str()),
|
||||
token_type: "bearer",
|
||||
expires_at,
|
||||
})
|
||||
.await?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Clear the pending login state.
|
||||
async fn clear_pending(&self) {
|
||||
let mut guard = self.pending.lock().await;
|
||||
*guard = None;
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Query parameter parsing
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Parse `code` and `state` from the first line of an HTTP request.
|
||||
///
|
||||
/// Expects: `GET /callback?code=xxx&state=yyy HTTP/1.1`
|
||||
fn parse_callback_query(request: &str) -> Result<(String, String), McpError> {
|
||||
let first_line = request
|
||||
.lines()
|
||||
.next()
|
||||
.ok_or_else(|| McpError::OAuth("Empty HTTP request".to_string()))?;
|
||||
|
||||
let path = first_line
|
||||
.split_whitespace()
|
||||
.nth(1)
|
||||
.ok_or_else(|| McpError::OAuth("Malformed HTTP request line".to_string()))?;
|
||||
|
||||
let query_str = path
|
||||
.split_once('?')
|
||||
.map(|(_, q)| q)
|
||||
.ok_or_else(|| McpError::OAuth("No query parameters in callback".to_string()))?;
|
||||
|
||||
let mut code = None;
|
||||
let mut state = None;
|
||||
|
||||
for pair in query_str.split('&') {
|
||||
if let Some((key, value)) = pair.split_once('=') {
|
||||
match key {
|
||||
"code" => code = Some(url_decode(value)),
|
||||
"state" => state = Some(url_decode(value)),
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let code = code.ok_or_else(|| McpError::OAuth("Missing 'code' in callback".to_string()))?;
|
||||
let state = state.ok_or_else(|| McpError::OAuth("Missing 'state' in callback".to_string()))?;
|
||||
|
||||
Ok((code, state))
|
||||
}
|
||||
|
||||
/// Minimal percent-decoding for query parameter values.
|
||||
fn url_decode(input: &str) -> String {
|
||||
let mut result = String::with_capacity(input.len());
|
||||
let mut chars = input.bytes();
|
||||
|
||||
while let Some(b) = chars.next() {
|
||||
if b == b'%' {
|
||||
let hi = chars.next();
|
||||
let lo = chars.next();
|
||||
if let (Some(h), Some(l)) = (hi, lo) {
|
||||
let hex = [h, l];
|
||||
if let Ok(s) = std::str::from_utf8(&hex)
|
||||
&& let Ok(byte) = u8::from_str_radix(s, 16)
|
||||
{
|
||||
result.push(byte as char);
|
||||
continue;
|
||||
}
|
||||
// Malformed percent-encoding: keep as-is.
|
||||
result.push('%');
|
||||
result.push(h as char);
|
||||
result.push(l as char);
|
||||
}
|
||||
} else if b == b'+' {
|
||||
result.push(' ');
|
||||
} else {
|
||||
result.push(b as char);
|
||||
}
|
||||
}
|
||||
|
||||
result
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Unit tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
// -- parse_callback_query ------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn parse_valid_callback_query() {
|
||||
let request = "GET /callback?code=abc123&state=xyz789 HTTP/1.1\r\nHost: localhost\r\n";
|
||||
let (code, state) = parse_callback_query(request).unwrap();
|
||||
assert_eq!(code, "abc123");
|
||||
assert_eq!(state, "xyz789");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_callback_query_reversed_params() {
|
||||
let request = "GET /callback?state=s1&code=c1 HTTP/1.1\r\n";
|
||||
let (code, state) = parse_callback_query(request).unwrap();
|
||||
assert_eq!(code, "c1");
|
||||
assert_eq!(state, "s1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_callback_query_with_extra_params() {
|
||||
let request = "GET /callback?code=c&foo=bar&state=s HTTP/1.1\r\n";
|
||||
let (code, state) = parse_callback_query(request).unwrap();
|
||||
assert_eq!(code, "c");
|
||||
assert_eq!(state, "s");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_callback_query_missing_code() {
|
||||
let request = "GET /callback?state=s HTTP/1.1\r\n";
|
||||
let err = parse_callback_query(request).unwrap_err();
|
||||
assert!(err.to_string().contains("Missing 'code'"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_callback_query_missing_state() {
|
||||
let request = "GET /callback?code=c HTTP/1.1\r\n";
|
||||
let err = parse_callback_query(request).unwrap_err();
|
||||
assert!(err.to_string().contains("Missing 'state'"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_callback_query_no_query_string() {
|
||||
let request = "GET /callback HTTP/1.1\r\n";
|
||||
let err = parse_callback_query(request).unwrap_err();
|
||||
assert!(err.to_string().contains("No query parameters"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_callback_query_empty_request() {
|
||||
let err = parse_callback_query("").unwrap_err();
|
||||
assert!(err.to_string().contains("Empty HTTP request"));
|
||||
}
|
||||
|
||||
// -- url_decode ----------------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn url_decode_no_encoding() {
|
||||
assert_eq!(url_decode("hello"), "hello");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn url_decode_percent_encoded() {
|
||||
assert_eq!(url_decode("hello%20world"), "hello world");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn url_decode_plus_sign() {
|
||||
assert_eq!(url_decode("hello+world"), "hello world");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn url_decode_special_characters() {
|
||||
assert_eq!(url_decode("%3D%26%3F"), "=&?");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn url_decode_mixed() {
|
||||
assert_eq!(url_decode("a%20b+c%3Dd"), "a b c=d");
|
||||
}
|
||||
|
||||
// -- McpOAuthService construction ----------------------------------------
|
||||
|
||||
#[test]
|
||||
fn service_clone_is_independent() {
|
||||
let repo: Arc<dyn IOAuthTokenRepository> = Arc::new(MockTokenRepo);
|
||||
let http = reqwest::Client::new();
|
||||
let svc = McpOAuthService::new(repo, http);
|
||||
let _clone = svc.clone();
|
||||
}
|
||||
|
||||
// -- Mock repositories ---------------------------------------------------
|
||||
|
||||
struct MockTokenRepo;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl IOAuthTokenRepository for MockTokenRepo {
|
||||
async fn get_by_url(&self, _: &str) -> Result<Option<nomifun_db::models::OAuthTokenRow>, nomifun_db::DbError> {
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn upsert(
|
||||
&self,
|
||||
_: UpsertOAuthTokenParams<'_>,
|
||||
) -> Result<nomifun_db::models::OAuthTokenRow, nomifun_db::DbError> {
|
||||
unimplemented!()
|
||||
}
|
||||
|
||||
async fn delete(&self, _: &str) -> Result<(), nomifun_db::DbError> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn list_authenticated_urls(&self) -> Result<Vec<String>, nomifun_db::DbError> {
|
||||
Ok(vec![])
|
||||
}
|
||||
}
|
||||
|
||||
struct IdempotentDeleteRepo;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl IOAuthTokenRepository for IdempotentDeleteRepo {
|
||||
async fn get_by_url(&self, _: &str) -> Result<Option<nomifun_db::models::OAuthTokenRow>, nomifun_db::DbError> {
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn upsert(
|
||||
&self,
|
||||
_: UpsertOAuthTokenParams<'_>,
|
||||
) -> Result<nomifun_db::models::OAuthTokenRow, nomifun_db::DbError> {
|
||||
unimplemented!()
|
||||
}
|
||||
|
||||
async fn delete(&self, url: &str) -> Result<(), nomifun_db::DbError> {
|
||||
Err(nomifun_db::DbError::NotFound(format!(
|
||||
"OAuth token for '{url}' not found"
|
||||
)))
|
||||
}
|
||||
|
||||
async fn list_authenticated_urls(&self) -> Result<Vec<String>, nomifun_db::DbError> {
|
||||
Ok(vec![])
|
||||
}
|
||||
}
|
||||
|
||||
struct ValidTokenRepo;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl IOAuthTokenRepository for ValidTokenRepo {
|
||||
async fn get_by_url(&self, _: &str) -> Result<Option<nomifun_db::models::OAuthTokenRow>, nomifun_db::DbError> {
|
||||
Ok(Some(nomifun_db::models::OAuthTokenRow {
|
||||
server_url: "https://example.com".to_string(),
|
||||
access_token: "valid_access_token".to_string(),
|
||||
refresh_token: None,
|
||||
token_type: "bearer".to_string(),
|
||||
expires_at: Some(now_ms() + 3_600_000),
|
||||
created_at: now_ms(),
|
||||
updated_at: now_ms(),
|
||||
}))
|
||||
}
|
||||
|
||||
async fn upsert(
|
||||
&self,
|
||||
_: UpsertOAuthTokenParams<'_>,
|
||||
) -> Result<nomifun_db::models::OAuthTokenRow, nomifun_db::DbError> {
|
||||
unimplemented!()
|
||||
}
|
||||
|
||||
async fn delete(&self, _: &str) -> Result<(), nomifun_db::DbError> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn list_authenticated_urls(&self) -> Result<Vec<String>, nomifun_db::DbError> {
|
||||
Ok(vec!["https://example.com".to_string()])
|
||||
}
|
||||
}
|
||||
|
||||
struct ExpiredTokenRepo;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl IOAuthTokenRepository for ExpiredTokenRepo {
|
||||
async fn get_by_url(&self, _: &str) -> Result<Option<nomifun_db::models::OAuthTokenRow>, nomifun_db::DbError> {
|
||||
Ok(Some(nomifun_db::models::OAuthTokenRow {
|
||||
server_url: "https://example.com".to_string(),
|
||||
access_token: "expired_token".to_string(),
|
||||
refresh_token: None,
|
||||
token_type: "bearer".to_string(),
|
||||
expires_at: Some(1000),
|
||||
created_at: 500,
|
||||
updated_at: 500,
|
||||
}))
|
||||
}
|
||||
|
||||
async fn upsert(
|
||||
&self,
|
||||
_: UpsertOAuthTokenParams<'_>,
|
||||
) -> Result<nomifun_db::models::OAuthTokenRow, nomifun_db::DbError> {
|
||||
unimplemented!()
|
||||
}
|
||||
|
||||
async fn delete(&self, _: &str) -> Result<(), nomifun_db::DbError> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn list_authenticated_urls(&self) -> Result<Vec<String>, nomifun_db::DbError> {
|
||||
Ok(vec![])
|
||||
}
|
||||
}
|
||||
|
||||
struct NoExpiryTokenRepo;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl IOAuthTokenRepository for NoExpiryTokenRepo {
|
||||
async fn get_by_url(&self, _: &str) -> Result<Option<nomifun_db::models::OAuthTokenRow>, nomifun_db::DbError> {
|
||||
Ok(Some(nomifun_db::models::OAuthTokenRow {
|
||||
server_url: "https://example.com".to_string(),
|
||||
access_token: "no_expiry_token".to_string(),
|
||||
refresh_token: None,
|
||||
token_type: "bearer".to_string(),
|
||||
expires_at: None,
|
||||
created_at: now_ms(),
|
||||
updated_at: now_ms(),
|
||||
}))
|
||||
}
|
||||
|
||||
async fn upsert(
|
||||
&self,
|
||||
_: UpsertOAuthTokenParams<'_>,
|
||||
) -> Result<nomifun_db::models::OAuthTokenRow, nomifun_db::DbError> {
|
||||
unimplemented!()
|
||||
}
|
||||
|
||||
async fn delete(&self, _: &str) -> Result<(), nomifun_db::DbError> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn list_authenticated_urls(&self) -> Result<Vec<String>, nomifun_db::DbError> {
|
||||
Ok(vec!["https://example.com".to_string()])
|
||||
}
|
||||
}
|
||||
|
||||
// -- Service behavior tests ----------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn check_status_no_token_returns_false() {
|
||||
let svc = McpOAuthService::new(Arc::new(MockTokenRepo), reqwest::Client::new());
|
||||
let status = svc.check_oauth_status("https://example.com").await.unwrap();
|
||||
assert!(!status.authenticated);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn check_status_with_valid_token() {
|
||||
let svc = McpOAuthService::new(Arc::new(ValidTokenRepo), reqwest::Client::new());
|
||||
let status = svc.check_oauth_status("https://example.com").await.unwrap();
|
||||
assert!(status.authenticated);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn check_status_with_expired_token() {
|
||||
let svc = McpOAuthService::new(Arc::new(ExpiredTokenRepo), reqwest::Client::new());
|
||||
let status = svc.check_oauth_status("https://example.com").await.unwrap();
|
||||
assert!(!status.authenticated);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn check_status_no_expiry_treated_as_valid() {
|
||||
let svc = McpOAuthService::new(Arc::new(NoExpiryTokenRepo), reqwest::Client::new());
|
||||
let status = svc.check_oauth_status("https://example.com").await.unwrap();
|
||||
assert!(status.authenticated);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn logout_idempotent_for_nonexistent() {
|
||||
let svc = McpOAuthService::new(Arc::new(IdempotentDeleteRepo), reqwest::Client::new());
|
||||
svc.logout("https://nonexistent.example.com").await.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn get_authenticated_servers_empty() {
|
||||
let svc = McpOAuthService::new(Arc::new(MockTokenRepo), reqwest::Client::new());
|
||||
let urls = svc.get_authenticated_servers().await.unwrap();
|
||||
assert!(urls.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn get_authenticated_servers_returns_urls() {
|
||||
let svc = McpOAuthService::new(Arc::new(ValidTokenRepo), reqwest::Client::new());
|
||||
let urls = svc.get_authenticated_servers().await.unwrap();
|
||||
assert_eq!(urls, vec!["https://example.com"]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn get_token_returns_none_when_no_token() {
|
||||
let svc = McpOAuthService::new(Arc::new(MockTokenRepo), reqwest::Client::new());
|
||||
let token = svc.get_token("https://example.com").await.unwrap();
|
||||
assert!(token.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn get_token_returns_access_token() {
|
||||
let svc = McpOAuthService::new(Arc::new(ValidTokenRepo), reqwest::Client::new());
|
||||
let token = svc.get_token("https://example.com").await.unwrap();
|
||||
assert_eq!(token.as_deref(), Some("valid_access_token"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn get_token_returns_expired_when_no_refresh() {
|
||||
let svc = McpOAuthService::new(Arc::new(ExpiredTokenRepo), reqwest::Client::new());
|
||||
// Expired token with no refresh_token: returns the expired token as-is.
|
||||
let token = svc.get_token("https://example.com").await.unwrap();
|
||||
assert_eq!(token.as_deref(), Some("expired_token"));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,244 @@
|
||||
use axum::Router;
|
||||
use axum::extract::rejection::JsonRejection;
|
||||
use axum::extract::{Json, Path, State};
|
||||
use axum::http::StatusCode;
|
||||
use axum::response::{IntoResponse, Response};
|
||||
use axum::routing::{get, post};
|
||||
|
||||
use nomifun_api_types::{
|
||||
ApiResponse, BatchImportMcpServersRequest, CreateMcpServerRequest, DetectedMcpServerResponse, ErrorResponse,
|
||||
McpConnectionTestErrorCode, McpServerResponse, OAuthCheckStatusRequest, OAuthLoginRequest, OAuthLoginResponse,
|
||||
OAuthLogoutRequest, OAuthStatusResponse, TestMcpConnectionRequest, UpdateMcpServerRequest,
|
||||
};
|
||||
use nomifun_common::AppError;
|
||||
|
||||
use crate::connection_test::McpConnectionTestService;
|
||||
use crate::oauth_service::McpOAuthService;
|
||||
use crate::service::McpConfigService;
|
||||
use crate::sync_service::McpSyncService;
|
||||
use crate::types::McpServerTransport;
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Router state
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Shared state for MCP route handlers.
|
||||
#[derive(Clone)]
|
||||
pub struct McpRouterState {
|
||||
pub config_service: McpConfigService,
|
||||
pub sync_service: McpSyncService,
|
||||
pub connection_test_service: McpConnectionTestService,
|
||||
pub oauth_service: McpOAuthService,
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Router builder
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Build the MCP router with all `/api/mcp/*` routes.
|
||||
///
|
||||
/// Includes CRUD routes, agent config detection, connection tests, and OAuth.
|
||||
/// All routes require authentication (applied by the caller).
|
||||
pub fn mcp_routes(state: McpRouterState) -> Router {
|
||||
Router::new()
|
||||
.route("/api/mcp/servers", get(list_servers).post(add_server))
|
||||
.route("/api/mcp/servers/import", post(batch_import))
|
||||
.route(
|
||||
"/api/mcp/servers/{id}",
|
||||
get(get_server).put(edit_server).delete(delete_server),
|
||||
)
|
||||
.route("/api/mcp/servers/{id}/toggle", post(toggle_server))
|
||||
// Connection test route
|
||||
.route("/api/mcp/test-connection", post(test_connection))
|
||||
// Agent config discovery route
|
||||
.route("/api/mcp/agent-configs", get(get_agent_configs))
|
||||
// OAuth routes
|
||||
.route("/api/mcp/oauth/check-status", post(oauth_check_status))
|
||||
.route("/api/mcp/oauth/login", post(oauth_login))
|
||||
.route("/api/mcp/oauth/logout", post(oauth_logout))
|
||||
.route("/api/mcp/oauth/authenticated", get(oauth_authenticated))
|
||||
.with_state(state)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// CRUD Handlers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// `GET /api/mcp/servers` — list all MCP servers.
|
||||
async fn list_servers(
|
||||
State(state): State<McpRouterState>,
|
||||
) -> Result<Json<ApiResponse<Vec<McpServerResponse>>>, AppError> {
|
||||
let servers = state.config_service.list_servers().await?;
|
||||
Ok(Json(ApiResponse::ok(servers)))
|
||||
}
|
||||
|
||||
/// `GET /api/mcp/servers/:id` — get a single MCP server.
|
||||
async fn get_server(
|
||||
State(state): State<McpRouterState>,
|
||||
Path(id): Path<String>,
|
||||
) -> Result<Json<ApiResponse<McpServerResponse>>, AppError> {
|
||||
let server = state.config_service.get_server(&id).await?;
|
||||
Ok(Json(ApiResponse::ok(server)))
|
||||
}
|
||||
|
||||
/// `POST /api/mcp/servers` — create (or upsert by name) an MCP server.
|
||||
async fn add_server(
|
||||
State(state): State<McpRouterState>,
|
||||
body: Result<Json<CreateMcpServerRequest>, JsonRejection>,
|
||||
) -> Result<(StatusCode, Json<ApiResponse<McpServerResponse>>), AppError> {
|
||||
let Json(req) = body.map_err(|e| AppError::BadRequest(e.to_string()))?;
|
||||
let server = state.config_service.add_server(req).await?;
|
||||
Ok((StatusCode::CREATED, Json(ApiResponse::ok(server))))
|
||||
}
|
||||
|
||||
/// `PUT /api/mcp/servers/:id` — partial update an MCP server.
|
||||
async fn edit_server(
|
||||
State(state): State<McpRouterState>,
|
||||
Path(id): Path<String>,
|
||||
body: Result<Json<UpdateMcpServerRequest>, JsonRejection>,
|
||||
) -> Result<Json<ApiResponse<McpServerResponse>>, AppError> {
|
||||
let Json(req) = body.map_err(|e| AppError::BadRequest(e.to_string()))?;
|
||||
let server = state.config_service.edit_server(&id, req).await?;
|
||||
Ok(Json(ApiResponse::ok(server)))
|
||||
}
|
||||
|
||||
/// `DELETE /api/mcp/servers/:id` — delete an MCP server.
|
||||
async fn delete_server(
|
||||
State(state): State<McpRouterState>,
|
||||
Path(id): Path<String>,
|
||||
) -> Result<Json<ApiResponse<()>>, AppError> {
|
||||
state.config_service.delete_server(&id).await?;
|
||||
Ok(Json(ApiResponse::success()))
|
||||
}
|
||||
|
||||
/// `POST /api/mcp/servers/:id/toggle` — toggle enabled state.
|
||||
async fn toggle_server(
|
||||
State(state): State<McpRouterState>,
|
||||
Path(id): Path<String>,
|
||||
) -> Result<Json<ApiResponse<McpServerResponse>>, AppError> {
|
||||
let server = state.config_service.toggle_server(&id).await?;
|
||||
Ok(Json(ApiResponse::ok(server)))
|
||||
}
|
||||
|
||||
/// `POST /api/mcp/servers/import` — batch import MCP servers.
|
||||
async fn batch_import(
|
||||
State(state): State<McpRouterState>,
|
||||
body: Result<Json<BatchImportMcpServersRequest>, JsonRejection>,
|
||||
) -> Result<Json<ApiResponse<Vec<McpServerResponse>>>, AppError> {
|
||||
let Json(req) = body.map_err(|e| AppError::BadRequest(e.to_string()))?;
|
||||
let servers = state.config_service.batch_import(req).await?;
|
||||
Ok(Json(ApiResponse::ok(servers)))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Connection Test Handler
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// `POST /api/mcp/test-connection` — test MCP server connectivity.
|
||||
///
|
||||
/// Creates a temporary MCP client, connects, lists tools, and closes.
|
||||
async fn test_connection(
|
||||
State(state): State<McpRouterState>,
|
||||
body: Result<Json<TestMcpConnectionRequest>, JsonRejection>,
|
||||
) -> Result<Response, AppError> {
|
||||
let Json(req) = body.map_err(|e| AppError::BadRequest(e.to_string()))?;
|
||||
let transport = McpServerTransport::from(req.transport);
|
||||
let result = state
|
||||
.connection_test_service
|
||||
.test_connection(&req.name, &transport)
|
||||
.await;
|
||||
if let Some(server_id) = req.id.as_deref() {
|
||||
state.config_service.persist_test_result(server_id, &result).await?;
|
||||
}
|
||||
if result.success || result.needs_auth == Some(true) {
|
||||
return Ok(Json(ApiResponse::ok(result)).into_response());
|
||||
}
|
||||
|
||||
let status = result
|
||||
.code
|
||||
.map(connection_test_failure_status)
|
||||
.unwrap_or(StatusCode::BAD_GATEWAY);
|
||||
let error = result
|
||||
.error
|
||||
.clone()
|
||||
.unwrap_or_else(|| "MCP connection test failed".to_string());
|
||||
let code = result
|
||||
.code
|
||||
.map(McpConnectionTestErrorCode::as_str)
|
||||
.unwrap_or("MCP_CONNECTION_FAILED");
|
||||
|
||||
Ok((
|
||||
status,
|
||||
Json(ErrorResponse::new_with_details(error, code, result.details.clone())),
|
||||
)
|
||||
.into_response())
|
||||
}
|
||||
|
||||
fn connection_test_failure_status(code: McpConnectionTestErrorCode) -> StatusCode {
|
||||
match code {
|
||||
McpConnectionTestErrorCode::CommandNotFound
|
||||
| McpConnectionTestErrorCode::CommandPermissionDenied
|
||||
| McpConnectionTestErrorCode::CommandStartFailed => StatusCode::UNPROCESSABLE_ENTITY,
|
||||
McpConnectionTestErrorCode::Timeout => StatusCode::GATEWAY_TIMEOUT,
|
||||
McpConnectionTestErrorCode::ConnectionFailed
|
||||
| McpConnectionTestErrorCode::HttpError
|
||||
| McpConnectionTestErrorCode::RpcError
|
||||
| McpConnectionTestErrorCode::ProtocolError => StatusCode::BAD_GATEWAY,
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Agent Sync Handlers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// `GET /api/mcp/agent-configs` — scan all installed Agent CLIs
|
||||
/// and return their current MCP server configurations.
|
||||
async fn get_agent_configs(
|
||||
State(state): State<McpRouterState>,
|
||||
) -> Result<Json<ApiResponse<Vec<DetectedMcpServerResponse>>>, AppError> {
|
||||
let configs = state.sync_service.get_agent_configs().await?;
|
||||
Ok(Json(ApiResponse::ok(configs)))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// OAuth Handlers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// `POST /api/mcp/oauth/check-status` — check OAuth authentication status.
|
||||
async fn oauth_check_status(
|
||||
State(state): State<McpRouterState>,
|
||||
body: Result<Json<OAuthCheckStatusRequest>, JsonRejection>,
|
||||
) -> Result<Json<ApiResponse<OAuthStatusResponse>>, AppError> {
|
||||
let Json(req) = body.map_err(|e| AppError::BadRequest(e.to_string()))?;
|
||||
let status = state.oauth_service.check_oauth_status(&req.server_url).await?;
|
||||
Ok(Json(ApiResponse::ok(status)))
|
||||
}
|
||||
|
||||
/// `POST /api/mcp/oauth/login` — start OAuth PKCE login flow.
|
||||
///
|
||||
/// Discovers endpoints, opens the browser for authorization, waits for
|
||||
/// the callback, and exchanges the code for tokens.
|
||||
async fn oauth_login(
|
||||
State(state): State<McpRouterState>,
|
||||
body: Result<Json<OAuthLoginRequest>, JsonRejection>,
|
||||
) -> Result<Json<ApiResponse<OAuthLoginResponse>>, AppError> {
|
||||
let Json(req) = body.map_err(|e| AppError::BadRequest(e.to_string()))?;
|
||||
let result = state.oauth_service.login(&req.server_url).await?;
|
||||
Ok(Json(ApiResponse::ok(result)))
|
||||
}
|
||||
|
||||
/// `POST /api/mcp/oauth/logout` — delete stored OAuth token.
|
||||
async fn oauth_logout(
|
||||
State(state): State<McpRouterState>,
|
||||
body: Result<Json<OAuthLogoutRequest>, JsonRejection>,
|
||||
) -> Result<Json<ApiResponse<()>>, AppError> {
|
||||
let Json(req) = body.map_err(|e| AppError::BadRequest(e.to_string()))?;
|
||||
state.oauth_service.logout(&req.server_url).await?;
|
||||
Ok(Json(ApiResponse::success()))
|
||||
}
|
||||
|
||||
/// `GET /api/mcp/oauth/authenticated` — list server URLs with stored tokens.
|
||||
async fn oauth_authenticated(State(state): State<McpRouterState>) -> Result<Json<ApiResponse<Vec<String>>>, AppError> {
|
||||
let urls = state.oauth_service.get_authenticated_servers().await?;
|
||||
Ok(Json(ApiResponse::ok(urls)))
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,713 @@
|
||||
use std::collections::HashMap;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::types::{McpServer, McpServerTransport};
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// NameValuePair
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// A name-value pair for environment variables and HTTP headers
|
||||
/// in the ACP session MCP server format.
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct NameValuePair {
|
||||
pub name: String,
|
||||
pub value: String,
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// AcpSessionMcpServer
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// ACP session MCP server configuration.
|
||||
///
|
||||
/// This is the wire format expected by the ACP backend when creating
|
||||
/// a new session with MCP servers injected. Two shapes:
|
||||
///
|
||||
/// - **Stdio**: command-based MCP servers (command + args + env)
|
||||
/// - **Http / Sse**: URL-based MCP servers (url + optional headers)
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(tag = "type", rename_all = "lowercase")]
|
||||
pub enum AcpSessionMcpServer {
|
||||
Stdio {
|
||||
name: String,
|
||||
command: String,
|
||||
#[serde(default)]
|
||||
args: Vec<String>,
|
||||
#[serde(default)]
|
||||
env: Vec<NameValuePair>,
|
||||
},
|
||||
Http {
|
||||
name: String,
|
||||
url: String,
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
headers: Vec<NameValuePair>,
|
||||
},
|
||||
Sse {
|
||||
name: String,
|
||||
url: String,
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
headers: Vec<NameValuePair>,
|
||||
},
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// AcpMcpCapabilities
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// ACP backend MCP capability declaration.
|
||||
///
|
||||
/// Describes which transport types the ACP backend supports for
|
||||
/// spawning MCP servers during a session.
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct AcpMcpCapabilities {
|
||||
pub stdio: bool,
|
||||
pub http: bool,
|
||||
pub sse: bool,
|
||||
}
|
||||
|
||||
impl AcpMcpCapabilities {
|
||||
/// Returns true if no transport type is supported.
|
||||
pub fn is_empty(&self) -> bool {
|
||||
!self.stdio && !self.http && !self.sse
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for AcpMcpCapabilities {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
stdio: true,
|
||||
http: false,
|
||||
sse: false,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// ImageGenConfig
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Configuration for the builtin image generation MCP server.
|
||||
///
|
||||
/// Values are injected as environment variables when building
|
||||
/// the builtin MCP server config for ACP sessions.
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct ImageGenConfig {
|
||||
pub model: Option<String>,
|
||||
pub api_url: Option<String>,
|
||||
pub api_key: Option<String>,
|
||||
pub size: Option<String>,
|
||||
pub quality: Option<String>,
|
||||
pub style: Option<String>,
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Public API
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Parse ACP MCP capabilities from an ACP backend response.
|
||||
///
|
||||
/// Looks for capabilities under `mcp_capabilities`, `mcpCapabilities`,
|
||||
/// or `mcp` keys. Returns default capabilities (stdio only) when the
|
||||
/// field is missing or not an object.
|
||||
pub fn parse_acp_mcp_capabilities(response: &serde_json::Value) -> AcpMcpCapabilities {
|
||||
let caps = response
|
||||
.get("mcp_capabilities")
|
||||
.or_else(|| response.get("mcpCapabilities"))
|
||||
.or_else(|| response.get("mcp"));
|
||||
|
||||
let Some(caps) = caps else {
|
||||
return AcpMcpCapabilities::default();
|
||||
};
|
||||
|
||||
let http = bool_field(caps, "http");
|
||||
let sse = bool_field(caps, "sse");
|
||||
let stdio = bool_field(caps, "stdio") || http || sse;
|
||||
|
||||
AcpMcpCapabilities { stdio, http, sse }
|
||||
}
|
||||
|
||||
/// Build ACP session MCP server configs from domain servers.
|
||||
///
|
||||
/// Filters to only enabled servers whose transport type is supported
|
||||
/// by the ACP backend, then converts to the ACP wire format.
|
||||
pub fn build_session_mcp_servers(servers: &[McpServer], capabilities: &AcpMcpCapabilities) -> Vec<AcpSessionMcpServer> {
|
||||
servers
|
||||
.iter()
|
||||
.filter(|s| s.enabled)
|
||||
.filter_map(|s| convert_server(s, capabilities))
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Build the builtin image generation MCP server config.
|
||||
///
|
||||
/// Returns `None` if the ACP backend doesn't support stdio transport
|
||||
/// or if `command` is empty.
|
||||
pub fn build_builtin_image_gen_server(
|
||||
capabilities: &AcpMcpCapabilities,
|
||||
command: &str,
|
||||
config: &ImageGenConfig,
|
||||
) -> Option<AcpSessionMcpServer> {
|
||||
if !capabilities.stdio || command.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let env = build_image_gen_env(config);
|
||||
|
||||
Some(AcpSessionMcpServer::Stdio {
|
||||
name: "nomifun-image-generation".into(),
|
||||
command: command.to_owned(),
|
||||
args: Vec::new(),
|
||||
env,
|
||||
})
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Internal helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Extract a boolean field from a JSON value, defaulting to false.
|
||||
fn bool_field(value: &serde_json::Value, key: &str) -> bool {
|
||||
value.get(key).and_then(|v| v.as_bool()).unwrap_or(false)
|
||||
}
|
||||
|
||||
/// Convert a domain `McpServer` to `AcpSessionMcpServer`.
|
||||
///
|
||||
/// Returns `None` if the server's transport type is not supported
|
||||
/// by the given capabilities.
|
||||
fn convert_server(server: &McpServer, capabilities: &AcpMcpCapabilities) -> Option<AcpSessionMcpServer> {
|
||||
match &server.transport {
|
||||
McpServerTransport::Stdio { command, args, env } if capabilities.stdio => Some(AcpSessionMcpServer::Stdio {
|
||||
name: server.name.clone(),
|
||||
command: command.clone(),
|
||||
args: args.clone(),
|
||||
env: hashmap_to_pairs(env),
|
||||
}),
|
||||
McpServerTransport::Http { url, headers } if capabilities.http => Some(AcpSessionMcpServer::Http {
|
||||
name: server.name.clone(),
|
||||
url: url.clone(),
|
||||
headers: hashmap_to_pairs(headers),
|
||||
}),
|
||||
McpServerTransport::Sse { url, headers } if capabilities.sse => Some(AcpSessionMcpServer::Sse {
|
||||
name: server.name.clone(),
|
||||
url: url.clone(),
|
||||
headers: hashmap_to_pairs(headers),
|
||||
}),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Convert a `HashMap<String, String>` to a sorted `Vec<NameValuePair>`.
|
||||
///
|
||||
/// Sorted by key for deterministic serialization.
|
||||
fn hashmap_to_pairs(map: &HashMap<String, String>) -> Vec<NameValuePair> {
|
||||
let mut pairs: Vec<NameValuePair> = map
|
||||
.iter()
|
||||
.map(|(k, v)| NameValuePair {
|
||||
name: k.clone(),
|
||||
value: v.clone(),
|
||||
})
|
||||
.collect();
|
||||
pairs.sort_by(|a, b| a.name.cmp(&b.name));
|
||||
pairs
|
||||
}
|
||||
|
||||
/// Build environment variable pairs for the image generation server.
|
||||
///
|
||||
/// Sorted by name for deterministic output, consistent with `hashmap_to_pairs`.
|
||||
fn build_image_gen_env(config: &ImageGenConfig) -> Vec<NameValuePair> {
|
||||
let entries: [(&str, &Option<String>); 6] = [
|
||||
("NOMIFUN_IMG_API_KEY", &config.api_key),
|
||||
("NOMIFUN_IMG_API_URL", &config.api_url),
|
||||
("NOMIFUN_IMG_MODEL", &config.model),
|
||||
("NOMIFUN_IMG_QUALITY", &config.quality),
|
||||
("NOMIFUN_IMG_SIZE", &config.size),
|
||||
("NOMIFUN_IMG_STYLE", &config.style),
|
||||
];
|
||||
|
||||
entries
|
||||
.into_iter()
|
||||
.filter_map(|(name, value)| {
|
||||
value.as_ref().map(|v| NameValuePair {
|
||||
name: name.into(),
|
||||
value: v.clone(),
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::types::McpServerTransport;
|
||||
use nomifun_common::McpServerStatus;
|
||||
|
||||
// -- helpers --
|
||||
|
||||
fn make_server(name: &str, enabled: bool, transport: McpServerTransport) -> McpServer {
|
||||
McpServer {
|
||||
// Injection keys on `name`, never `id`; any stable value works here.
|
||||
id: name.bytes().map(i64::from).sum::<i64>().max(1),
|
||||
name: name.into(),
|
||||
description: None,
|
||||
enabled,
|
||||
transport,
|
||||
tools: vec![],
|
||||
last_test_status: McpServerStatus::Disconnected,
|
||||
last_connected: None,
|
||||
original_json: None,
|
||||
builtin: false,
|
||||
created_at: 0,
|
||||
updated_at: 0,
|
||||
}
|
||||
}
|
||||
|
||||
fn stdio_transport(cmd: &str) -> McpServerTransport {
|
||||
McpServerTransport::Stdio {
|
||||
command: cmd.into(),
|
||||
args: vec!["-y".into(), "@test/server".into()],
|
||||
env: HashMap::from([("NODE_ENV".into(), "production".into())]),
|
||||
}
|
||||
}
|
||||
|
||||
fn http_transport(url: &str) -> McpServerTransport {
|
||||
McpServerTransport::Http {
|
||||
url: url.into(),
|
||||
headers: HashMap::from([("Authorization".into(), "Bearer tok".into())]),
|
||||
}
|
||||
}
|
||||
|
||||
fn sse_transport(url: &str) -> McpServerTransport {
|
||||
McpServerTransport::Sse {
|
||||
url: url.into(),
|
||||
headers: HashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn all_caps() -> AcpMcpCapabilities {
|
||||
AcpMcpCapabilities {
|
||||
stdio: true,
|
||||
http: true,
|
||||
sse: true,
|
||||
}
|
||||
}
|
||||
|
||||
// -- AcpMcpCapabilities ---------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn capabilities_default_is_stdio_only() {
|
||||
let caps = AcpMcpCapabilities::default();
|
||||
assert!(caps.stdio);
|
||||
assert!(!caps.http);
|
||||
assert!(!caps.sse);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn capabilities_is_empty() {
|
||||
let empty = AcpMcpCapabilities {
|
||||
stdio: false,
|
||||
http: false,
|
||||
sse: false,
|
||||
};
|
||||
assert!(empty.is_empty());
|
||||
assert!(!AcpMcpCapabilities::default().is_empty());
|
||||
}
|
||||
|
||||
// -- parse_acp_mcp_capabilities -------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn parse_full_capabilities() {
|
||||
let resp = serde_json::json!({
|
||||
"mcp_capabilities": { "stdio": true, "http": true, "sse": true }
|
||||
});
|
||||
let caps = parse_acp_mcp_capabilities(&resp);
|
||||
assert_eq!(caps, all_caps());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_camel_case_key() {
|
||||
let resp = serde_json::json!({
|
||||
"mcpCapabilities": { "stdio": true, "http": false, "sse": true }
|
||||
});
|
||||
let caps = parse_acp_mcp_capabilities(&resp);
|
||||
assert!(caps.stdio);
|
||||
assert!(!caps.http);
|
||||
assert!(caps.sse);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_mcp_shorthand_key() {
|
||||
let resp = serde_json::json!({
|
||||
"mcp": { "stdio": false, "http": true, "sse": false }
|
||||
});
|
||||
let caps = parse_acp_mcp_capabilities(&resp);
|
||||
assert!(caps.stdio);
|
||||
assert!(caps.http);
|
||||
assert!(!caps.sse);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_missing_capabilities_returns_default() {
|
||||
let resp = serde_json::json!({ "other": "data" });
|
||||
let caps = parse_acp_mcp_capabilities(&resp);
|
||||
assert_eq!(caps, AcpMcpCapabilities::default());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_partial_capabilities_defaults_missing_to_false() {
|
||||
let resp = serde_json::json!({
|
||||
"mcp_capabilities": { "stdio": true }
|
||||
});
|
||||
let caps = parse_acp_mcp_capabilities(&resp);
|
||||
assert!(caps.stdio);
|
||||
assert!(!caps.http);
|
||||
assert!(!caps.sse);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_http_support_implies_stdio() {
|
||||
let resp = serde_json::json!({
|
||||
"mcp_capabilities": { "http": true, "sse": false }
|
||||
});
|
||||
let caps = parse_acp_mcp_capabilities(&resp);
|
||||
assert!(caps.stdio);
|
||||
assert!(caps.http);
|
||||
assert!(!caps.sse);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_priority_mcp_capabilities_over_mcp() {
|
||||
let resp = serde_json::json!({
|
||||
"mcp_capabilities": { "stdio": true, "http": true, "sse": true },
|
||||
"mcp": { "stdio": false, "http": false, "sse": false }
|
||||
});
|
||||
let caps = parse_acp_mcp_capabilities(&resp);
|
||||
assert_eq!(caps, all_caps());
|
||||
}
|
||||
|
||||
// -- convert_server -------------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn convert_stdio_server() {
|
||||
let server = make_server("test", true, stdio_transport("npx"));
|
||||
let result = convert_server(&server, &all_caps());
|
||||
assert!(result.is_some());
|
||||
let acp = result.unwrap();
|
||||
match acp {
|
||||
AcpSessionMcpServer::Stdio {
|
||||
name,
|
||||
command,
|
||||
args,
|
||||
env,
|
||||
} => {
|
||||
assert_eq!(name, "test");
|
||||
assert_eq!(command, "npx");
|
||||
assert_eq!(args, vec!["-y", "@test/server"]);
|
||||
assert_eq!(env.len(), 1);
|
||||
assert_eq!(env[0].name, "NODE_ENV");
|
||||
assert_eq!(env[0].value, "production");
|
||||
}
|
||||
_ => panic!("expected Stdio"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn convert_http_server() {
|
||||
let server = make_server("http-test", true, http_transport("https://example.com/mcp"));
|
||||
let result = convert_server(&server, &all_caps());
|
||||
assert!(result.is_some());
|
||||
match result.unwrap() {
|
||||
AcpSessionMcpServer::Http { name, url, headers } => {
|
||||
assert_eq!(name, "http-test");
|
||||
assert_eq!(url, "https://example.com/mcp");
|
||||
assert_eq!(headers.len(), 1);
|
||||
assert_eq!(headers[0].name, "Authorization");
|
||||
assert_eq!(headers[0].value, "Bearer tok");
|
||||
}
|
||||
_ => panic!("expected Http"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn convert_sse_server() {
|
||||
let server = make_server("sse-test", true, sse_transport("https://example.com/sse"));
|
||||
let result = convert_server(&server, &all_caps());
|
||||
assert!(result.is_some());
|
||||
match result.unwrap() {
|
||||
AcpSessionMcpServer::Sse { name, url, headers, .. } => {
|
||||
assert_eq!(name, "sse-test");
|
||||
assert_eq!(url, "https://example.com/sse");
|
||||
assert!(headers.is_empty());
|
||||
}
|
||||
_ => panic!("expected Sse"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn convert_skips_unsupported_transport() {
|
||||
let stdio_only = AcpMcpCapabilities {
|
||||
stdio: true,
|
||||
http: false,
|
||||
sse: false,
|
||||
};
|
||||
let http_server = make_server("http-test", true, http_transport("https://example.com/mcp"));
|
||||
assert!(convert_server(&http_server, &stdio_only).is_none());
|
||||
|
||||
let sse_server = make_server("sse-test", true, sse_transport("https://example.com/sse"));
|
||||
assert!(convert_server(&sse_server, &stdio_only).is_none());
|
||||
}
|
||||
|
||||
// -- build_session_mcp_servers --------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn build_filters_disabled_servers() {
|
||||
let servers = vec![
|
||||
make_server("enabled", true, stdio_transport("npx")),
|
||||
make_server("disabled", false, stdio_transport("node")),
|
||||
];
|
||||
let result = build_session_mcp_servers(&servers, &all_caps());
|
||||
assert_eq!(result.len(), 1);
|
||||
match &result[0] {
|
||||
AcpSessionMcpServer::Stdio { name, .. } => assert_eq!(name, "enabled"),
|
||||
_ => panic!("expected Stdio"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_filters_by_capabilities() {
|
||||
let caps = AcpMcpCapabilities {
|
||||
stdio: true,
|
||||
http: false,
|
||||
sse: true,
|
||||
};
|
||||
let servers = vec![
|
||||
make_server("s1", true, stdio_transport("npx")),
|
||||
make_server("s2", true, http_transport("https://example.com/mcp")),
|
||||
make_server("s3", true, sse_transport("https://example.com/sse")),
|
||||
];
|
||||
let result = build_session_mcp_servers(&servers, &caps);
|
||||
assert_eq!(result.len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_empty_servers_returns_empty() {
|
||||
let result = build_session_mcp_servers(&[], &all_caps());
|
||||
assert!(result.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_no_capabilities_returns_empty() {
|
||||
let no_caps = AcpMcpCapabilities {
|
||||
stdio: false,
|
||||
http: false,
|
||||
sse: false,
|
||||
};
|
||||
let servers = vec![
|
||||
make_server("s1", true, stdio_transport("npx")),
|
||||
make_server("s2", true, http_transport("https://example.com")),
|
||||
];
|
||||
let result = build_session_mcp_servers(&servers, &no_caps);
|
||||
assert!(result.is_empty());
|
||||
}
|
||||
|
||||
// -- build_builtin_image_gen_server ---------------------------------------
|
||||
|
||||
#[test]
|
||||
fn builtin_image_gen_with_full_config() {
|
||||
let caps = all_caps();
|
||||
let config = ImageGenConfig {
|
||||
model: Some("dall-e-3".into()),
|
||||
api_url: Some("https://api.openai.com".into()),
|
||||
api_key: Some("sk-test".into()),
|
||||
size: Some("1024x1024".into()),
|
||||
quality: Some("hd".into()),
|
||||
style: Some("natural".into()),
|
||||
};
|
||||
let result = build_builtin_image_gen_server(&caps, "/usr/bin/img-gen", &config);
|
||||
assert!(result.is_some());
|
||||
match result.unwrap() {
|
||||
AcpSessionMcpServer::Stdio {
|
||||
name,
|
||||
command,
|
||||
args,
|
||||
env,
|
||||
} => {
|
||||
assert_eq!(name, "nomifun-image-generation");
|
||||
assert_eq!(command, "/usr/bin/img-gen");
|
||||
assert!(args.is_empty());
|
||||
assert_eq!(env.len(), 6);
|
||||
assert_eq!(env[0].name, "NOMIFUN_IMG_API_KEY");
|
||||
assert_eq!(env[0].value, "sk-test");
|
||||
}
|
||||
_ => panic!("expected Stdio"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builtin_image_gen_with_partial_config() {
|
||||
let caps = all_caps();
|
||||
let config = ImageGenConfig {
|
||||
model: Some("dall-e-3".into()),
|
||||
api_url: Some("https://api.openai.com".into()),
|
||||
..Default::default()
|
||||
};
|
||||
let result = build_builtin_image_gen_server(&caps, "img-gen", &config);
|
||||
assert!(result.is_some());
|
||||
match result.unwrap() {
|
||||
AcpSessionMcpServer::Stdio { env, .. } => {
|
||||
assert_eq!(env.len(), 2);
|
||||
// Sorted alphabetically: API_URL before MODEL
|
||||
assert_eq!(env[0].name, "NOMIFUN_IMG_API_URL");
|
||||
assert_eq!(env[0].value, "https://api.openai.com");
|
||||
assert_eq!(env[1].name, "NOMIFUN_IMG_MODEL");
|
||||
assert_eq!(env[1].value, "dall-e-3");
|
||||
}
|
||||
_ => panic!("expected Stdio"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builtin_image_gen_no_stdio_returns_none() {
|
||||
let caps = AcpMcpCapabilities {
|
||||
stdio: false,
|
||||
http: true,
|
||||
sse: true,
|
||||
};
|
||||
let config = ImageGenConfig {
|
||||
model: Some("dall-e-3".into()),
|
||||
..Default::default()
|
||||
};
|
||||
assert!(build_builtin_image_gen_server(&caps, "img-gen", &config).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builtin_image_gen_empty_command_returns_none() {
|
||||
let caps = all_caps();
|
||||
let config = ImageGenConfig::default();
|
||||
assert!(build_builtin_image_gen_server(&caps, "", &config).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builtin_image_gen_empty_config() {
|
||||
let caps = all_caps();
|
||||
let config = ImageGenConfig::default();
|
||||
let result = build_builtin_image_gen_server(&caps, "img-gen", &config);
|
||||
assert!(result.is_some());
|
||||
match result.unwrap() {
|
||||
AcpSessionMcpServer::Stdio { env, .. } => {
|
||||
assert!(env.is_empty());
|
||||
}
|
||||
_ => panic!("expected Stdio"),
|
||||
}
|
||||
}
|
||||
|
||||
// -- hashmap_to_pairs -----------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn hashmap_to_pairs_sorted() {
|
||||
let map = HashMap::from([
|
||||
("Z_KEY".into(), "z_val".into()),
|
||||
("A_KEY".into(), "a_val".into()),
|
||||
("M_KEY".into(), "m_val".into()),
|
||||
]);
|
||||
let pairs = hashmap_to_pairs(&map);
|
||||
assert_eq!(pairs.len(), 3);
|
||||
assert_eq!(pairs[0].name, "A_KEY");
|
||||
assert_eq!(pairs[1].name, "M_KEY");
|
||||
assert_eq!(pairs[2].name, "Z_KEY");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn hashmap_to_pairs_empty() {
|
||||
let map = HashMap::new();
|
||||
let pairs = hashmap_to_pairs(&map);
|
||||
assert!(pairs.is_empty());
|
||||
}
|
||||
|
||||
// -- Serialization roundtrip ----------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn stdio_serialization_roundtrip() {
|
||||
let server = AcpSessionMcpServer::Stdio {
|
||||
name: "test".into(),
|
||||
command: "npx".into(),
|
||||
args: vec!["-y".into()],
|
||||
env: vec![NameValuePair {
|
||||
name: "K".into(),
|
||||
value: "V".into(),
|
||||
}],
|
||||
};
|
||||
let json = serde_json::to_string(&server).unwrap();
|
||||
let parsed: AcpSessionMcpServer = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(server, parsed);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn http_serialization_roundtrip() {
|
||||
let server = AcpSessionMcpServer::Http {
|
||||
name: "http-test".into(),
|
||||
url: "https://example.com/mcp".into(),
|
||||
headers: vec![NameValuePair {
|
||||
name: "Auth".into(),
|
||||
value: "Bearer tok".into(),
|
||||
}],
|
||||
};
|
||||
let json = serde_json::to_string(&server).unwrap();
|
||||
let parsed: AcpSessionMcpServer = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(server, parsed);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sse_serialization_roundtrip() {
|
||||
let server = AcpSessionMcpServer::Sse {
|
||||
name: "sse-test".into(),
|
||||
url: "https://example.com/sse".into(),
|
||||
headers: vec![],
|
||||
};
|
||||
let json = serde_json::to_string(&server).unwrap();
|
||||
assert!(!json.contains("headers")); // skip_serializing_if
|
||||
let parsed: AcpSessionMcpServer = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(server, parsed);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stdio_json_has_type_field() {
|
||||
let server = AcpSessionMcpServer::Stdio {
|
||||
name: "test".into(),
|
||||
command: "npx".into(),
|
||||
args: vec![],
|
||||
env: vec![],
|
||||
};
|
||||
let value: serde_json::Value = serde_json::to_value(&server).unwrap();
|
||||
assert_eq!(value["type"], "stdio");
|
||||
assert_eq!(value["name"], "test");
|
||||
assert_eq!(value["command"], "npx");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn http_json_has_type_field() {
|
||||
let server = AcpSessionMcpServer::Http {
|
||||
name: "h".into(),
|
||||
url: "https://example.com".into(),
|
||||
headers: vec![],
|
||||
};
|
||||
let value: serde_json::Value = serde_json::to_value(&server).unwrap();
|
||||
assert_eq!(value["type"], "http");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sse_json_has_type_field() {
|
||||
let server = AcpSessionMcpServer::Sse {
|
||||
name: "s".into(),
|
||||
url: "https://example.com".into(),
|
||||
headers: vec![],
|
||||
};
|
||||
let value: serde_json::Value = serde_json::to_value(&server).unwrap();
|
||||
assert_eq!(value["type"], "sse");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,293 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use dashmap::DashMap;
|
||||
use nomifun_api_types::{DetectedMcpServerEntry, DetectedMcpServerResponse};
|
||||
use nomifun_common::McpSource;
|
||||
use nomifun_db::IMcpServerRepository;
|
||||
use tokio::sync::Mutex;
|
||||
use tracing::warn;
|
||||
|
||||
use crate::adapter::{DetectedServer, McpAgentAdapter};
|
||||
use crate::error::McpError;
|
||||
|
||||
/// Discovers MCP configuration currently installed in external Agent CLIs.
|
||||
///
|
||||
/// This service is intentionally read-only. It serializes detection work to
|
||||
/// avoid concurrent CLI scans from spawning overlapping child processes.
|
||||
#[derive(Clone)]
|
||||
pub struct McpSyncService {
|
||||
adapters: Arc<Vec<Arc<dyn McpAgentAdapter>>>,
|
||||
service_lock: Arc<Mutex<()>>,
|
||||
agent_locks: Arc<DashMap<McpSource, Arc<Mutex<()>>>>,
|
||||
}
|
||||
|
||||
impl McpSyncService {
|
||||
pub fn new(_repo: Arc<dyn IMcpServerRepository>, adapters: Vec<Arc<dyn McpAgentAdapter>>) -> Self {
|
||||
Self {
|
||||
adapters: Arc::new(adapters),
|
||||
service_lock: Arc::new(Mutex::new(())),
|
||||
agent_locks: Arc::new(DashMap::new()),
|
||||
}
|
||||
}
|
||||
|
||||
/// Scan all installed Agent CLIs and return each one's current MCP
|
||||
/// server configurations.
|
||||
///
|
||||
/// Agents that are not installed are silently skipped.
|
||||
pub async fn get_agent_configs(&self) -> Result<Vec<DetectedMcpServerResponse>, McpError> {
|
||||
let _guard = self.service_lock.lock().await;
|
||||
|
||||
let mut results = Vec::new();
|
||||
for adapter in self.adapters.iter() {
|
||||
let _agent_guard = self.agent_lock(adapter.source()).await;
|
||||
|
||||
let installed = adapter.is_installed().await.unwrap_or(false);
|
||||
if !installed {
|
||||
continue;
|
||||
}
|
||||
|
||||
match adapter.detect_existing().await {
|
||||
Ok(detected) => {
|
||||
let servers = detected.into_iter().map(detected_to_response).collect();
|
||||
results.push(DetectedMcpServerResponse {
|
||||
source: adapter.source(),
|
||||
servers,
|
||||
});
|
||||
}
|
||||
Err(e) => {
|
||||
warn!(
|
||||
agent = ?adapter.source(),
|
||||
error = %e,
|
||||
"failed to detect existing MCP servers"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(results)
|
||||
}
|
||||
|
||||
async fn agent_lock(&self, source: McpSource) -> tokio::sync::OwnedMutexGuard<()> {
|
||||
let lock = self
|
||||
.agent_locks
|
||||
.entry(source)
|
||||
.or_insert_with(|| Arc::new(Mutex::new(())))
|
||||
.clone();
|
||||
lock.lock_owned().await
|
||||
}
|
||||
}
|
||||
|
||||
fn detected_to_response(detected: DetectedServer) -> DetectedMcpServerEntry {
|
||||
let normalized_skip_reason = detected.import_skip_reason.as_deref().map(normalize_import_skip_reason);
|
||||
let importable = detected.importable || normalized_skip_reason.as_deref() == Some("Connected");
|
||||
|
||||
DetectedMcpServerEntry {
|
||||
server: nomifun_api_types::McpServerResponse {
|
||||
// Detected servers are not DB entities and have no host-local
|
||||
// primary key. Clients match/import them by `name`/`transport`, so
|
||||
// a `0` sentinel id is sufficient and never read for identity.
|
||||
id: 0,
|
||||
name: detected.name,
|
||||
description: None,
|
||||
enabled: false,
|
||||
transport: detected.transport.into(),
|
||||
tools: None,
|
||||
last_test_status: nomifun_common::McpServerStatus::Disconnected,
|
||||
last_connected: None,
|
||||
original_json: None,
|
||||
builtin: false,
|
||||
created_at: 0,
|
||||
updated_at: 0,
|
||||
},
|
||||
importable,
|
||||
import_skip_reason: if importable { None } else { normalized_skip_reason },
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_import_skip_reason(reason: &str) -> String {
|
||||
reason
|
||||
.trim()
|
||||
.trim_start_matches(|c: char| {
|
||||
matches!(c, '✓' | '✗' | '!' | '•' | '-' | '*' | '✔' | '✘' | ':' | '[' | ']') || c.is_whitespace()
|
||||
})
|
||||
.trim()
|
||||
.to_owned()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::types::McpServerTransport;
|
||||
use nomifun_common::{McpServerStatus, TimestampMs};
|
||||
use nomifun_db::models::McpServerRow;
|
||||
use nomifun_db::{CreateMcpServerParams, DbError, UpdateMcpServerParams};
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Mutex as StdMutex;
|
||||
|
||||
struct MockAdapter {
|
||||
source: McpSource,
|
||||
installed: bool,
|
||||
servers: Arc<StdMutex<Vec<DetectedServer>>>,
|
||||
}
|
||||
|
||||
impl MockAdapter {
|
||||
fn new(source: McpSource, installed: bool) -> Self {
|
||||
Self {
|
||||
source,
|
||||
installed,
|
||||
servers: Arc::new(StdMutex::new(Vec::new())),
|
||||
}
|
||||
}
|
||||
|
||||
fn with_existing(mut self, servers: Vec<DetectedServer>) -> Self {
|
||||
self.servers = Arc::new(StdMutex::new(servers));
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl McpAgentAdapter for MockAdapter {
|
||||
fn source(&self) -> McpSource {
|
||||
self.source
|
||||
}
|
||||
|
||||
async fn is_installed(&self) -> Result<bool, McpError> {
|
||||
Ok(self.installed)
|
||||
}
|
||||
|
||||
async fn detect_existing(&self) -> Result<Vec<DetectedServer>, McpError> {
|
||||
if !self.installed {
|
||||
return Err(McpError::AgentNotInstalled(format!("{:?}", self.source)));
|
||||
}
|
||||
Ok(self.servers.lock().unwrap().clone())
|
||||
}
|
||||
|
||||
async fn install_server(&self, _name: &str, _transport: &McpServerTransport) -> Result<(), McpError> {
|
||||
unreachable!("write-to-CLI is no longer supported")
|
||||
}
|
||||
|
||||
async fn remove_server(&self, _name: &str) -> Result<(), McpError> {
|
||||
unreachable!("write-to-CLI is no longer supported")
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct MockRepo;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl IMcpServerRepository for MockRepo {
|
||||
async fn list(&self) -> Result<Vec<McpServerRow>, DbError> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
|
||||
async fn find_by_id(&self, _id: i64) -> Result<Option<McpServerRow>, DbError> {
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn find_by_name(&self, _name: &str) -> Result<Option<McpServerRow>, DbError> {
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn create(&self, _params: CreateMcpServerParams<'_>) -> Result<McpServerRow, DbError> {
|
||||
unimplemented!("not needed for detection tests")
|
||||
}
|
||||
|
||||
async fn update(&self, _id: i64, _params: UpdateMcpServerParams<'_>) -> Result<McpServerRow, DbError> {
|
||||
unimplemented!("not needed for detection tests")
|
||||
}
|
||||
|
||||
async fn delete(&self, _id: i64) -> Result<(), DbError> {
|
||||
unimplemented!("not needed for detection tests")
|
||||
}
|
||||
|
||||
async fn batch_upsert(&self, _params_list: &[CreateMcpServerParams<'_>]) -> Result<Vec<McpServerRow>, DbError> {
|
||||
unimplemented!("not needed for detection tests")
|
||||
}
|
||||
|
||||
async fn update_status(
|
||||
&self,
|
||||
_id: i64,
|
||||
_status: &str,
|
||||
_last_connected: Option<TimestampMs>,
|
||||
) -> Result<(), DbError> {
|
||||
unimplemented!("not needed for detection tests")
|
||||
}
|
||||
|
||||
async fn update_tools(&self, _id: i64, _tools: Option<&str>) -> Result<(), DbError> {
|
||||
unimplemented!("not needed for detection tests")
|
||||
}
|
||||
}
|
||||
|
||||
fn stdio_transport() -> McpServerTransport {
|
||||
McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec!["-y".into(), "@test/server".into()],
|
||||
env: HashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn make_service(adapters: Vec<Arc<dyn McpAgentAdapter>>) -> McpSyncService {
|
||||
McpSyncService::new(Arc::new(MockRepo), adapters)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn get_agent_configs_returns_installed_only() {
|
||||
let adapter_a = Arc::new(
|
||||
MockAdapter::new(McpSource::Claude, true).with_existing(vec![DetectedServer {
|
||||
name: "srv1".into(),
|
||||
transport: stdio_transport(),
|
||||
importable: true,
|
||||
import_skip_reason: None,
|
||||
}]),
|
||||
);
|
||||
let adapter_b = Arc::new(MockAdapter::new(McpSource::Gemini, false));
|
||||
let adapter_c = Arc::new(MockAdapter::new(McpSource::Qwen, true).with_existing(vec![]));
|
||||
|
||||
let svc = make_service(vec![adapter_a, adapter_b, adapter_c]);
|
||||
let configs = svc.get_agent_configs().await.unwrap();
|
||||
|
||||
assert_eq!(configs.len(), 2);
|
||||
assert_eq!(configs[0].source, McpSource::Claude);
|
||||
assert_eq!(configs[0].servers.len(), 1);
|
||||
assert_eq!(configs[0].servers[0].server.name, "srv1");
|
||||
assert_eq!(configs[1].source, McpSource::Qwen);
|
||||
assert!(configs[1].servers.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn detected_to_response_normalizes_connected_skip_reason() {
|
||||
let resp = detected_to_response(DetectedServer {
|
||||
name: "sentry".into(),
|
||||
transport: stdio_transport(),
|
||||
importable: false,
|
||||
import_skip_reason: Some("✓ Connected".into()),
|
||||
});
|
||||
|
||||
assert!(resp.importable);
|
||||
assert_eq!(resp.import_skip_reason, None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn get_agent_configs_no_adapters() {
|
||||
let svc = make_service(vec![]);
|
||||
let configs = svc.get_agent_configs().await.unwrap();
|
||||
assert!(configs.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn detected_to_response_fields() {
|
||||
let detected = DetectedServer {
|
||||
name: "my-srv".into(),
|
||||
transport: stdio_transport(),
|
||||
importable: false,
|
||||
import_skip_reason: Some("Needs authentication".into()),
|
||||
};
|
||||
let resp = detected_to_response(detected);
|
||||
assert_eq!(resp.server.name, "my-srv");
|
||||
assert_eq!(resp.server.id, 0);
|
||||
assert!(!resp.server.enabled);
|
||||
assert_eq!(resp.server.last_test_status, McpServerStatus::Disconnected);
|
||||
assert!(!resp.importable);
|
||||
assert_eq!(resp.import_skip_reason.as_deref(), Some("Needs authentication"));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,606 @@
|
||||
use std::collections::HashMap;
|
||||
|
||||
use nomifun_api_types::{McpServerResponse, McpToolResponse, McpTransport};
|
||||
use nomifun_common::{McpServerStatus, TimestampMs};
|
||||
use nomifun_db::models::McpServerRow;
|
||||
|
||||
use crate::error::McpError;
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// McpServerTransport — domain transport enum
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Domain-layer MCP server transport configuration.
|
||||
///
|
||||
/// Mirrors `McpTransport` from `nomifun-api-types` but lives in the business
|
||||
/// layer. Conversions are provided in both directions.
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub enum McpServerTransport {
|
||||
Stdio {
|
||||
command: String,
|
||||
args: Vec<String>,
|
||||
env: HashMap<String, String>,
|
||||
},
|
||||
Sse {
|
||||
url: String,
|
||||
headers: HashMap<String, String>,
|
||||
},
|
||||
Http {
|
||||
url: String,
|
||||
headers: HashMap<String, String>,
|
||||
},
|
||||
}
|
||||
|
||||
impl McpServerTransport {
|
||||
/// Returns the transport type as a string for DB storage.
|
||||
pub fn transport_type(&self) -> &'static str {
|
||||
match self {
|
||||
Self::Stdio { .. } => "stdio",
|
||||
Self::Sse { .. } => "sse",
|
||||
Self::Http { .. } => "http",
|
||||
}
|
||||
}
|
||||
|
||||
/// Serializes the transport config to a JSON string for DB storage.
|
||||
///
|
||||
/// Only serializes the variant-specific fields (command/args/env or
|
||||
/// url/headers), not the type discriminant.
|
||||
pub fn to_config_json(&self) -> Result<String, McpError> {
|
||||
let value = match self {
|
||||
Self::Stdio { command, args, env } => {
|
||||
serde_json::json!({ "command": command, "args": args, "env": env })
|
||||
}
|
||||
Self::Sse { url, headers } => {
|
||||
serde_json::json!({ "url": url, "headers": headers })
|
||||
}
|
||||
Self::Http { url, headers } => {
|
||||
serde_json::json!({ "url": url, "headers": headers })
|
||||
}
|
||||
};
|
||||
serde_json::to_string(&value).map_err(McpError::from)
|
||||
}
|
||||
|
||||
/// Parses transport from DB fields (type string + config JSON).
|
||||
pub fn from_db(transport_type: &str, config_json: &str) -> Result<Self, McpError> {
|
||||
let value: serde_json::Value = serde_json::from_str(config_json).map_err(McpError::from)?;
|
||||
|
||||
match transport_type {
|
||||
"stdio" => {
|
||||
let command = value["command"]
|
||||
.as_str()
|
||||
.ok_or_else(|| McpError::InvalidTransport("stdio: missing command".into()))?
|
||||
.to_owned();
|
||||
let args = value["args"]
|
||||
.as_array()
|
||||
.map(|arr| arr.iter().filter_map(|v| v.as_str().map(String::from)).collect())
|
||||
.unwrap_or_default();
|
||||
let env = value["env"]
|
||||
.as_object()
|
||||
.map(|obj| {
|
||||
obj.iter()
|
||||
.filter_map(|(k, v)| v.as_str().map(|s| (k.clone(), s.to_owned())))
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
Ok(Self::Stdio { command, args, env })
|
||||
}
|
||||
"sse" => {
|
||||
let url = value["url"]
|
||||
.as_str()
|
||||
.ok_or_else(|| McpError::InvalidTransport("sse: missing url".into()))?
|
||||
.to_owned();
|
||||
let headers = parse_headers_object(&value["headers"]);
|
||||
Ok(Self::Sse { url, headers })
|
||||
}
|
||||
"http" => {
|
||||
let url = value["url"]
|
||||
.as_str()
|
||||
.ok_or_else(|| McpError::InvalidTransport("http: missing url".into()))?
|
||||
.to_owned();
|
||||
let headers = parse_headers_object(&value["headers"]);
|
||||
Ok(Self::Http { url, headers })
|
||||
}
|
||||
other => Err(McpError::InvalidTransport(format!("unknown transport type: {other}"))),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Helper: extract `HashMap<String, String>` from a JSON object value.
|
||||
fn parse_headers_object(value: &serde_json::Value) -> HashMap<String, String> {
|
||||
value
|
||||
.as_object()
|
||||
.map(|obj| {
|
||||
obj.iter()
|
||||
.filter_map(|(k, v)| v.as_str().map(|s| (k.clone(), s.to_owned())))
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
// -- Conversions between domain and API transport --
|
||||
|
||||
impl From<McpTransport> for McpServerTransport {
|
||||
fn from(t: McpTransport) -> Self {
|
||||
match t {
|
||||
McpTransport::Stdio { command, args, env } => Self::Stdio { command, args, env },
|
||||
McpTransport::Sse { url, headers } => Self::Sse { url, headers },
|
||||
McpTransport::Http { url, headers } => Self::Http { url, headers },
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<McpServerTransport> for McpTransport {
|
||||
fn from(t: McpServerTransport) -> Self {
|
||||
match t {
|
||||
McpServerTransport::Stdio { command, args, env } => McpTransport::Stdio { command, args, env },
|
||||
McpServerTransport::Sse { url, headers } => McpTransport::Sse { url, headers },
|
||||
McpServerTransport::Http { url, headers } => McpTransport::Http { url, headers },
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// McpTool — domain tool description
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Domain-layer MCP tool description.
|
||||
///
|
||||
/// Populated after a successful connection test (`tools/list` response).
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct McpTool {
|
||||
pub name: String,
|
||||
pub description: Option<String>,
|
||||
pub input_schema: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
impl From<McpToolResponse> for McpTool {
|
||||
fn from(r: McpToolResponse) -> Self {
|
||||
Self {
|
||||
name: r.name,
|
||||
description: r.description,
|
||||
input_schema: r.input_schema,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<McpTool> for McpToolResponse {
|
||||
fn from(t: McpTool) -> Self {
|
||||
McpToolResponse {
|
||||
name: t.name,
|
||||
description: t.description,
|
||||
input_schema: t.input_schema,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// McpServer — domain server model
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Domain-layer MCP server model.
|
||||
///
|
||||
/// Constructed from `McpServerRow` (DB) by parsing JSON fields into
|
||||
/// structured types. This is the primary type used across business logic.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct McpServer {
|
||||
/// Local-only integer primary key (cross-device classification: MCP is a
|
||||
/// host-local INTEGER entity). Carried through to `McpServerResponse.id`
|
||||
/// unchanged (number on the API boundary).
|
||||
pub id: i64,
|
||||
pub name: String,
|
||||
pub description: Option<String>,
|
||||
pub enabled: bool,
|
||||
pub transport: McpServerTransport,
|
||||
pub tools: Vec<McpTool>,
|
||||
pub last_test_status: McpServerStatus,
|
||||
pub last_connected: Option<TimestampMs>,
|
||||
pub original_json: Option<String>,
|
||||
pub builtin: bool,
|
||||
pub created_at: TimestampMs,
|
||||
pub updated_at: TimestampMs,
|
||||
}
|
||||
|
||||
impl McpServer {
|
||||
/// Converts a DB row into a domain model by parsing JSON fields.
|
||||
pub fn from_row(row: McpServerRow) -> Result<Self, McpError> {
|
||||
let transport = McpServerTransport::from_db(&row.transport_type, &row.transport_config)?;
|
||||
|
||||
let tools = match row.tools.as_deref() {
|
||||
Some(json_str) if !json_str.is_empty() => {
|
||||
let tool_responses: Vec<McpToolResponse> = serde_json::from_str(json_str).map_err(McpError::from)?;
|
||||
tool_responses.into_iter().map(McpTool::from).collect()
|
||||
}
|
||||
_ => Vec::new(),
|
||||
};
|
||||
|
||||
let last_test_status = parse_server_status(&row.last_test_status);
|
||||
|
||||
Ok(Self {
|
||||
id: row.id,
|
||||
name: row.name,
|
||||
description: row.description,
|
||||
enabled: row.enabled,
|
||||
transport,
|
||||
tools,
|
||||
last_test_status,
|
||||
last_connected: row.last_connected,
|
||||
original_json: row.original_json,
|
||||
builtin: row.builtin,
|
||||
created_at: row.created_at,
|
||||
updated_at: row.updated_at,
|
||||
})
|
||||
}
|
||||
|
||||
/// Converts to the API response DTO.
|
||||
pub fn into_response(self) -> McpServerResponse {
|
||||
let tools = if self.tools.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(self.tools.into_iter().map(McpToolResponse::from).collect())
|
||||
};
|
||||
|
||||
McpServerResponse {
|
||||
id: self.id,
|
||||
name: self.name,
|
||||
description: self.description,
|
||||
enabled: self.enabled,
|
||||
transport: self.transport.into(),
|
||||
tools,
|
||||
last_test_status: self.last_test_status,
|
||||
last_connected: self.last_connected,
|
||||
original_json: self.original_json,
|
||||
builtin: self.builtin,
|
||||
created_at: self.created_at,
|
||||
updated_at: self.updated_at,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Parse status string to enum, defaulting to `Disconnected` for unknown values.
|
||||
fn parse_server_status(s: &str) -> McpServerStatus {
|
||||
match s {
|
||||
"connected" => McpServerStatus::Connected,
|
||||
"disconnected" => McpServerStatus::Disconnected,
|
||||
"error" => McpServerStatus::Error,
|
||||
"testing" => McpServerStatus::Testing,
|
||||
_ => McpServerStatus::Disconnected,
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
// -- McpServerTransport ---------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn transport_type_string() {
|
||||
let stdio = McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
};
|
||||
assert_eq!(stdio.transport_type(), "stdio");
|
||||
|
||||
let sse = McpServerTransport::Sse {
|
||||
url: "http://x".into(),
|
||||
headers: HashMap::new(),
|
||||
};
|
||||
assert_eq!(sse.transport_type(), "sse");
|
||||
|
||||
let http = McpServerTransport::Http {
|
||||
url: "http://x".into(),
|
||||
headers: HashMap::new(),
|
||||
};
|
||||
assert_eq!(http.transport_type(), "http");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stdio_roundtrip_via_db() {
|
||||
let original = McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec!["-y".into(), "@test/server".into()],
|
||||
env: HashMap::from([("NODE_ENV".into(), "production".into())]),
|
||||
};
|
||||
let json = original.to_config_json().unwrap();
|
||||
let parsed = McpServerTransport::from_db("stdio", &json).unwrap();
|
||||
assert_eq!(parsed, original);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sse_roundtrip_via_db() {
|
||||
let original = McpServerTransport::Sse {
|
||||
url: "https://example.com/sse".into(),
|
||||
headers: HashMap::from([("Authorization".into(), "Bearer tok".into())]),
|
||||
};
|
||||
let json = original.to_config_json().unwrap();
|
||||
let parsed = McpServerTransport::from_db("sse", &json).unwrap();
|
||||
assert_eq!(parsed, original);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn http_roundtrip_via_db() {
|
||||
let original = McpServerTransport::Http {
|
||||
url: "https://example.com/mcp".into(),
|
||||
headers: HashMap::new(),
|
||||
};
|
||||
let json = original.to_config_json().unwrap();
|
||||
let parsed = McpServerTransport::from_db("http", &json).unwrap();
|
||||
assert_eq!(parsed, original);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_db_stdio_minimal() {
|
||||
let json = r#"{"command":"node"}"#;
|
||||
let t = McpServerTransport::from_db("stdio", json).unwrap();
|
||||
assert_eq!(
|
||||
t,
|
||||
McpServerTransport::Stdio {
|
||||
command: "node".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_db_unknown_type_fails() {
|
||||
let result = McpServerTransport::from_db("websocket", "{}");
|
||||
assert!(matches!(result, Err(McpError::InvalidTransport(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_db_stdio_missing_command_fails() {
|
||||
let result = McpServerTransport::from_db("stdio", r#"{"args":[]}"#);
|
||||
assert!(matches!(result, Err(McpError::InvalidTransport(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_db_sse_missing_url_fails() {
|
||||
let result = McpServerTransport::from_db("sse", r#"{"headers":{}}"#);
|
||||
assert!(matches!(result, Err(McpError::InvalidTransport(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_db_http_missing_url_fails() {
|
||||
let result = McpServerTransport::from_db("http", r#"{"headers":{}}"#);
|
||||
assert!(matches!(result, Err(McpError::InvalidTransport(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_db_invalid_json_fails() {
|
||||
let result = McpServerTransport::from_db("stdio", "not json");
|
||||
assert!(matches!(result, Err(McpError::Json(_))));
|
||||
}
|
||||
|
||||
// -- API transport conversions --------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn api_transport_roundtrip_stdio() {
|
||||
let domain = McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec!["-y".into()],
|
||||
env: HashMap::from([("K".into(), "V".into())]),
|
||||
};
|
||||
let api: McpTransport = domain.clone().into();
|
||||
let back: McpServerTransport = api.into();
|
||||
assert_eq!(back, domain);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn api_transport_roundtrip_sse() {
|
||||
let domain = McpServerTransport::Sse {
|
||||
url: "http://x".into(),
|
||||
headers: HashMap::from([("H".into(), "V".into())]),
|
||||
};
|
||||
let api: McpTransport = domain.clone().into();
|
||||
let back: McpServerTransport = api.into();
|
||||
assert_eq!(back, domain);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn api_transport_roundtrip_http() {
|
||||
let domain = McpServerTransport::Http {
|
||||
url: "http://x".into(),
|
||||
headers: HashMap::new(),
|
||||
};
|
||||
let api: McpTransport = domain.clone().into();
|
||||
let back: McpServerTransport = api.into();
|
||||
assert_eq!(back, domain);
|
||||
}
|
||||
|
||||
// -- McpTool conversions --------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn tool_from_response() {
|
||||
let resp = McpToolResponse {
|
||||
name: "read_file".into(),
|
||||
description: Some("Read a file".into()),
|
||||
input_schema: Some(serde_json::json!({"type": "object"})),
|
||||
};
|
||||
let tool = McpTool::from(resp);
|
||||
assert_eq!(tool.name, "read_file");
|
||||
assert_eq!(tool.description.as_deref(), Some("Read a file"));
|
||||
assert!(tool.input_schema.is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tool_to_response() {
|
||||
let tool = McpTool {
|
||||
name: "write_file".into(),
|
||||
description: None,
|
||||
input_schema: None,
|
||||
};
|
||||
let resp = McpToolResponse::from(tool);
|
||||
assert_eq!(resp.name, "write_file");
|
||||
assert!(resp.description.is_none());
|
||||
}
|
||||
|
||||
// -- McpServer::from_row --------------------------------------------------
|
||||
|
||||
fn make_test_row(transport_type: &str, transport_config: &str, tools: Option<&str>, status: &str) -> McpServerRow {
|
||||
McpServerRow {
|
||||
id: 123,
|
||||
name: "test-server".into(),
|
||||
description: Some("A test server".into()),
|
||||
enabled: true,
|
||||
transport_type: transport_type.into(),
|
||||
transport_config: transport_config.into(),
|
||||
tools: tools.map(String::from),
|
||||
last_test_status: status.into(),
|
||||
last_connected: Some(1000),
|
||||
original_json: None,
|
||||
builtin: false,
|
||||
deleted_at: None,
|
||||
created_at: 500,
|
||||
updated_at: 600,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_row_stdio_with_tools() {
|
||||
let row = make_test_row(
|
||||
"stdio",
|
||||
r#"{"command":"npx","args":["-y","@test/server"],"env":{"K":"V"}}"#,
|
||||
Some(r#"[{"name":"read","description":"Read file"}]"#),
|
||||
"connected",
|
||||
);
|
||||
let server = McpServer::from_row(row).unwrap();
|
||||
|
||||
assert_eq!(server.id, 123);
|
||||
assert_eq!(server.name, "test-server");
|
||||
assert!(server.enabled);
|
||||
assert_eq!(server.last_test_status, McpServerStatus::Connected);
|
||||
assert_eq!(server.tools.len(), 1);
|
||||
assert_eq!(server.tools[0].name, "read");
|
||||
match &server.transport {
|
||||
McpServerTransport::Stdio { command, args, env } => {
|
||||
assert_eq!(command, "npx");
|
||||
assert_eq!(args, &["-y", "@test/server"]);
|
||||
assert_eq!(env.get("K").unwrap(), "V");
|
||||
}
|
||||
_ => panic!("expected Stdio"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_row_http_no_tools() {
|
||||
let row = make_test_row(
|
||||
"http",
|
||||
r#"{"url":"https://example.com/mcp","headers":{}}"#,
|
||||
None,
|
||||
"disconnected",
|
||||
);
|
||||
let server = McpServer::from_row(row).unwrap();
|
||||
|
||||
assert!(server.tools.is_empty());
|
||||
assert_eq!(server.last_test_status, McpServerStatus::Disconnected);
|
||||
match &server.transport {
|
||||
McpServerTransport::Http { url, headers } => {
|
||||
assert_eq!(url, "https://example.com/mcp");
|
||||
assert!(headers.is_empty());
|
||||
}
|
||||
_ => panic!("expected Http"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_row_empty_tools_string() {
|
||||
let row = make_test_row("stdio", r#"{"command":"node"}"#, Some(""), "error");
|
||||
let server = McpServer::from_row(row).unwrap();
|
||||
assert!(server.tools.is_empty());
|
||||
assert_eq!(server.last_test_status, McpServerStatus::Error);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_row_unknown_status_defaults_to_disconnected() {
|
||||
let row = make_test_row("stdio", r#"{"command":"node"}"#, None, "unknown_status");
|
||||
let server = McpServer::from_row(row).unwrap();
|
||||
assert_eq!(server.last_test_status, McpServerStatus::Disconnected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_row_invalid_transport_config_fails() {
|
||||
let row = make_test_row("stdio", "not json", None, "disconnected");
|
||||
let result = McpServer::from_row(row);
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_row_invalid_tools_json_fails() {
|
||||
let row = make_test_row("stdio", r#"{"command":"node"}"#, Some("not json"), "disconnected");
|
||||
let result = McpServer::from_row(row);
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
// -- McpServer::into_response ---------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn into_response_with_tools() {
|
||||
let server = McpServer {
|
||||
id: 1,
|
||||
name: "test".into(),
|
||||
description: None,
|
||||
enabled: true,
|
||||
transport: McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
},
|
||||
tools: vec![McpTool {
|
||||
name: "read_file".into(),
|
||||
description: None,
|
||||
input_schema: None,
|
||||
}],
|
||||
last_test_status: McpServerStatus::Connected,
|
||||
last_connected: Some(1000),
|
||||
original_json: None,
|
||||
builtin: false,
|
||||
created_at: 500,
|
||||
updated_at: 600,
|
||||
};
|
||||
let resp = server.into_response();
|
||||
assert_eq!(resp.id, 1);
|
||||
assert!(resp.tools.is_some());
|
||||
assert_eq!(resp.tools.unwrap().len(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn into_response_empty_tools_is_none() {
|
||||
let server = McpServer {
|
||||
id: 2,
|
||||
name: "test".into(),
|
||||
description: Some("desc".into()),
|
||||
enabled: false,
|
||||
transport: McpServerTransport::Http {
|
||||
url: "http://x".into(),
|
||||
headers: HashMap::new(),
|
||||
},
|
||||
tools: vec![],
|
||||
last_test_status: McpServerStatus::Disconnected,
|
||||
last_connected: None,
|
||||
original_json: None,
|
||||
builtin: false,
|
||||
created_at: 500,
|
||||
updated_at: 600,
|
||||
};
|
||||
let resp = server.into_response();
|
||||
assert!(resp.tools.is_none());
|
||||
assert_eq!(resp.description.as_deref(), Some("desc"));
|
||||
}
|
||||
|
||||
// -- parse_server_status --------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn parse_all_statuses() {
|
||||
assert_eq!(parse_server_status("connected"), McpServerStatus::Connected);
|
||||
assert_eq!(parse_server_status("disconnected"), McpServerStatus::Disconnected);
|
||||
assert_eq!(parse_server_status("error"), McpServerStatus::Error);
|
||||
assert_eq!(parse_server_status("testing"), McpServerStatus::Testing);
|
||||
assert_eq!(parse_server_status("garbage"), McpServerStatus::Disconnected);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,192 @@
|
||||
//! Integration tests for McpAgentAdapter trait and DetectedServer.
|
||||
//!
|
||||
//! Uses a mock adapter to verify the trait's public API contract:
|
||||
//! object safety, install/detect/remove lifecycle, and error cases.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use nomifun_common::McpSource;
|
||||
use nomifun_mcp::{DetectedServer, McpAgentAdapter, McpError, McpServerTransport};
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Mock adapter (in-memory, for integration tests)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
struct InMemoryAdapter {
|
||||
source: McpSource,
|
||||
installed: bool,
|
||||
servers: Mutex<Vec<DetectedServer>>,
|
||||
}
|
||||
|
||||
impl InMemoryAdapter {
|
||||
fn new(source: McpSource, installed: bool) -> Self {
|
||||
Self {
|
||||
source,
|
||||
installed,
|
||||
servers: Mutex::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl McpAgentAdapter for InMemoryAdapter {
|
||||
fn source(&self) -> McpSource {
|
||||
self.source
|
||||
}
|
||||
|
||||
async fn is_installed(&self) -> Result<bool, McpError> {
|
||||
Ok(self.installed)
|
||||
}
|
||||
|
||||
async fn detect_existing(&self) -> Result<Vec<DetectedServer>, McpError> {
|
||||
if !self.installed {
|
||||
return Err(McpError::AgentNotInstalled(format!("{:?}", self.source)));
|
||||
}
|
||||
Ok(self.servers.lock().unwrap().clone())
|
||||
}
|
||||
|
||||
async fn install_server(&self, name: &str, transport: &McpServerTransport) -> Result<(), McpError> {
|
||||
if !self.installed {
|
||||
return Err(McpError::AgentNotInstalled(format!("{:?}", self.source)));
|
||||
}
|
||||
let mut servers = self.servers.lock().unwrap();
|
||||
servers.retain(|s| s.name != name);
|
||||
servers.push(DetectedServer {
|
||||
name: name.to_owned(),
|
||||
transport: transport.clone(),
|
||||
importable: true,
|
||||
import_skip_reason: None,
|
||||
});
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn remove_server(&self, name: &str) -> Result<(), McpError> {
|
||||
if !self.installed {
|
||||
return Err(McpError::AgentNotInstalled(format!("{:?}", self.source)));
|
||||
}
|
||||
let mut servers = self.servers.lock().unwrap();
|
||||
servers.retain(|s| s.name != name);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn trait_object_safety_with_arc() {
|
||||
let adapter: Arc<dyn McpAgentAdapter> = Arc::new(InMemoryAdapter::new(McpSource::Claude, true));
|
||||
|
||||
assert_eq!(adapter.source(), McpSource::Claude);
|
||||
assert!(adapter.is_installed().await.unwrap());
|
||||
assert!(adapter.detect_existing().await.unwrap().is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn full_lifecycle_install_detect_remove() {
|
||||
let adapter = InMemoryAdapter::new(McpSource::Gemini, true);
|
||||
|
||||
// Install two servers
|
||||
let t1 = McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec!["-y".into(), "server-a".into()],
|
||||
env: HashMap::new(),
|
||||
};
|
||||
let t2 = McpServerTransport::Http {
|
||||
url: "https://example.com/mcp".into(),
|
||||
headers: HashMap::from([("Auth".into(), "Bearer x".into())]),
|
||||
};
|
||||
|
||||
adapter.install_server("server-a", &t1).await.unwrap();
|
||||
adapter.install_server("server-b", &t2).await.unwrap();
|
||||
|
||||
// Detect both
|
||||
let detected = adapter.detect_existing().await.unwrap();
|
||||
assert_eq!(detected.len(), 2);
|
||||
|
||||
let names: Vec<&str> = detected.iter().map(|s| s.name.as_str()).collect();
|
||||
assert!(names.contains(&"server-a"));
|
||||
assert!(names.contains(&"server-b"));
|
||||
|
||||
// Remove one
|
||||
adapter.remove_server("server-a").await.unwrap();
|
||||
let detected = adapter.detect_existing().await.unwrap();
|
||||
assert_eq!(detected.len(), 1);
|
||||
assert_eq!(detected[0].name, "server-b");
|
||||
|
||||
// Remove the other
|
||||
adapter.remove_server("server-b").await.unwrap();
|
||||
let detected = adapter.detect_existing().await.unwrap();
|
||||
assert!(detected.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn install_replaces_existing_by_name() {
|
||||
let adapter = InMemoryAdapter::new(McpSource::Qwen, true);
|
||||
|
||||
let t1 = McpServerTransport::Stdio {
|
||||
command: "old-cmd".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
};
|
||||
let t2 = McpServerTransport::Stdio {
|
||||
command: "new-cmd".into(),
|
||||
args: vec!["--flag".into()],
|
||||
env: HashMap::new(),
|
||||
};
|
||||
|
||||
adapter.install_server("my-server", &t1).await.unwrap();
|
||||
adapter.install_server("my-server", &t2).await.unwrap();
|
||||
|
||||
let detected = adapter.detect_existing().await.unwrap();
|
||||
assert_eq!(detected.len(), 1);
|
||||
assert_eq!(detected[0].transport, t2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn remove_nonexistent_is_idempotent() {
|
||||
let adapter = InMemoryAdapter::new(McpSource::Nomifun, true);
|
||||
// Should succeed without error
|
||||
adapter.remove_server("does-not-exist").await.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn not_installed_errors() {
|
||||
let adapter = InMemoryAdapter::new(McpSource::Codex, false);
|
||||
|
||||
assert!(!adapter.is_installed().await.unwrap());
|
||||
|
||||
let err = adapter.detect_existing().await.unwrap_err();
|
||||
assert!(matches!(err, McpError::AgentNotInstalled(_)));
|
||||
|
||||
let transport = McpServerTransport::Stdio {
|
||||
command: "x".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
};
|
||||
let err = adapter.install_server("s", &transport).await.unwrap_err();
|
||||
assert!(matches!(err, McpError::AgentNotInstalled(_)));
|
||||
|
||||
let err = adapter.remove_server("s").await.unwrap_err();
|
||||
assert!(matches!(err, McpError::AgentNotInstalled(_)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn multiple_adapters_independent() {
|
||||
let claude: Arc<dyn McpAgentAdapter> = Arc::new(InMemoryAdapter::new(McpSource::Claude, true));
|
||||
let gemini: Arc<dyn McpAgentAdapter> = Arc::new(InMemoryAdapter::new(McpSource::Gemini, true));
|
||||
|
||||
let transport = McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
};
|
||||
|
||||
claude.install_server("shared-server", &transport).await.unwrap();
|
||||
|
||||
// Claude has the server, Gemini does not
|
||||
assert_eq!(claude.detect_existing().await.unwrap().len(), 1);
|
||||
assert!(gemini.detect_existing().await.unwrap().is_empty());
|
||||
}
|
||||
@@ -0,0 +1,303 @@
|
||||
//! Integration tests for McpConnectionTestService.
|
||||
//!
|
||||
//! Tests from test-plan §2 (Connection Test):
|
||||
//! - CT-3: Command not found (ENOENT)
|
||||
//! - CT-4: URL not reachable
|
||||
//! - CT-5: Needs OAuth authentication (401)
|
||||
//! - CT-6: Timeout
|
||||
//! - SSE auth probe (M-33 coverage)
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::time::Duration;
|
||||
|
||||
use nomifun_mcp::McpConnectionTestService;
|
||||
use nomifun_mcp::McpServerTransport;
|
||||
|
||||
fn make_service() -> McpConnectionTestService {
|
||||
McpConnectionTestService::new(reqwest::Client::new())
|
||||
}
|
||||
|
||||
fn make_service_with_timeout(timeout: Duration) -> McpConnectionTestService {
|
||||
McpConnectionTestService::new(reqwest::Client::new()).with_timeout(timeout)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// CT-3: Command not found (ENOENT)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn stdio_nonexistent_command_returns_not_found_error() {
|
||||
let svc = make_service();
|
||||
let transport = McpServerTransport::Stdio {
|
||||
command: "nonexistent-mcp-cmd-xyz-12345".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
};
|
||||
|
||||
let result = svc.test_connection("test-server", &transport).await;
|
||||
|
||||
assert!(!result.success);
|
||||
let error = result.error.as_deref().unwrap();
|
||||
assert!(
|
||||
error.contains("Command not found"),
|
||||
"expected 'Command not found' in: {error}"
|
||||
);
|
||||
assert!(result.tools.is_none());
|
||||
assert!(result.needs_auth.is_none());
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// CT-4: URL not reachable
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn http_unreachable_url_returns_connection_error() {
|
||||
let svc = make_service_with_timeout(Duration::from_secs(5));
|
||||
let transport = McpServerTransport::Http {
|
||||
url: "http://127.0.0.1:1/mcp-unreachable".into(),
|
||||
headers: HashMap::new(),
|
||||
};
|
||||
|
||||
let result = svc.test_connection("test-http", &transport).await;
|
||||
|
||||
assert!(!result.success);
|
||||
let error = result.error.as_deref().unwrap();
|
||||
assert!(
|
||||
error.contains("Connection failed"),
|
||||
"expected connection failure in: {error}"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn sse_unreachable_url_returns_connection_error() {
|
||||
let svc = make_service_with_timeout(Duration::from_secs(5));
|
||||
let transport = McpServerTransport::Sse {
|
||||
url: "http://127.0.0.1:1/sse-unreachable".into(),
|
||||
headers: HashMap::new(),
|
||||
};
|
||||
|
||||
let result = svc.test_connection("test-sse", &transport).await;
|
||||
|
||||
assert!(!result.success);
|
||||
let error = result.error.as_deref().unwrap();
|
||||
assert!(
|
||||
error.contains("Connection failed"),
|
||||
"expected connection failure in: {error}"
|
||||
);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// CT-5: HTTP 401 Unauthorized -> needsAuth
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn http_401_returns_needs_auth() {
|
||||
// Spin up a mock server that returns 401 with WWW-Authenticate
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let addr = listener.local_addr().unwrap();
|
||||
|
||||
let server_handle = tokio::spawn(async move {
|
||||
let app = axum::Router::new().route(
|
||||
"/mcp",
|
||||
axum::routing::post(|| async {
|
||||
(
|
||||
axum::http::StatusCode::UNAUTHORIZED,
|
||||
[(axum::http::header::WWW_AUTHENTICATE, "Bearer realm=\"mcp-server\"")],
|
||||
"",
|
||||
)
|
||||
}),
|
||||
);
|
||||
axum::serve(listener, app).await.unwrap();
|
||||
});
|
||||
|
||||
let svc = make_service();
|
||||
let transport = McpServerTransport::Http {
|
||||
url: format!("http://{}/mcp", addr),
|
||||
headers: HashMap::new(),
|
||||
};
|
||||
|
||||
let result = svc.test_connection("auth-server", &transport).await;
|
||||
|
||||
assert!(!result.success);
|
||||
assert_eq!(result.needs_auth, Some(true));
|
||||
assert!(result.auth_method.is_some());
|
||||
assert!(result.www_authenticate.is_some());
|
||||
assert!(result.error.is_none());
|
||||
|
||||
server_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn sse_401_returns_needs_auth() {
|
||||
// Spin up a mock server that returns 401 for GET
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let addr = listener.local_addr().unwrap();
|
||||
|
||||
let server_handle = tokio::spawn(async move {
|
||||
let app = axum::Router::new().route(
|
||||
"/sse",
|
||||
axum::routing::get(|| async {
|
||||
(
|
||||
axum::http::StatusCode::UNAUTHORIZED,
|
||||
[(axum::http::header::WWW_AUTHENTICATE, "Bearer realm=\"mcp-sse\"")],
|
||||
"",
|
||||
)
|
||||
}),
|
||||
);
|
||||
axum::serve(listener, app).await.unwrap();
|
||||
});
|
||||
|
||||
let svc = make_service();
|
||||
let transport = McpServerTransport::Sse {
|
||||
url: format!("http://{}/sse", addr),
|
||||
headers: HashMap::new(),
|
||||
};
|
||||
|
||||
let result = svc.test_connection("sse-auth", &transport).await;
|
||||
|
||||
assert!(!result.success);
|
||||
assert_eq!(result.needs_auth, Some(true));
|
||||
assert!(result.www_authenticate.is_some());
|
||||
|
||||
server_handle.abort();
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// CT-6: Timeout
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn stdio_timeout_returns_timeout_error() {
|
||||
// Use `sleep` which produces no stdout — our protocol read will block
|
||||
let svc = make_service_with_timeout(Duration::from_secs(1));
|
||||
let transport = McpServerTransport::Stdio {
|
||||
command: "sleep".into(),
|
||||
args: vec!["60".into()],
|
||||
env: HashMap::new(),
|
||||
};
|
||||
|
||||
let result = svc.test_connection("timeout-server", &transport).await;
|
||||
|
||||
assert!(!result.success);
|
||||
let error = result.error.as_deref().unwrap();
|
||||
assert!(error.contains("timed out"), "expected timeout in: {error}");
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// HTTP non-success status
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn http_500_returns_error_with_status() {
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let addr = listener.local_addr().unwrap();
|
||||
|
||||
let server_handle = tokio::spawn(async move {
|
||||
let app = axum::Router::new().route(
|
||||
"/mcp",
|
||||
axum::routing::post(|| async { axum::http::StatusCode::INTERNAL_SERVER_ERROR }),
|
||||
);
|
||||
axum::serve(listener, app).await.unwrap();
|
||||
});
|
||||
|
||||
let svc = make_service();
|
||||
let transport = McpServerTransport::Http {
|
||||
url: format!("http://{}/mcp", addr),
|
||||
headers: HashMap::new(),
|
||||
};
|
||||
|
||||
let result = svc.test_connection("error-server", &transport).await;
|
||||
|
||||
assert!(!result.success);
|
||||
let error = result.error.as_deref().unwrap();
|
||||
assert!(error.contains("500"), "expected HTTP 500 in: {error}");
|
||||
|
||||
server_handle.abort();
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// HTTP transport with custom headers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn http_custom_headers_are_sent() {
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let addr = listener.local_addr().unwrap();
|
||||
|
||||
let server_handle = tokio::spawn(async move {
|
||||
let app = axum::Router::new().route(
|
||||
"/mcp",
|
||||
axum::routing::post(|headers: axum::http::HeaderMap| async move {
|
||||
// Verify the custom header was received
|
||||
if headers.get("x-api-key").and_then(|v| v.to_str().ok()) == Some("secret") {
|
||||
// Return a valid initialize response
|
||||
axum::Json(serde_json::json!({
|
||||
"jsonrpc": "2.0",
|
||||
"id": 1,
|
||||
"result": {
|
||||
"protocolVersion": "2024-11-05",
|
||||
"capabilities": {},
|
||||
"serverInfo": { "name": "test", "version": "1.0" }
|
||||
}
|
||||
}))
|
||||
} else {
|
||||
// Return error if header missing
|
||||
axum::Json(serde_json::json!({
|
||||
"jsonrpc": "2.0",
|
||||
"id": 1,
|
||||
"error": { "code": -1, "message": "Missing API key" }
|
||||
}))
|
||||
}
|
||||
}),
|
||||
);
|
||||
axum::serve(listener, app).await.unwrap();
|
||||
});
|
||||
|
||||
let svc = make_service();
|
||||
let mut headers = HashMap::new();
|
||||
headers.insert("X-Api-Key".into(), "secret".into());
|
||||
let transport = McpServerTransport::Http {
|
||||
url: format!("http://{}/mcp", addr),
|
||||
headers,
|
||||
};
|
||||
|
||||
let result = svc.test_connection("header-server", &transport).await;
|
||||
|
||||
// The server returns a valid initialize response for request id=1,
|
||||
// but the subsequent tools/list (id=2) will also hit the same handler.
|
||||
// Either way, the first request should succeed (no initialize error).
|
||||
// The tools/list might succeed or fail depending on how the mock handles id=2.
|
||||
// For this test, we just verify the custom header was sent (no "Missing API key" error).
|
||||
if let Some(ref error) = result.error {
|
||||
assert!(
|
||||
!error.contains("Missing API key"),
|
||||
"Custom header should have been sent"
|
||||
);
|
||||
}
|
||||
|
||||
server_handle.abort();
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Stdio with args and env
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn stdio_with_args_spawns_correctly() {
|
||||
// Use echo as a simple command that exits immediately
|
||||
// Since echo doesn't speak MCP, we expect a protocol error (not a spawn error)
|
||||
let svc = make_service_with_timeout(Duration::from_secs(3));
|
||||
let transport = McpServerTransport::Stdio {
|
||||
command: "echo".into(),
|
||||
args: vec!["hello".into()],
|
||||
env: HashMap::new(),
|
||||
};
|
||||
|
||||
let result = svc.test_connection("echo-server", &transport).await;
|
||||
|
||||
// echo outputs "hello\n" then exits — not valid JSON-RPC
|
||||
assert!(!result.success);
|
||||
let error = result.error.as_deref().unwrap();
|
||||
// Should be a protocol error, not a spawn error
|
||||
assert!(!error.contains("Command not found"), "echo should be found");
|
||||
}
|
||||
+66
@@ -0,0 +1,66 @@
|
||||
//! Isolated PATH-resolution coverage for stdio MCP connection tests.
|
||||
//!
|
||||
//! This file intentionally contains one test because it mutates process PATH
|
||||
//! to model the startup-enhanced GUI environment.
|
||||
|
||||
#![cfg(unix)]
|
||||
|
||||
use std::collections::HashMap;
|
||||
|
||||
use nomifun_mcp::{McpConnectionTestService, McpServerTransport};
|
||||
|
||||
#[tokio::test]
|
||||
async fn stdio_npx_resolves_from_enhanced_process_path() {
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
|
||||
let tmp = tempfile::TempDir::new().unwrap();
|
||||
let bin_dir = tmp.path().join("bin");
|
||||
std::fs::create_dir(&bin_dir).unwrap();
|
||||
|
||||
let fake_npx = bin_dir.join("npx");
|
||||
std::fs::write(
|
||||
&fake_npx,
|
||||
r#"#!/bin/sh
|
||||
while IFS= read -r line; do
|
||||
case "$line" in
|
||||
*'"id":1'*)
|
||||
printf '%s\n' '{"jsonrpc":"2.0","id":1,"result":{"protocolVersion":"2024-11-05","capabilities":{},"serverInfo":{"name":"fake-npx","version":"1.0.0"}}}'
|
||||
;;
|
||||
*'"id":2'*)
|
||||
printf '%s\n' '{"jsonrpc":"2.0","id":2,"result":{"tools":[]}}'
|
||||
exit 0
|
||||
;;
|
||||
esac
|
||||
done
|
||||
"#,
|
||||
)
|
||||
.unwrap();
|
||||
let mut perms = std::fs::metadata(&fake_npx).unwrap().permissions();
|
||||
perms.set_mode(0o755);
|
||||
std::fs::set_permissions(&fake_npx, perms).unwrap();
|
||||
|
||||
let original_path = std::env::var_os("PATH");
|
||||
unsafe {
|
||||
std::env::set_var("PATH", &bin_dir);
|
||||
}
|
||||
|
||||
let svc = McpConnectionTestService::new(reqwest::Client::new());
|
||||
let transport = McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
};
|
||||
|
||||
let result = svc.test_connection("fake-npx", &transport).await;
|
||||
|
||||
unsafe {
|
||||
if let Some(path) = original_path {
|
||||
std::env::set_var("PATH", path);
|
||||
} else {
|
||||
std::env::remove_var("PATH");
|
||||
}
|
||||
}
|
||||
|
||||
assert!(result.success, "expected fake npx MCP server to connect: {result:?}");
|
||||
assert!(result.tools.unwrap().is_empty());
|
||||
}
|
||||
@@ -0,0 +1,230 @@
|
||||
//! Integration tests for file-based MCP Agent adapters (Opencode, Nomi, Nomi).
|
||||
//!
|
||||
//! These tests exercise the real filesystem read/write logic using temp
|
||||
//! directories. CLI detection (`is_installed`, `which`) is NOT tested here
|
||||
//! because it depends on the host environment.
|
||||
//!
|
||||
//! For Nomi, we use a mock repository since it reads from the DB.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use nomifun_common::McpSource;
|
||||
use nomifun_mcp::{McpAgentAdapter, McpServerTransport, NomifunAdapter};
|
||||
|
||||
// ===========================================================================
|
||||
// Nomi adapter (DB-backed)
|
||||
// ===========================================================================
|
||||
|
||||
mod nomifun {
|
||||
use super::*;
|
||||
use nomifun_db::models::McpServerRow;
|
||||
use nomifun_db::{CreateMcpServerParams, DbError, IMcpServerRepository, UpdateMcpServerParams};
|
||||
|
||||
struct MockRepo {
|
||||
servers: tokio::sync::Mutex<Vec<McpServerRow>>,
|
||||
}
|
||||
|
||||
impl MockRepo {
|
||||
fn new(servers: Vec<McpServerRow>) -> Self {
|
||||
Self {
|
||||
servers: tokio::sync::Mutex::new(servers),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl IMcpServerRepository for MockRepo {
|
||||
async fn list(&self) -> Result<Vec<McpServerRow>, DbError> {
|
||||
Ok(self.servers.lock().await.clone())
|
||||
}
|
||||
|
||||
async fn find_by_id(&self, id: i64) -> Result<Option<McpServerRow>, DbError> {
|
||||
Ok(self.servers.lock().await.iter().find(|s| s.id == id).cloned())
|
||||
}
|
||||
|
||||
async fn find_by_name(&self, name: &str) -> Result<Option<McpServerRow>, DbError> {
|
||||
Ok(self.servers.lock().await.iter().find(|s| s.name == name).cloned())
|
||||
}
|
||||
|
||||
async fn create(&self, _p: CreateMcpServerParams<'_>) -> Result<McpServerRow, DbError> {
|
||||
unimplemented!()
|
||||
}
|
||||
|
||||
async fn update(&self, _id: i64, _p: UpdateMcpServerParams<'_>) -> Result<McpServerRow, DbError> {
|
||||
unimplemented!()
|
||||
}
|
||||
|
||||
async fn delete(&self, _id: i64) -> Result<(), DbError> {
|
||||
unimplemented!()
|
||||
}
|
||||
|
||||
async fn batch_upsert(&self, _s: &[CreateMcpServerParams<'_>]) -> Result<Vec<McpServerRow>, DbError> {
|
||||
unimplemented!()
|
||||
}
|
||||
|
||||
async fn update_status(
|
||||
&self,
|
||||
_id: i64,
|
||||
_s: &str,
|
||||
_lc: Option<nomifun_common::TimestampMs>,
|
||||
) -> Result<(), DbError> {
|
||||
unimplemented!()
|
||||
}
|
||||
|
||||
async fn update_tools(&self, _id: i64, _t: Option<&str>) -> Result<(), DbError> {
|
||||
unimplemented!()
|
||||
}
|
||||
}
|
||||
|
||||
fn make_row(name: &str, t_type: &str, t_config: &str) -> McpServerRow {
|
||||
McpServerRow {
|
||||
id: name.bytes().map(i64::from).sum::<i64>().max(1),
|
||||
name: name.to_owned(),
|
||||
description: None,
|
||||
enabled: true,
|
||||
transport_type: t_type.into(),
|
||||
transport_config: t_config.into(),
|
||||
tools: None,
|
||||
last_test_status: "disconnected".into(),
|
||||
last_connected: None,
|
||||
original_json: None,
|
||||
builtin: false,
|
||||
deleted_at: None,
|
||||
created_at: 1000,
|
||||
updated_at: 1000,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn source_is_nomifun() {
|
||||
let repo = Arc::new(MockRepo::new(vec![]));
|
||||
let adapter = NomifunAdapter::new(repo);
|
||||
assert_eq!(adapter.source(), McpSource::Nomifun);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn always_installed() {
|
||||
let repo = Arc::new(MockRepo::new(vec![]));
|
||||
let adapter = NomifunAdapter::new(repo);
|
||||
assert!(adapter.is_installed().await.unwrap());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn detect_returns_all_db_servers() {
|
||||
let rows = vec![
|
||||
make_row("stdio-srv", "stdio", r#"{"command":"npx","args":[]}"#),
|
||||
make_row("http-srv", "http", r#"{"url":"https://example.com/mcp","headers":{}}"#),
|
||||
make_row("sse-srv", "sse", r#"{"url":"https://example.com/sse","headers":{}}"#),
|
||||
];
|
||||
let repo = Arc::new(MockRepo::new(rows));
|
||||
let adapter = NomifunAdapter::new(repo);
|
||||
|
||||
let servers = adapter.detect_existing().await.unwrap();
|
||||
assert_eq!(servers.len(), 3);
|
||||
assert_eq!(servers[0].name, "stdio-srv");
|
||||
assert_eq!(servers[1].name, "http-srv");
|
||||
assert_eq!(servers[2].name, "sse-srv");
|
||||
|
||||
assert!(matches!(servers[0].transport, McpServerTransport::Stdio { .. }));
|
||||
assert!(matches!(servers[1].transport, McpServerTransport::Http { .. }));
|
||||
assert!(matches!(servers[2].transport, McpServerTransport::Sse { .. }));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn detect_empty_db_returns_empty() {
|
||||
let repo = Arc::new(MockRepo::new(vec![]));
|
||||
let adapter = NomifunAdapter::new(repo);
|
||||
let servers = adapter.detect_existing().await.unwrap();
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn install_is_noop() {
|
||||
let repo = Arc::new(MockRepo::new(vec![]));
|
||||
let adapter = NomifunAdapter::new(repo.clone());
|
||||
|
||||
let transport = McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
};
|
||||
adapter.install_server("test", &transport).await.unwrap();
|
||||
|
||||
// DB should still be empty since install is a no-op
|
||||
let servers = adapter.detect_existing().await.unwrap();
|
||||
assert!(servers.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn remove_is_noop() {
|
||||
let rows = vec![make_row("srv", "stdio", r#"{"command":"npx","args":[]}"#)];
|
||||
let repo = Arc::new(MockRepo::new(rows));
|
||||
let adapter = NomifunAdapter::new(repo);
|
||||
|
||||
adapter.remove_server("srv").await.unwrap();
|
||||
|
||||
// Server should still be in DB since remove is a no-op
|
||||
let servers = adapter.detect_existing().await.unwrap();
|
||||
assert_eq!(servers.len(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn trait_object_safety() {
|
||||
let repo = Arc::new(MockRepo::new(vec![]));
|
||||
let adapter: Arc<dyn McpAgentAdapter> = Arc::new(NomifunAdapter::new(repo));
|
||||
assert_eq!(adapter.source(), McpSource::Nomifun);
|
||||
assert!(adapter.is_installed().await.unwrap());
|
||||
}
|
||||
}
|
||||
|
||||
// ===========================================================================
|
||||
// Opencode adapter (filesystem-backed)
|
||||
// ===========================================================================
|
||||
|
||||
// Note: Full lifecycle tests for Opencode require controlling the config
|
||||
// directory path, which the adapter currently derives from `dirs::config_dir()`.
|
||||
// The unit tests in opencode.rs thoroughly cover parsing and serialization.
|
||||
// Here we verify that the adapter implements the trait correctly and that
|
||||
// the public API surface is accessible from outside the crate.
|
||||
|
||||
mod opencode {
|
||||
use super::*;
|
||||
use nomifun_mcp::OpencodeAdapter;
|
||||
|
||||
#[test]
|
||||
fn source_is_opencode() {
|
||||
assert_eq!(OpencodeAdapter.source(), McpSource::OpenCode);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn trait_object_safety() {
|
||||
let adapter: Box<dyn McpAgentAdapter> = Box::new(OpencodeAdapter);
|
||||
assert_eq!(adapter.source(), McpSource::OpenCode);
|
||||
}
|
||||
}
|
||||
|
||||
// ===========================================================================
|
||||
// Nomi adapter (CLI + TOML-backed)
|
||||
// ===========================================================================
|
||||
|
||||
// Note: Full lifecycle tests for Nomi require the `nomi` CLI to be
|
||||
// installed (for `--config-path`). The unit tests in nomi.rs thoroughly
|
||||
// cover TOML parsing, serialization, and roundtrip behavior. Here we
|
||||
// verify the public API surface.
|
||||
|
||||
mod nomi {
|
||||
use super::*;
|
||||
use nomifun_mcp::NomiAdapter;
|
||||
|
||||
#[test]
|
||||
fn source_is_nomi() {
|
||||
assert_eq!(NomiAdapter.source(), McpSource::Nomi);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn trait_object_safety() {
|
||||
let adapter: Box<dyn McpAgentAdapter> = Box::new(NomiAdapter);
|
||||
assert_eq!(adapter.source(), McpSource::Nomi);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,261 @@
|
||||
//! Integration tests for McpOAuthService with real SQLite.
|
||||
//!
|
||||
//! Tests from test-plan §4 (OAuth) at the service layer.
|
||||
//! These tests exercise check_status, logout, get_authenticated_servers,
|
||||
//! and get_token with a real DB. The full login flow (browser + callback)
|
||||
//! cannot be tested end-to-end here; it requires a mock OAuth server.
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use nomifun_db::{IOAuthTokenRepository, SqliteOAuthTokenRepository, UpsertOAuthTokenParams};
|
||||
use nomifun_mcp::McpOAuthService;
|
||||
|
||||
async fn make_service() -> (McpOAuthService, Arc<dyn IOAuthTokenRepository>) {
|
||||
let db = nomifun_db::init_database_memory().await.unwrap();
|
||||
let repo: Arc<dyn IOAuthTokenRepository> = Arc::new(SqliteOAuthTokenRepository::new(db.pool().clone()));
|
||||
let svc = McpOAuthService::new(repo.clone(), reqwest::Client::new());
|
||||
// Keep db alive by leaking it (integration test only).
|
||||
std::mem::forget(db);
|
||||
(svc, repo)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// OA-1: Unauthenticated server returns false
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn check_status_unauthenticated_returns_false() {
|
||||
let (svc, _repo) = make_service().await;
|
||||
let status = svc.check_oauth_status("https://new-server.example.com").await.unwrap();
|
||||
assert!(!status.authenticated);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// OA-2: Authenticated server returns true
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn check_status_authenticated_returns_true() {
|
||||
let (svc, repo) = make_service().await;
|
||||
|
||||
// Seed a valid token.
|
||||
repo.upsert(UpsertOAuthTokenParams {
|
||||
server_url: "https://mcp.example.com",
|
||||
access_token: "access_123",
|
||||
refresh_token: Some("refresh_456"),
|
||||
token_type: "bearer",
|
||||
// Expires in the far future.
|
||||
expires_at: Some(nomifun_common::now_ms() + 3_600_000),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let status = svc.check_oauth_status("https://mcp.example.com").await.unwrap();
|
||||
assert!(status.authenticated);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// OA-2b: Expired token treated as unauthenticated
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn check_status_expired_token_returns_false() {
|
||||
let (svc, repo) = make_service().await;
|
||||
|
||||
repo.upsert(UpsertOAuthTokenParams {
|
||||
server_url: "https://expired.example.com",
|
||||
access_token: "old_token",
|
||||
refresh_token: None,
|
||||
token_type: "bearer",
|
||||
// Already expired.
|
||||
expires_at: Some(1000),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let status = svc.check_oauth_status("https://expired.example.com").await.unwrap();
|
||||
assert!(!status.authenticated);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// OA-2c: Token with no expiry treated as valid
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn check_status_no_expiry_treated_as_valid() {
|
||||
let (svc, repo) = make_service().await;
|
||||
|
||||
repo.upsert(UpsertOAuthTokenParams {
|
||||
server_url: "https://no-expiry.example.com",
|
||||
access_token: "no_exp_token",
|
||||
refresh_token: None,
|
||||
token_type: "bearer",
|
||||
expires_at: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let status = svc.check_oauth_status("https://no-expiry.example.com").await.unwrap();
|
||||
assert!(status.authenticated);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// OA-3: Get all authenticated URLs
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn get_authenticated_servers_returns_all_urls() {
|
||||
let (svc, repo) = make_service().await;
|
||||
|
||||
repo.upsert(UpsertOAuthTokenParams {
|
||||
server_url: "https://a.example.com",
|
||||
access_token: "tok_a",
|
||||
refresh_token: None,
|
||||
token_type: "bearer",
|
||||
expires_at: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
repo.upsert(UpsertOAuthTokenParams {
|
||||
server_url: "https://b.example.com",
|
||||
access_token: "tok_b",
|
||||
refresh_token: None,
|
||||
token_type: "bearer",
|
||||
expires_at: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let urls = svc.get_authenticated_servers().await.unwrap();
|
||||
assert_eq!(urls.len(), 2);
|
||||
assert!(urls.contains(&"https://a.example.com".to_string()));
|
||||
assert!(urls.contains(&"https://b.example.com".to_string()));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn get_authenticated_servers_empty_when_no_tokens() {
|
||||
let (svc, _repo) = make_service().await;
|
||||
let urls = svc.get_authenticated_servers().await.unwrap();
|
||||
assert!(urls.is_empty());
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// OA-5: Login with invalid URL (no OAuth endpoints discoverable)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn login_invalid_url_returns_error() {
|
||||
let (svc, _repo) = make_service().await;
|
||||
// This URL won't have .well-known endpoints.
|
||||
let result = svc.login("https://127.0.0.1:1").await;
|
||||
// Should return an McpError::OAuth about discovery failure.
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// OA-6: Logout deletes stored token
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn logout_deletes_stored_token() {
|
||||
let (svc, repo) = make_service().await;
|
||||
|
||||
repo.upsert(UpsertOAuthTokenParams {
|
||||
server_url: "https://logout.example.com",
|
||||
access_token: "to_delete",
|
||||
refresh_token: None,
|
||||
token_type: "bearer",
|
||||
expires_at: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Verify token exists.
|
||||
let status = svc.check_oauth_status("https://logout.example.com").await.unwrap();
|
||||
assert!(status.authenticated);
|
||||
|
||||
// Logout.
|
||||
svc.logout("https://logout.example.com").await.unwrap();
|
||||
|
||||
// Verify token is gone.
|
||||
let status = svc.check_oauth_status("https://logout.example.com").await.unwrap();
|
||||
assert!(!status.authenticated);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// OA-7: Logout is idempotent for non-authenticated URL
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn logout_idempotent_for_unauthenticated() {
|
||||
let (svc, _repo) = make_service().await;
|
||||
// Should not error.
|
||||
svc.logout("https://never-authed.example.com").await.unwrap();
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// get_token tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn get_token_returns_none_for_unknown_url() {
|
||||
let (svc, _repo) = make_service().await;
|
||||
let token = svc.get_token("https://unknown.example.com").await.unwrap();
|
||||
assert!(token.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn get_token_returns_access_token_when_valid() {
|
||||
let (svc, repo) = make_service().await;
|
||||
|
||||
repo.upsert(UpsertOAuthTokenParams {
|
||||
server_url: "https://valid.example.com",
|
||||
access_token: "my_access_token",
|
||||
refresh_token: None,
|
||||
token_type: "bearer",
|
||||
expires_at: Some(nomifun_common::now_ms() + 3_600_000),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let token = svc.get_token("https://valid.example.com").await.unwrap();
|
||||
assert_eq!(token.as_deref(), Some("my_access_token"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn get_token_returns_expired_token_when_no_refresh_token() {
|
||||
let (svc, repo) = make_service().await;
|
||||
|
||||
repo.upsert(UpsertOAuthTokenParams {
|
||||
server_url: "https://expired.example.com",
|
||||
access_token: "old_access",
|
||||
refresh_token: None,
|
||||
token_type: "bearer",
|
||||
expires_at: Some(1000),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// With no refresh_token, returns the expired token as-is.
|
||||
let token = svc.get_token("https://expired.example.com").await.unwrap();
|
||||
assert_eq!(token.as_deref(), Some("old_access"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn get_token_returns_no_expiry_token() {
|
||||
let (svc, repo) = make_service().await;
|
||||
|
||||
repo.upsert(UpsertOAuthTokenParams {
|
||||
server_url: "https://noexp.example.com",
|
||||
access_token: "forever_token",
|
||||
refresh_token: None,
|
||||
token_type: "bearer",
|
||||
expires_at: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let token = svc.get_token("https://noexp.example.com").await.unwrap();
|
||||
assert_eq!(token.as_deref(), Some("forever_token"));
|
||||
}
|
||||
@@ -0,0 +1,351 @@
|
||||
//! Integration tests for McpConfigService with real SQLite.
|
||||
//!
|
||||
//! Tests from test-plan §1 (CRUD) at the service layer.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use nomifun_api_types::{
|
||||
BatchImportMcpServersRequest, CreateMcpServerRequest, ImportMcpServerRequest, McpTransport, UpdateMcpServerRequest,
|
||||
};
|
||||
use nomifun_db::SqliteMcpServerRepository;
|
||||
use nomifun_mcp::{McpConfigService, McpError};
|
||||
|
||||
async fn make_service() -> McpConfigService {
|
||||
let db = nomifun_db::init_database_memory().await.unwrap();
|
||||
let repo = Arc::new(SqliteMcpServerRepository::new(db.pool().clone()));
|
||||
McpConfigService::new(repo)
|
||||
}
|
||||
|
||||
fn stdio_req(name: &str) -> CreateMcpServerRequest {
|
||||
CreateMcpServerRequest {
|
||||
name: name.to_owned(),
|
||||
description: Some("test".to_owned()),
|
||||
transport: McpTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec!["-y".into(), "@test/server".into()],
|
||||
env: HashMap::new(),
|
||||
},
|
||||
original_json: None,
|
||||
builtin: false,
|
||||
}
|
||||
}
|
||||
|
||||
fn http_req(name: &str) -> CreateMcpServerRequest {
|
||||
CreateMcpServerRequest {
|
||||
name: name.to_owned(),
|
||||
description: None,
|
||||
transport: McpTransport::Http {
|
||||
url: "https://example.com/mcp".into(),
|
||||
headers: HashMap::from([("Auth".into(), "Bearer tok".into())]),
|
||||
},
|
||||
original_json: None,
|
||||
builtin: false,
|
||||
}
|
||||
}
|
||||
|
||||
fn stdio_import_req(name: &str) -> ImportMcpServerRequest {
|
||||
ImportMcpServerRequest {
|
||||
name: name.to_owned(),
|
||||
description: Some("test".to_owned()),
|
||||
transport: McpTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec!["-y".into(), "@test/server".into()],
|
||||
env: HashMap::new(),
|
||||
},
|
||||
original_json: None,
|
||||
builtin: false,
|
||||
enabled: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn http_import_req(name: &str) -> ImportMcpServerRequest {
|
||||
ImportMcpServerRequest {
|
||||
name: name.to_owned(),
|
||||
description: None,
|
||||
transport: McpTransport::Http {
|
||||
url: "https://example.com/mcp".into(),
|
||||
headers: HashMap::from([("Auth".into(), "Bearer tok".into())]),
|
||||
},
|
||||
original_json: None,
|
||||
builtin: false,
|
||||
enabled: None,
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Create
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn create_and_get_stdio_server() {
|
||||
let svc = make_service().await;
|
||||
let resp = svc.add_server(stdio_req("test-stdio")).await.unwrap();
|
||||
|
||||
// Host-local INTEGER primary key, surfaced as a number on the DTO.
|
||||
assert!(resp.id > 0);
|
||||
assert_eq!(resp.name, "test-stdio");
|
||||
assert!(!resp.enabled);
|
||||
assert_eq!(resp.description.as_deref(), Some("test"));
|
||||
|
||||
let found = svc.get_server(&resp.id.to_string()).await.unwrap();
|
||||
assert_eq!(found.id, resp.id);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn create_http_with_headers() {
|
||||
let svc = make_service().await;
|
||||
let resp = svc.add_server(http_req("test-http")).await.unwrap();
|
||||
|
||||
match resp.transport {
|
||||
McpTransport::Http { ref url, ref headers } => {
|
||||
assert_eq!(url, "https://example.com/mcp");
|
||||
assert_eq!(headers.get("Auth").unwrap(), "Bearer tok");
|
||||
}
|
||||
_ => panic!("expected Http"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn create_same_name_upserts() {
|
||||
let svc = make_service().await;
|
||||
let first = svc.add_server(stdio_req("dup")).await.unwrap();
|
||||
let second = svc.add_server(http_req("dup")).await.unwrap();
|
||||
|
||||
assert_eq!(first.id, second.id);
|
||||
match second.transport {
|
||||
McpTransport::Http { .. } => {}
|
||||
_ => panic!("expected Http after upsert"),
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Read
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn list_empty() {
|
||||
let svc = make_service().await;
|
||||
assert!(svc.list_servers().await.unwrap().is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn list_returns_all() {
|
||||
let svc = make_service().await;
|
||||
svc.add_server(stdio_req("a")).await.unwrap();
|
||||
svc.add_server(http_req("b")).await.unwrap();
|
||||
assert_eq!(svc.list_servers().await.unwrap().len(), 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn get_not_found() {
|
||||
let svc = make_service().await;
|
||||
let err = svc.get_server("nonexistent").await.unwrap_err();
|
||||
assert!(matches!(err, McpError::NotFound(_)));
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Update
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn edit_name_is_rejected() {
|
||||
let svc = make_service().await;
|
||||
let created = svc.add_server(stdio_req("old")).await.unwrap();
|
||||
|
||||
let err = svc
|
||||
.edit_server(
|
||||
&created.id.to_string(),
|
||||
UpdateMcpServerRequest {
|
||||
name: Some("new".into()),
|
||||
description: None,
|
||||
transport: None,
|
||||
original_json: None,
|
||||
builtin: None,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(err, McpError::InvalidEdit(_)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn edit_transport() {
|
||||
let svc = make_service().await;
|
||||
let created = svc.add_server(stdio_req("test")).await.unwrap();
|
||||
|
||||
let updated = svc
|
||||
.edit_server(
|
||||
&created.id.to_string(),
|
||||
UpdateMcpServerRequest {
|
||||
name: None,
|
||||
description: None,
|
||||
transport: Some(McpTransport::Sse {
|
||||
url: "https://new.url".into(),
|
||||
headers: HashMap::new(),
|
||||
}),
|
||||
original_json: None,
|
||||
builtin: None,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
match updated.transport {
|
||||
McpTransport::Sse { ref url, .. } => assert_eq!(url, "https://new.url"),
|
||||
_ => panic!("expected Sse"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn edit_clears_description() {
|
||||
let svc = make_service().await;
|
||||
let created = svc.add_server(stdio_req("test")).await.unwrap();
|
||||
assert!(created.description.is_some());
|
||||
|
||||
let updated = svc
|
||||
.edit_server(
|
||||
&created.id.to_string(),
|
||||
UpdateMcpServerRequest {
|
||||
name: None,
|
||||
description: Some(None),
|
||||
transport: None,
|
||||
original_json: None,
|
||||
builtin: None,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(updated.description.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn edit_not_found() {
|
||||
let svc = make_service().await;
|
||||
let err = svc
|
||||
.edit_server(
|
||||
"nonexistent",
|
||||
UpdateMcpServerRequest {
|
||||
name: Some("x".into()),
|
||||
description: None,
|
||||
transport: None,
|
||||
original_json: None,
|
||||
builtin: None,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(err, McpError::NotFound(_)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn edit_name_conflict() {
|
||||
let svc = make_service().await;
|
||||
svc.add_server(stdio_req("a")).await.unwrap();
|
||||
let b = svc.add_server(stdio_req("b")).await.unwrap();
|
||||
|
||||
let err = svc
|
||||
.edit_server(
|
||||
&b.id.to_string(),
|
||||
UpdateMcpServerRequest {
|
||||
name: Some("a".into()),
|
||||
description: None,
|
||||
transport: None,
|
||||
original_json: None,
|
||||
builtin: None,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(err, McpError::InvalidEdit(_)));
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Delete
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn delete_removes_server() {
|
||||
let svc = make_service().await;
|
||||
let created = svc.add_server(stdio_req("del")).await.unwrap();
|
||||
let was_enabled = svc.delete_server(&created.id.to_string()).await.unwrap();
|
||||
assert!(!was_enabled);
|
||||
|
||||
let err = svc.get_server(&created.id.to_string()).await.unwrap_err();
|
||||
assert!(matches!(err, McpError::NotFound(_)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn delete_enabled_returns_true() {
|
||||
let svc = make_service().await;
|
||||
let created = svc.add_server(stdio_req("del-en")).await.unwrap();
|
||||
svc.toggle_server(&created.id.to_string()).await.unwrap();
|
||||
|
||||
let was_enabled = svc.delete_server(&created.id.to_string()).await.unwrap();
|
||||
assert!(was_enabled);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn delete_not_found() {
|
||||
let svc = make_service().await;
|
||||
let err = svc.delete_server("nonexistent").await.unwrap_err();
|
||||
assert!(matches!(err, McpError::NotFound(_)));
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Toggle
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn toggle_enables_then_disables() {
|
||||
let svc = make_service().await;
|
||||
let created = svc.add_server(stdio_req("tog")).await.unwrap();
|
||||
assert!(!created.enabled);
|
||||
|
||||
let toggled = svc.toggle_server(&created.id.to_string()).await.unwrap();
|
||||
assert!(toggled.enabled);
|
||||
|
||||
let toggled_back = svc.toggle_server(&created.id.to_string()).await.unwrap();
|
||||
assert!(!toggled_back.enabled);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Batch import
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn batch_import_creates_and_upserts() {
|
||||
let svc = make_service().await;
|
||||
svc.add_server(stdio_req("existing")).await.unwrap();
|
||||
|
||||
let req = BatchImportMcpServersRequest {
|
||||
servers: vec![
|
||||
http_import_req("existing"), // upsert
|
||||
stdio_import_req("new"), // create
|
||||
],
|
||||
};
|
||||
let results = svc.batch_import(req).await.unwrap();
|
||||
assert_eq!(results.len(), 2);
|
||||
|
||||
let all = svc.list_servers().await.unwrap();
|
||||
assert_eq!(all.len(), 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn batch_import_preserves_enabled_in_database() {
|
||||
let svc = make_service().await;
|
||||
let mut req = stdio_import_req("enabled-db-mcp");
|
||||
req.enabled = Some(true);
|
||||
|
||||
let result = svc
|
||||
.batch_import(BatchImportMcpServersRequest { servers: vec![req] })
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(result.len(), 1);
|
||||
assert_eq!(result[0].name, "enabled-db-mcp");
|
||||
assert!(result[0].enabled);
|
||||
|
||||
let listed = svc.list_servers().await.unwrap();
|
||||
assert_eq!(listed.len(), 1);
|
||||
assert!(listed[0].enabled);
|
||||
}
|
||||
@@ -0,0 +1,536 @@
|
||||
//! Integration tests for ACP session MCP injection.
|
||||
//!
|
||||
//! Covers test-plan items SI-1 through SI-7: capability parsing,
|
||||
//! format conversion, enabled-only filtering, and builtin server injection.
|
||||
|
||||
use std::collections::HashMap;
|
||||
|
||||
use nomifun_common::McpServerStatus;
|
||||
use nomifun_mcp::{
|
||||
AcpMcpCapabilities, AcpSessionMcpServer, ImageGenConfig, McpServer, McpServerTransport, NameValuePair,
|
||||
build_builtin_image_gen_server, build_session_mcp_servers, parse_acp_mcp_capabilities,
|
||||
};
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
fn make_server(name: &str, enabled: bool, transport: McpServerTransport) -> McpServer {
|
||||
McpServer {
|
||||
// Injection keys on `name`, never `id`; any stable value works here.
|
||||
id: name.bytes().map(i64::from).sum::<i64>().max(1),
|
||||
name: name.into(),
|
||||
description: None,
|
||||
enabled,
|
||||
transport,
|
||||
tools: vec![],
|
||||
last_test_status: McpServerStatus::Disconnected,
|
||||
last_connected: None,
|
||||
original_json: None,
|
||||
builtin: false,
|
||||
created_at: 0,
|
||||
updated_at: 0,
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// SI-1: Full capabilities → all transports retained
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn si_1_full_capabilities_retains_all_transports() {
|
||||
let caps = AcpMcpCapabilities {
|
||||
stdio: true,
|
||||
http: true,
|
||||
sse: true,
|
||||
};
|
||||
let servers = vec![
|
||||
make_server(
|
||||
"stdio-mcp",
|
||||
true,
|
||||
McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec!["-y".into(), "test-server".into()],
|
||||
env: HashMap::from([("KEY".into(), "VAL".into())]),
|
||||
},
|
||||
),
|
||||
make_server(
|
||||
"http-mcp",
|
||||
true,
|
||||
McpServerTransport::Http {
|
||||
url: "https://example.com/mcp".into(),
|
||||
headers: HashMap::from([("Authorization".into(), "Bearer tok".into())]),
|
||||
},
|
||||
),
|
||||
make_server(
|
||||
"sse-mcp",
|
||||
true,
|
||||
McpServerTransport::Sse {
|
||||
url: "https://example.com/sse".into(),
|
||||
headers: HashMap::new(),
|
||||
},
|
||||
),
|
||||
];
|
||||
|
||||
let result = build_session_mcp_servers(&servers, &caps);
|
||||
assert_eq!(result.len(), 3, "all 3 transports should be retained");
|
||||
|
||||
// Verify each type is present
|
||||
let has_stdio = result
|
||||
.iter()
|
||||
.any(|s| matches!(s, AcpSessionMcpServer::Stdio { name, .. } if name == "stdio-mcp"));
|
||||
let has_http = result
|
||||
.iter()
|
||||
.any(|s| matches!(s, AcpSessionMcpServer::Http { name, .. } if name == "http-mcp"));
|
||||
let has_sse = result
|
||||
.iter()
|
||||
.any(|s| matches!(s, AcpSessionMcpServer::Sse { name, .. } if name == "sse-mcp"));
|
||||
|
||||
assert!(has_stdio, "stdio server missing");
|
||||
assert!(has_http, "http server missing");
|
||||
assert!(has_sse, "sse server missing");
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// SI-2: stdio-only capabilities → only stdio retained
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn si_2_stdio_only_keeps_stdio_servers() {
|
||||
let caps = AcpMcpCapabilities {
|
||||
stdio: true,
|
||||
http: false,
|
||||
sse: false,
|
||||
};
|
||||
let servers = vec![
|
||||
make_server(
|
||||
"stdio-mcp",
|
||||
true,
|
||||
McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
},
|
||||
),
|
||||
make_server(
|
||||
"http-mcp",
|
||||
true,
|
||||
McpServerTransport::Http {
|
||||
url: "https://example.com/mcp".into(),
|
||||
headers: HashMap::new(),
|
||||
},
|
||||
),
|
||||
make_server(
|
||||
"sse-mcp",
|
||||
true,
|
||||
McpServerTransport::Sse {
|
||||
url: "https://example.com/sse".into(),
|
||||
headers: HashMap::new(),
|
||||
},
|
||||
),
|
||||
];
|
||||
|
||||
let result = build_session_mcp_servers(&servers, &caps);
|
||||
assert_eq!(result.len(), 1);
|
||||
assert!(matches!(&result[0], AcpSessionMcpServer::Stdio { name, .. } if name == "stdio-mcp"));
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// SI-3: No capabilities → empty list
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn si_3_no_capabilities_returns_empty() {
|
||||
let caps = AcpMcpCapabilities {
|
||||
stdio: false,
|
||||
http: false,
|
||||
sse: false,
|
||||
};
|
||||
let servers = vec![
|
||||
make_server(
|
||||
"s1",
|
||||
true,
|
||||
McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
},
|
||||
),
|
||||
make_server(
|
||||
"s2",
|
||||
true,
|
||||
McpServerTransport::Http {
|
||||
url: "https://example.com".into(),
|
||||
headers: HashMap::new(),
|
||||
},
|
||||
),
|
||||
];
|
||||
|
||||
let result = build_session_mcp_servers(&servers, &caps);
|
||||
assert!(result.is_empty(), "no capabilities → empty result");
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// SI-4: stdio server format conversion (env Record → Vec<{name,value}>)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn si_4_stdio_format_conversion() {
|
||||
let caps = AcpMcpCapabilities {
|
||||
stdio: true,
|
||||
http: false,
|
||||
sse: false,
|
||||
};
|
||||
let servers = vec![make_server(
|
||||
"test-stdio",
|
||||
true,
|
||||
McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec!["-y".into(), "test-server".into()],
|
||||
env: HashMap::from([
|
||||
("NODE_ENV".into(), "production".into()),
|
||||
("DEBUG".into(), "true".into()),
|
||||
]),
|
||||
},
|
||||
)];
|
||||
|
||||
let result = build_session_mcp_servers(&servers, &caps);
|
||||
assert_eq!(result.len(), 1);
|
||||
|
||||
match &result[0] {
|
||||
AcpSessionMcpServer::Stdio {
|
||||
name,
|
||||
command,
|
||||
args,
|
||||
env,
|
||||
} => {
|
||||
assert_eq!(name, "test-stdio");
|
||||
assert_eq!(command, "npx");
|
||||
assert_eq!(args, &["-y", "test-server"]);
|
||||
assert_eq!(env.len(), 2);
|
||||
// Sorted by name
|
||||
assert_eq!(env[0].name, "DEBUG");
|
||||
assert_eq!(env[0].value, "true");
|
||||
assert_eq!(env[1].name, "NODE_ENV");
|
||||
assert_eq!(env[1].value, "production");
|
||||
}
|
||||
_ => panic!("expected Stdio variant"),
|
||||
}
|
||||
|
||||
// Verify JSON wire format
|
||||
let json = serde_json::to_value(&result[0]).unwrap();
|
||||
assert_eq!(json["type"], "stdio");
|
||||
assert_eq!(json["name"], "test-stdio");
|
||||
assert_eq!(json["command"], "npx");
|
||||
assert!(json["env"].is_array());
|
||||
assert_eq!(json["env"][0]["name"], "DEBUG");
|
||||
assert_eq!(json["env"][0]["value"], "true");
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// SI-5: http server format conversion (headers Record → Vec<{name,value}>)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn si_5_http_format_conversion() {
|
||||
let caps = AcpMcpCapabilities {
|
||||
stdio: false,
|
||||
http: true,
|
||||
sse: false,
|
||||
};
|
||||
let servers = vec![make_server(
|
||||
"test-http",
|
||||
true,
|
||||
McpServerTransport::Http {
|
||||
url: "https://example.com/mcp".into(),
|
||||
headers: HashMap::from([
|
||||
("Authorization".into(), "Bearer tok".into()),
|
||||
("X-Custom".into(), "val".into()),
|
||||
]),
|
||||
},
|
||||
)];
|
||||
|
||||
let result = build_session_mcp_servers(&servers, &caps);
|
||||
assert_eq!(result.len(), 1);
|
||||
|
||||
match &result[0] {
|
||||
AcpSessionMcpServer::Http { name, url, headers } => {
|
||||
assert_eq!(name, "test-http");
|
||||
assert_eq!(url, "https://example.com/mcp");
|
||||
assert_eq!(headers.len(), 2);
|
||||
// Sorted by name
|
||||
assert_eq!(headers[0].name, "Authorization");
|
||||
assert_eq!(headers[0].value, "Bearer tok");
|
||||
assert_eq!(headers[1].name, "X-Custom");
|
||||
assert_eq!(headers[1].value, "val");
|
||||
}
|
||||
_ => panic!("expected Http variant"),
|
||||
}
|
||||
|
||||
// Verify JSON wire format
|
||||
let json = serde_json::to_value(&result[0]).unwrap();
|
||||
assert_eq!(json["type"], "http");
|
||||
assert_eq!(json["name"], "test-http");
|
||||
assert_eq!(json["url"], "https://example.com/mcp");
|
||||
assert!(json["headers"].is_array());
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// SI-6: only enabled servers appear
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn si_6_only_enabled_servers_in_result() {
|
||||
let caps = AcpMcpCapabilities {
|
||||
stdio: true,
|
||||
http: true,
|
||||
sse: true,
|
||||
};
|
||||
let servers = vec![
|
||||
make_server(
|
||||
"enabled-stdio",
|
||||
true,
|
||||
McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
},
|
||||
),
|
||||
make_server(
|
||||
"disabled-stdio",
|
||||
false,
|
||||
McpServerTransport::Stdio {
|
||||
command: "node".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
},
|
||||
),
|
||||
make_server(
|
||||
"enabled-http",
|
||||
true,
|
||||
McpServerTransport::Http {
|
||||
url: "https://example.com".into(),
|
||||
headers: HashMap::new(),
|
||||
},
|
||||
),
|
||||
make_server(
|
||||
"disabled-http",
|
||||
false,
|
||||
McpServerTransport::Http {
|
||||
url: "https://other.com".into(),
|
||||
headers: HashMap::new(),
|
||||
},
|
||||
),
|
||||
];
|
||||
|
||||
let result = build_session_mcp_servers(&servers, &caps);
|
||||
assert_eq!(result.len(), 2, "only 2 enabled servers should appear");
|
||||
|
||||
let names: Vec<&str> = result
|
||||
.iter()
|
||||
.map(|s| match s {
|
||||
AcpSessionMcpServer::Stdio { name, .. } => name.as_str(),
|
||||
AcpSessionMcpServer::Http { name, .. } => name.as_str(),
|
||||
AcpSessionMcpServer::Sse { name, .. } => name.as_str(),
|
||||
})
|
||||
.collect();
|
||||
assert!(names.contains(&"enabled-stdio"));
|
||||
assert!(names.contains(&"enabled-http"));
|
||||
assert!(!names.contains(&"disabled-stdio"));
|
||||
assert!(!names.contains(&"disabled-http"));
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// SI-7: builtin MCP injection (image generation with env vars)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn si_7_builtin_image_gen_injection() {
|
||||
let caps = AcpMcpCapabilities {
|
||||
stdio: true,
|
||||
http: true,
|
||||
sse: true,
|
||||
};
|
||||
|
||||
let img_config = ImageGenConfig {
|
||||
model: Some("dall-e-3".into()),
|
||||
api_url: Some("https://api.openai.com/v1".into()),
|
||||
api_key: Some("sk-test-key".into()),
|
||||
size: Some("1024x1024".into()),
|
||||
quality: Some("hd".into()),
|
||||
style: Some("natural".into()),
|
||||
};
|
||||
|
||||
// Build user servers
|
||||
let user_servers = vec![make_server(
|
||||
"user-mcp",
|
||||
true,
|
||||
McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec!["-y".into(), "user-server".into()],
|
||||
env: HashMap::new(),
|
||||
},
|
||||
)];
|
||||
|
||||
let mut session_servers = build_session_mcp_servers(&user_servers, &caps);
|
||||
|
||||
// Inject builtin image gen server
|
||||
if let Some(builtin) = build_builtin_image_gen_server(&caps, "/usr/local/bin/nomifun-img-gen", &img_config) {
|
||||
session_servers.push(builtin);
|
||||
}
|
||||
|
||||
assert_eq!(session_servers.len(), 2, "user server + builtin image gen");
|
||||
|
||||
// Verify the builtin server
|
||||
let builtin = &session_servers[1];
|
||||
match builtin {
|
||||
AcpSessionMcpServer::Stdio { name, command, env, .. } => {
|
||||
assert_eq!(name, "nomifun-image-generation");
|
||||
assert_eq!(command, "/usr/local/bin/nomifun-img-gen");
|
||||
|
||||
// Verify all 6 env vars are present
|
||||
assert_eq!(env.len(), 6);
|
||||
|
||||
let env_map: HashMap<&str, &str> = env.iter().map(|p| (p.name.as_str(), p.value.as_str())).collect();
|
||||
assert_eq!(env_map["NOMIFUN_IMG_MODEL"], "dall-e-3");
|
||||
assert_eq!(env_map["NOMIFUN_IMG_API_URL"], "https://api.openai.com/v1");
|
||||
assert_eq!(env_map["NOMIFUN_IMG_API_KEY"], "sk-test-key");
|
||||
assert_eq!(env_map["NOMIFUN_IMG_SIZE"], "1024x1024");
|
||||
assert_eq!(env_map["NOMIFUN_IMG_QUALITY"], "hd");
|
||||
assert_eq!(env_map["NOMIFUN_IMG_STYLE"], "natural");
|
||||
}
|
||||
_ => panic!("expected Stdio variant for builtin"),
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Capability parsing from various response formats
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn parse_capabilities_from_real_response_shape() {
|
||||
// Simulates a realistic ACP backend response with nested capabilities
|
||||
let response = serde_json::json!({
|
||||
"status": "ok",
|
||||
"version": "1.2.3",
|
||||
"mcp_capabilities": {
|
||||
"stdio": true,
|
||||
"http": true,
|
||||
"sse": false
|
||||
}
|
||||
});
|
||||
|
||||
let caps = parse_acp_mcp_capabilities(&response);
|
||||
assert!(caps.stdio);
|
||||
assert!(caps.http);
|
||||
assert!(!caps.sse);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_capabilities_empty_response() {
|
||||
let response = serde_json::json!({});
|
||||
let caps = parse_acp_mcp_capabilities(&response);
|
||||
// Default: stdio only
|
||||
assert!(caps.stdio);
|
||||
assert!(!caps.http);
|
||||
assert!(!caps.sse);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// End-to-end: parse capabilities + build servers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn end_to_end_parse_then_build() {
|
||||
let acp_response = serde_json::json!({
|
||||
"mcp_capabilities": { "stdio": true, "http": true, "sse": false }
|
||||
});
|
||||
let caps = parse_acp_mcp_capabilities(&acp_response);
|
||||
|
||||
let servers = vec![
|
||||
make_server(
|
||||
"stdio-srv",
|
||||
true,
|
||||
McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec![],
|
||||
env: HashMap::new(),
|
||||
},
|
||||
),
|
||||
make_server(
|
||||
"http-srv",
|
||||
true,
|
||||
McpServerTransport::Http {
|
||||
url: "https://example.com".into(),
|
||||
headers: HashMap::new(),
|
||||
},
|
||||
),
|
||||
make_server(
|
||||
"sse-srv",
|
||||
true,
|
||||
McpServerTransport::Sse {
|
||||
url: "https://example.com/sse".into(),
|
||||
headers: HashMap::new(),
|
||||
},
|
||||
),
|
||||
];
|
||||
|
||||
let result = build_session_mcp_servers(&servers, &caps);
|
||||
assert_eq!(result.len(), 2, "sse should be filtered out");
|
||||
|
||||
let names: Vec<&str> = result
|
||||
.iter()
|
||||
.map(|s| match s {
|
||||
AcpSessionMcpServer::Stdio { name, .. } => name.as_str(),
|
||||
AcpSessionMcpServer::Http { name, .. } => name.as_str(),
|
||||
AcpSessionMcpServer::Sse { name, .. } => name.as_str(),
|
||||
})
|
||||
.collect();
|
||||
assert!(names.contains(&"stdio-srv"));
|
||||
assert!(names.contains(&"http-srv"));
|
||||
assert!(!names.contains(&"sse-srv"));
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// JSON wire format verification
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn wire_format_is_acp_compatible() {
|
||||
let servers = vec![
|
||||
AcpSessionMcpServer::Stdio {
|
||||
name: "test-stdio".into(),
|
||||
command: "npx".into(),
|
||||
args: vec!["-y".into(), "server".into()],
|
||||
env: vec![NameValuePair {
|
||||
name: "K".into(),
|
||||
value: "V".into(),
|
||||
}],
|
||||
},
|
||||
AcpSessionMcpServer::Http {
|
||||
name: "test-http".into(),
|
||||
url: "https://example.com/mcp".into(),
|
||||
headers: vec![NameValuePair {
|
||||
name: "Auth".into(),
|
||||
value: "Bearer x".into(),
|
||||
}],
|
||||
},
|
||||
];
|
||||
|
||||
let json = serde_json::to_value(&servers).unwrap();
|
||||
let arr = json.as_array().unwrap();
|
||||
|
||||
// stdio variant
|
||||
assert_eq!(arr[0]["type"], "stdio");
|
||||
assert_eq!(arr[0]["name"], "test-stdio");
|
||||
assert_eq!(arr[0]["command"], "npx");
|
||||
assert_eq!(arr[0]["args"][0], "-y");
|
||||
assert_eq!(arr[0]["env"][0]["name"], "K");
|
||||
assert_eq!(arr[0]["env"][0]["value"], "V");
|
||||
|
||||
// http variant
|
||||
assert_eq!(arr[1]["type"], "http");
|
||||
assert_eq!(arr[1]["name"], "test-http");
|
||||
assert_eq!(arr[1]["url"], "https://example.com/mcp");
|
||||
assert_eq!(arr[1]["headers"][0]["name"], "Auth");
|
||||
}
|
||||
@@ -0,0 +1,111 @@
|
||||
//! Integration tests for read-only Agent MCP config discovery.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use nomifun_common::McpSource;
|
||||
use nomifun_db::SqliteMcpServerRepository;
|
||||
use nomifun_mcp::{DetectedServer, McpAgentAdapter, McpError, McpServerTransport, McpSyncService};
|
||||
|
||||
struct MockAdapter {
|
||||
source: McpSource,
|
||||
installed: bool,
|
||||
servers: std::sync::Mutex<Vec<DetectedServer>>,
|
||||
}
|
||||
|
||||
impl MockAdapter {
|
||||
fn new(source: McpSource, installed: bool) -> Self {
|
||||
Self {
|
||||
source,
|
||||
installed,
|
||||
servers: std::sync::Mutex::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
fn with_servers(source: McpSource, servers: Vec<DetectedServer>) -> Self {
|
||||
Self {
|
||||
source,
|
||||
installed: true,
|
||||
servers: std::sync::Mutex::new(servers),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl McpAgentAdapter for MockAdapter {
|
||||
fn source(&self) -> McpSource {
|
||||
self.source
|
||||
}
|
||||
|
||||
async fn is_installed(&self) -> Result<bool, McpError> {
|
||||
Ok(self.installed)
|
||||
}
|
||||
|
||||
async fn detect_existing(&self) -> Result<Vec<DetectedServer>, McpError> {
|
||||
if !self.installed {
|
||||
return Err(McpError::AgentNotInstalled(format!("{:?}", self.source)));
|
||||
}
|
||||
Ok(self.servers.lock().unwrap().clone())
|
||||
}
|
||||
|
||||
async fn install_server(&self, _name: &str, _transport: &McpServerTransport) -> Result<(), McpError> {
|
||||
unreachable!("write-to-CLI is no longer supported")
|
||||
}
|
||||
|
||||
async fn remove_server(&self, _name: &str) -> Result<(), McpError> {
|
||||
unreachable!("write-to-CLI is no longer supported")
|
||||
}
|
||||
}
|
||||
|
||||
async fn make_service(adapters: Vec<Arc<dyn McpAgentAdapter>>) -> McpSyncService {
|
||||
let db = nomifun_db::init_database_memory().await.unwrap();
|
||||
let repo: Arc<dyn nomifun_db::IMcpServerRepository> = Arc::new(SqliteMcpServerRepository::new(db.pool().clone()));
|
||||
McpSyncService::new(repo, adapters)
|
||||
}
|
||||
|
||||
fn stdio_transport() -> McpServerTransport {
|
||||
McpServerTransport::Stdio {
|
||||
command: "npx".into(),
|
||||
args: vec!["-y".into(), "@test/server".into()],
|
||||
env: HashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn get_agent_configs_returns_installed_agents() {
|
||||
let adapter_claude = Arc::new(MockAdapter::with_servers(
|
||||
McpSource::Claude,
|
||||
vec![DetectedServer {
|
||||
name: "existing-srv".into(),
|
||||
transport: stdio_transport(),
|
||||
importable: true,
|
||||
import_skip_reason: None,
|
||||
}],
|
||||
));
|
||||
let adapter_gemini = Arc::new(MockAdapter::new(McpSource::Gemini, false));
|
||||
let adapter_qwen = Arc::new(MockAdapter::new(McpSource::Qwen, true));
|
||||
|
||||
let sync_svc = make_service(vec![
|
||||
adapter_claude as Arc<dyn McpAgentAdapter>,
|
||||
adapter_gemini,
|
||||
adapter_qwen,
|
||||
])
|
||||
.await;
|
||||
let configs = sync_svc.get_agent_configs().await.unwrap();
|
||||
|
||||
assert_eq!(configs.len(), 2);
|
||||
assert_eq!(configs[0].source, McpSource::Claude);
|
||||
assert_eq!(configs[0].servers.len(), 1);
|
||||
assert_eq!(configs[0].servers[0].server.name, "existing-srv");
|
||||
assert_eq!(configs[1].source, McpSource::Qwen);
|
||||
assert!(configs[1].servers.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn get_agent_configs_empty_when_none_installed() {
|
||||
let adapter = Arc::new(MockAdapter::new(McpSource::Claude, false));
|
||||
let sync_svc = make_service(vec![adapter as Arc<dyn McpAgentAdapter>]).await;
|
||||
|
||||
let configs = sync_svc.get_agent_configs().await.unwrap();
|
||||
assert!(configs.is_empty());
|
||||
}
|
||||
@@ -0,0 +1,221 @@
|
||||
//! Integration tests for nomifun-mcp core types.
|
||||
//!
|
||||
//! Tests the public API surface: McpServer construction from DB rows,
|
||||
//! transport parsing/serialization, and response conversion.
|
||||
|
||||
use std::collections::HashMap;
|
||||
|
||||
use nomifun_common::McpServerStatus;
|
||||
use nomifun_db::models::McpServerRow;
|
||||
use nomifun_mcp::{McpServer, McpServerTransport, McpTool};
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// McpServer::from_row — full pipeline tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
fn row(transport_type: &str, transport_config: &str, tools: Option<&str>, status: &str) -> McpServerRow {
|
||||
McpServerRow {
|
||||
id: 42,
|
||||
name: "integration-test".into(),
|
||||
description: Some("Integration test server".into()),
|
||||
enabled: true,
|
||||
transport_type: transport_type.into(),
|
||||
transport_config: transport_config.into(),
|
||||
tools: tools.map(String::from),
|
||||
last_test_status: status.into(),
|
||||
last_connected: Some(9999),
|
||||
original_json: Some(r#"{"name":"integration-test"}"#.into()),
|
||||
builtin: false,
|
||||
deleted_at: None,
|
||||
created_at: 1000,
|
||||
updated_at: 2000,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stdio_server_full_pipeline() {
|
||||
let config = serde_json::json!({
|
||||
"command": "npx",
|
||||
"args": ["-y", "@modelcontextprotocol/server-everything"],
|
||||
"env": { "NODE_ENV": "test", "DEBUG": "mcp:*" }
|
||||
});
|
||||
let tools_json = serde_json::json!([
|
||||
{ "name": "echo", "description": "Echo input back" },
|
||||
{ "name": "add", "description": "Add numbers", "input_schema": { "type": "object" } }
|
||||
]);
|
||||
|
||||
let r = row("stdio", &config.to_string(), Some(&tools_json.to_string()), "connected");
|
||||
let server = McpServer::from_row(r).unwrap();
|
||||
|
||||
// Verify all fields
|
||||
assert_eq!(server.id, 42);
|
||||
assert_eq!(server.name, "integration-test");
|
||||
assert_eq!(server.description.as_deref(), Some("Integration test server"));
|
||||
assert!(server.enabled);
|
||||
assert_eq!(server.last_test_status, McpServerStatus::Connected);
|
||||
assert_eq!(server.last_connected, Some(9999));
|
||||
assert!(!server.builtin);
|
||||
|
||||
// Verify transport
|
||||
match &server.transport {
|
||||
McpServerTransport::Stdio { command, args, env } => {
|
||||
assert_eq!(command, "npx");
|
||||
assert_eq!(args.len(), 2);
|
||||
assert_eq!(env.len(), 2);
|
||||
assert_eq!(env["NODE_ENV"], "test");
|
||||
}
|
||||
_ => panic!("expected Stdio transport"),
|
||||
}
|
||||
|
||||
// Verify tools
|
||||
assert_eq!(server.tools.len(), 2);
|
||||
assert_eq!(server.tools[0].name, "echo");
|
||||
assert!(server.tools[1].input_schema.is_some());
|
||||
|
||||
// Convert to API response and verify
|
||||
let resp = server.into_response();
|
||||
assert_eq!(resp.id, 42);
|
||||
assert_eq!(resp.last_test_status, McpServerStatus::Connected);
|
||||
assert!(resp.tools.is_some());
|
||||
assert_eq!(resp.tools.unwrap().len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn http_server_with_headers() {
|
||||
let config = serde_json::json!({
|
||||
"url": "https://mcp.example.com/v1",
|
||||
"headers": { "Authorization": "Bearer secret123", "X-Custom": "value" }
|
||||
});
|
||||
|
||||
let r = row("http", &config.to_string(), None, "disconnected");
|
||||
let server = McpServer::from_row(r).unwrap();
|
||||
|
||||
match &server.transport {
|
||||
McpServerTransport::Http { url, headers } => {
|
||||
assert_eq!(url, "https://mcp.example.com/v1");
|
||||
assert_eq!(headers.len(), 2);
|
||||
assert_eq!(headers["Authorization"], "Bearer secret123");
|
||||
}
|
||||
_ => panic!("expected Http transport"),
|
||||
}
|
||||
|
||||
// Response should have no tools
|
||||
let resp = server.into_response();
|
||||
assert!(resp.tools.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sse_server_minimal() {
|
||||
let config = serde_json::json!({ "url": "https://sse.example.com/events" });
|
||||
let r = row("sse", &config.to_string(), None, "testing");
|
||||
let server = McpServer::from_row(r).unwrap();
|
||||
|
||||
assert_eq!(server.last_test_status, McpServerStatus::Testing);
|
||||
match &server.transport {
|
||||
McpServerTransport::Sse { url, headers } => {
|
||||
assert_eq!(url, "https://sse.example.com/events");
|
||||
assert!(headers.is_empty());
|
||||
}
|
||||
_ => panic!("expected Sse transport"),
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Transport DB roundtrip: domain -> JSON -> DB -> domain
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn transport_db_roundtrip_preserves_all_fields() {
|
||||
let transports = vec![
|
||||
McpServerTransport::Stdio {
|
||||
command: "python3".into(),
|
||||
args: vec!["-m".into(), "mcp_server".into()],
|
||||
env: HashMap::from([
|
||||
("PYTHONPATH".into(), "/usr/lib/python3".into()),
|
||||
("LOG_LEVEL".into(), "debug".into()),
|
||||
]),
|
||||
},
|
||||
McpServerTransport::Sse {
|
||||
url: "https://sse.example.com/mcp".into(),
|
||||
headers: HashMap::from([
|
||||
("Authorization".into(), "Bearer tok".into()),
|
||||
("Accept".into(), "text/event-stream".into()),
|
||||
]),
|
||||
},
|
||||
McpServerTransport::Http {
|
||||
url: "https://http.example.com/mcp".into(),
|
||||
headers: HashMap::new(),
|
||||
},
|
||||
];
|
||||
|
||||
for original in transports {
|
||||
let ttype = original.transport_type();
|
||||
let json = original.to_config_json().unwrap();
|
||||
let reconstructed = McpServerTransport::from_db(ttype, &json).unwrap();
|
||||
assert_eq!(reconstructed, original, "roundtrip failed for {ttype}");
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Error scenarios
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn invalid_transport_type_is_rejected() {
|
||||
let r = row("grpc", r#"{"endpoint":"localhost:50051"}"#, None, "disconnected");
|
||||
let err = McpServer::from_row(r).unwrap_err();
|
||||
let msg = err.to_string();
|
||||
assert!(msg.contains("unknown transport type"), "got: {msg}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn malformed_transport_json_is_rejected() {
|
||||
let r = row("stdio", "{broken json", None, "disconnected");
|
||||
let err = McpServer::from_row(r).unwrap_err();
|
||||
let msg = err.to_string();
|
||||
assert!(msg.contains("JSON"), "got: {msg}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn malformed_tools_json_is_rejected() {
|
||||
let r = row(
|
||||
"stdio",
|
||||
r#"{"command":"node"}"#,
|
||||
Some("[{not valid json}]"),
|
||||
"connected",
|
||||
);
|
||||
let err = McpServer::from_row(r).unwrap_err();
|
||||
let msg = err.to_string();
|
||||
assert!(msg.contains("JSON"), "got: {msg}");
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// McpTool construction
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[test]
|
||||
fn tool_fields_preserved_through_conversion() {
|
||||
let schema = serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"path": { "type": "string" }
|
||||
},
|
||||
"required": ["path"]
|
||||
});
|
||||
|
||||
let tool = McpTool {
|
||||
name: "read_file".into(),
|
||||
description: Some("Read a file from disk".into()),
|
||||
input_schema: Some(schema.clone()),
|
||||
};
|
||||
|
||||
// Domain -> API response
|
||||
let resp: nomifun_api_types::McpToolResponse = tool.clone().into();
|
||||
assert_eq!(resp.name, "read_file");
|
||||
assert_eq!(resp.description.as_deref(), Some("Read a file from disk"));
|
||||
assert_eq!(resp.input_schema, Some(schema.clone()));
|
||||
|
||||
// API response -> domain
|
||||
let back: McpTool = resp.into();
|
||||
assert_eq!(back, tool);
|
||||
}
|
||||
Reference in New Issue
Block a user