use axum::http::StatusCode;
use serde_json::{json, Value};
use std::collections::HashMap;
use tempfile::NamedTempFile;
use tokio::net::TcpListener;
use warmplane::{
config::{save_config, McpConfig, ProfileConfig, ServerConfig},
daemon::{
server::{build_router, initialize_state},
CapabilityMeta, PromptMeta, ResourceMeta,
},
};
#[tokio::test]
async fn test_profile_http_filtering_and_etag_caching() {
let temp_config = NamedTempFile::new().unwrap();
let config_path = temp_config.path().to_str().unwrap().to_string();
let mut initial_config = McpConfig::default();
initial_config.mcp_servers.insert(
"server_a".to_string(),
ServerConfig {
command: Some("echo".to_string()),
args: vec![],
env: HashMap::new(),
url: None,
protocol_version: None,
allow_stateless: None,
headers: HashMap::new(),
auth: None,
resilience: None,
},
);
initial_config.mcp_servers.insert(
"server_b".to_string(),
ServerConfig {
command: Some("echo".to_string()),
args: vec![],
env: HashMap::new(),
url: None,
protocol_version: None,
allow_stateless: None,
headers: HashMap::new(),
auth: None,
resilience: None,
},
);
initial_config.profiles.insert(
"only_a".to_string(),
ProfileConfig {
servers: vec!["server_a".to_string()],
description: Some("Profile with only server A".to_string()),
policy: None,
},
);
save_config(&config_path, &initial_config).unwrap();
let state = initialize_state(initial_config, &config_path)
.await
.unwrap();
{
let mut caps = state.capabilities.write().await;
caps.insert(
"a.tool".to_string(),
CapabilityMeta::new(
"server_a",
"tool_1",
"Tool from server A",
"Description A",
json!({"type": "object"}),
),
);
caps.insert(
"b.tool".to_string(),
CapabilityMeta::new(
"server_b",
"tool_2",
"Tool from server B",
"Description B",
json!({"type": "object"}),
),
);
}
{
let mut res = state.resources.write().await;
res.insert(
"res.a".to_string(),
ResourceMeta {
server: "server_a".to_string(),
uri: "file:///a.txt".to_string(),
name: "resource_a".to_string(),
description: Some("Resource A".to_string()),
mime_type: Some("text/plain".to_string()),
tags: vec![],
},
);
res.insert(
"res.b".to_string(),
ResourceMeta {
server: "server_b".to_string(),
uri: "file:///b.txt".to_string(),
name: "resource_b".to_string(),
description: Some("Resource B".to_string()),
mime_type: Some("text/plain".to_string()),
tags: vec![],
},
);
}
{
let mut prompts = state.prompts.write().await;
prompts.insert(
"prompt.a".to_string(),
PromptMeta {
server: "server_a".to_string(),
name: "prompt_a".to_string(),
title: Some("Prompt A".to_string()),
description: Some("Prompt on server A".to_string()),
arguments: vec![],
tags: vec![],
},
);
prompts.insert(
"prompt.b".to_string(),
PromptMeta {
server: "server_b".to_string(),
name: "prompt_b".to_string(),
title: Some("Prompt B".to_string()),
description: Some("Prompt on server B".to_string()),
arguments: vec![],
tags: vec![],
},
);
}
let app = build_router(state.clone());
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
let client = reqwest::Client::new();
let base_url = format!("http://127.0.0.1:{}", port);
let all_resp = client
.get(format!("{}/v1/capabilities", base_url))
.send()
.await
.unwrap();
assert_eq!(all_resp.status(), StatusCode::OK);
let all_etag = all_resp
.headers()
.get("etag")
.unwrap()
.to_str()
.unwrap()
.to_string();
let all_json: Value = all_resp.json().await.unwrap();
assert_eq!(all_json["capabilities"].as_array().unwrap().len(), 2);
let prof_resp = client
.get(format!("{}/v1/capabilities", base_url))
.header("x-warmplane-profile", "only_a")
.send()
.await
.unwrap();
assert_eq!(prof_resp.status(), StatusCode::OK);
let prof_etag = prof_resp
.headers()
.get("etag")
.unwrap()
.to_str()
.unwrap()
.to_string();
assert_ne!(all_etag, prof_etag);
assert!(prof_etag.contains("-p:only_a"));
let prof_json: Value = prof_resp.json().await.unwrap();
let caps_arr = prof_json["capabilities"].as_array().unwrap();
assert_eq!(caps_arr.len(), 1);
assert_eq!(caps_arr[0]["id"], "a.tool");
let not_mod_resp = client
.get(format!("{}/v1/capabilities", base_url))
.header("x-warmplane-profile", "only_a")
.header("if-none-match", &prof_etag)
.send()
.await
.unwrap();
assert_eq!(not_mod_resp.status(), StatusCode::NOT_MODIFIED);
let query_resp = client
.get(format!("{}/v1/capabilities?profile=only_a", base_url))
.send()
.await
.unwrap();
assert_eq!(query_resp.status(), StatusCode::OK);
let query_json: Value = query_resp.json().await.unwrap();
assert_eq!(query_json["capabilities"].as_array().unwrap().len(), 1);
let unknown_resp = client
.get(format!("{}/v1/capabilities", base_url))
.header("x-warmplane-profile", "non_existent")
.send()
.await
.unwrap();
assert_eq!(unknown_resp.status(), StatusCode::NOT_FOUND);
let unknown_json: Value = unknown_resp.json().await.unwrap();
assert_eq!(unknown_json["error"]["code"], "PROFILE_NOT_FOUND");
let desc_resp = client
.get(format!("{}/v1/capabilities/b.tool", base_url))
.header("x-warmplane-profile", "only_a")
.send()
.await
.unwrap();
assert_eq!(desc_resp.status(), StatusCode::NOT_FOUND);
let search_resp = client
.post(format!("{}/v1/capabilities/search", base_url))
.header("x-warmplane-profile", "only_a")
.json(&json!({"query": "Tool"}))
.send()
.await
.unwrap();
assert_eq!(search_resp.status(), StatusCode::OK);
let search_json: Value = search_resp.json().await.unwrap();
let hits = search_json["capabilities"].as_array().unwrap();
assert_eq!(hits.len(), 1);
assert_eq!(hits[0]["id"], "a.tool");
let call_resp = client
.post(format!("{}/v1/tools/call", base_url))
.header("x-warmplane-profile", "only_a")
.json(&json!({
"capability_id": "b.tool",
"args": {}
}))
.send()
.await
.unwrap();
assert_eq!(call_resp.status(), StatusCode::FORBIDDEN);
let call_json: Value = call_resp.json().await.unwrap();
assert_eq!(call_json["error"]["code"], "TOOL_NOT_IN_PROFILE");
let batch_resp = client
.post(format!("{}/v1/tools/batch_call", base_url))
.header("x-warmplane-profile", "only_a")
.json(&json!({
"steps": [
{
"id": "step1",
"capability_id": "b.tool",
"args": {},
"continue_on_error": false
}
]
}))
.send()
.await
.unwrap();
assert_eq!(batch_resp.status(), StatusCode::OK);
let batch_json: Value = batch_resp.json().await.unwrap();
assert_eq!(batch_json["ok"], false);
assert!(batch_json["results"][0]["error"]
.as_str()
.unwrap()
.contains("not in active profile"));
let res_resp = client
.get(format!("{}/v1/resources", base_url))
.header("x-warmplane-profile", "only_a")
.send()
.await
.unwrap();
assert_eq!(res_resp.status(), StatusCode::OK);
let res_json: Value = res_resp.json().await.unwrap();
let res_arr = res_json["resources"].as_array().unwrap();
assert_eq!(res_arr.len(), 1);
assert_eq!(res_arr[0]["id"], "res.a");
let read_resp = client
.post(format!("{}/v1/resources/read", base_url))
.header("x-warmplane-profile", "only_a")
.json(&json!({
"resource_id": "res.b"
}))
.send()
.await
.unwrap();
assert_eq!(read_resp.status(), StatusCode::NOT_FOUND);
let prompts_resp = client
.get(format!("{}/v1/prompts", base_url))
.header("x-warmplane-profile", "only_a")
.send()
.await
.unwrap();
assert_eq!(prompts_resp.status(), StatusCode::OK);
let prompts_json: Value = prompts_resp.json().await.unwrap();
let prompts_arr = prompts_json["prompts"].as_array().unwrap();
assert_eq!(prompts_arr.len(), 1);
assert_eq!(prompts_arr[0]["id"], "prompt.a");
let get_prompt_resp = client
.post(format!("{}/v1/prompts/get", base_url))
.header("x-warmplane-profile", "only_a")
.json(&json!({
"prompt_id": "prompt.b",
"arguments": {}
}))
.send()
.await
.unwrap();
assert_eq!(get_prompt_resp.status(), StatusCode::NOT_FOUND);
}
#[tokio::test]
async fn test_mcp_http_server_supported_protocol_versions_configuration() {
use rmcp::transport::streamable_http_server::{
session::local::LocalSessionManager, tower::StreamableHttpServerConfig,
StreamableHttpService,
};
use std::sync::Arc;
let temp_config = NamedTempFile::new().unwrap();
let config_path = temp_config.path().to_str().unwrap().to_string();
let mut mcp_config = McpConfig::default();
let http_cfg = warmplane::config::McpHttpServerConfig {
supported_protocol_versions: vec![
"2024-11-05".to_string(),
"2025-11-25".to_string(),
"2026-07-28".to_string(),
],
..Default::default()
};
mcp_config.mcp_http_server = Some(http_cfg.clone());
save_config(&config_path, &mcp_config).unwrap();
let state = initialize_state(mcp_config, &config_path).await.unwrap();
let parsed_protocol_versions: Vec<rmcp::model::ProtocolVersion> = http_cfg
.supported_protocol_versions
.iter()
.map(|s| match s.as_str() {
"2024-11-05" => rmcp::model::ProtocolVersion::V_2024_11_05,
"2025-03-26" => rmcp::model::ProtocolVersion::V_2025_03_26,
"2025-06-18" => rmcp::model::ProtocolVersion::V_2025_06_18,
"2025-11-25" => rmcp::model::ProtocolVersion::V_2025_11_25,
_ => rmcp::model::ProtocolVersion::V_2026_07_28,
})
.collect();
let state_for_factory = state.clone();
let mcp_service = StreamableHttpService::new(
move || {
let s = state_for_factory.clone();
let server = warmplane::mcp_server::FacadeMcpServer::new(s, None)
.with_supported_protocol_versions(parsed_protocol_versions.clone());
Ok(server)
},
Arc::new(LocalSessionManager::default()),
StreamableHttpServerConfig::default(),
);
let mcp_router = axum::Router::new().route_service("/mcp", mcp_service);
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
tokio::spawn(async move {
axum::serve(listener, mcp_router).await.unwrap();
});
let client = reqwest::Client::new();
let mcp_url = format!("http://127.0.0.1:{}/mcp", port);
let discover_resp = client
.post(&mcp_url)
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header("Mcp-Method", "server/discover")
.header("MCP-Protocol-Version", "2026-07-28")
.json(&json!({
"jsonrpc": "2.0",
"id": 1,
"method": "server/discover",
"params": {
"_meta": {
"io.modelcontextprotocol/protocolVersion": "2026-07-28",
"io.modelcontextprotocol/clientCapabilities": {}
}
}
}))
.send()
.await
.unwrap();
assert_eq!(discover_resp.status(), StatusCode::OK);
let body_text = discover_resp.text().await.unwrap();
let resp_json: Value = if body_text.starts_with("data:") {
let json_str = body_text
.lines()
.find_map(|l| l.strip_prefix("data: "))
.unwrap_or(&body_text);
serde_json::from_str(json_str).unwrap()
} else {
serde_json::from_str(&body_text).unwrap()
};
let supported = resp_json["result"]["supportedVersions"].as_array().unwrap();
let versions: Vec<&str> = supported.iter().filter_map(Value::as_str).collect();
assert_eq!(versions, vec!["2024-11-05", "2025-11-25", "2026-07-28"]);
}