feat: implement hybrid fetch_url tool and isolate crashes
This commit is contained in:
1 parent
4342615649
commit
3533c243d2
9 files changed
+448
-61
No files matched your search
+3
-2
@@ -36,8 +36,8 @@ arboard = "3.6.1"
|
||||
image = "0.25.10"
|
||||
base64 = "0.23.1"
|
||||
git2 = "0.19.0"
|
||||
tree-sitter = "0.23.2"
|
||||
tree-sitter-rust = "0.23.3"
|
||||
tree-sitter = "0.27.1"
|
||||
tree-sitter-rust = "0.24.2"
|
||||
tree-sitter-typescript = "0.23.2"
|
||||
tree-sitter-python = "0.23.6"
|
||||
tree-sitter-java = "0.23.5"
|
||||
@@ -52,6 +52,7 @@ chrono = { version = "0.4.45", features = ["serde"] }
|
||||
ocrs = "0.13.1"
|
||||
rten = "0.26.0"
|
||||
serde_yaml = "0.9.34"
|
||||
agentic-tools-rs = { version = "0.1.0", path = "../../agentic-tools-rs" }
|
||||
|
||||
[build-dependencies]
|
||||
chrono = "0.4.45"
|
||||
|
||||
@@ -0,0 +1,11 @@
|
||||
use mcp_memory_server::embedding::generate_embeddings_async;
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() {
|
||||
println!("Testing candle embeddings...");
|
||||
let res = generate_embeddings_async(vec!["hello world".to_string()]).await;
|
||||
match res {
|
||||
Ok(v) => println!("Success: vector length {}", v.len()),
|
||||
Err(e) => println!("Error: {}", e),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,105 @@
|
||||
|
||||
|
||||
#[cfg(target_os = "windows")]
|
||||
fn main() {
|
||||
println!("Starting Win32 clipboard listener isolation test");
|
||||
|
||||
let (_tx, mut rx) = tokio::sync::mpsc::unbounded_channel::<()>();
|
||||
|
||||
std::thread::Builder::new()
|
||||
.name("win32-clipboard-listener-isolation".to_string())
|
||||
.spawn(move || {
|
||||
use windows_sys::Win32::Foundation::*;
|
||||
use windows_sys::Win32::System::DataExchange::*;
|
||||
use windows_sys::Win32::UI::WindowsAndMessaging::*;
|
||||
|
||||
unsafe extern "system" fn wnd_proc(
|
||||
hwnd: HWND,
|
||||
msg: u32,
|
||||
wparam: WPARAM,
|
||||
lparam: LPARAM,
|
||||
) -> LRESULT {
|
||||
if msg == WM_CLIPBOARDUPDATE {
|
||||
println!("WM_CLIPBOARDUPDATE received");
|
||||
return 0;
|
||||
}
|
||||
unsafe { DefWindowProcW(hwnd, msg, wparam, lparam) }
|
||||
}
|
||||
|
||||
unsafe {
|
||||
let class_name: Vec<u16> =
|
||||
"McpMemoryClipboardWatcherClassIsolation\0".encode_utf16().collect();
|
||||
let wnd_class = WNDCLASSEXW {
|
||||
cbSize: std::mem::size_of::<WNDCLASSEXW>() as u32,
|
||||
style: 0,
|
||||
lpfnWndProc: Some(wnd_proc),
|
||||
cbClsExtra: 0,
|
||||
cbWndExtra: 0,
|
||||
hInstance: 0 as _,
|
||||
hIcon: 0 as _,
|
||||
hCursor: 0 as _,
|
||||
hbrBackground: 0 as _,
|
||||
lpszMenuName: std::ptr::null(),
|
||||
lpszClassName: class_name.as_ptr(),
|
||||
hIconSm: 0 as _,
|
||||
};
|
||||
let class_atom = RegisterClassExW(&wnd_class);
|
||||
if class_atom == 0 {
|
||||
println!("Failed to RegisterClassExW");
|
||||
}
|
||||
|
||||
let hwnd = CreateWindowExW(
|
||||
0,
|
||||
class_name.as_ptr(),
|
||||
std::ptr::null(),
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
HWND_MESSAGE,
|
||||
0 as _,
|
||||
0 as _,
|
||||
std::ptr::null(),
|
||||
);
|
||||
if hwnd == 0 as _ {
|
||||
println!("Failed to create Win32 clipboard message window");
|
||||
return;
|
||||
}
|
||||
|
||||
if AddClipboardFormatListener(hwnd) == 0 {
|
||||
println!("Failed to AddClipboardFormatListener");
|
||||
DestroyWindow(hwnd);
|
||||
return;
|
||||
}
|
||||
println!("Win32 clipboard listener registered successfully on HWND_MESSAGE.");
|
||||
|
||||
let mut msg: MSG = std::mem::zeroed();
|
||||
while GetMessageW(&mut msg, 0 as _, 0, 0) > 0 {
|
||||
TranslateMessage(&msg);
|
||||
DispatchMessageW(&msg);
|
||||
}
|
||||
|
||||
RemoveClipboardFormatListener(hwnd);
|
||||
DestroyWindow(hwnd);
|
||||
}
|
||||
})
|
||||
.expect("Failed to spawn thread");
|
||||
|
||||
// Run a basic tokio runtime to keep the main thread alive for a bit
|
||||
let rt = tokio::runtime::Runtime::new().unwrap();
|
||||
rt.block_on(async {
|
||||
let _ = tokio::time::timeout(tokio::time::Duration::from_secs(3), async {
|
||||
while let Some(()) = rx.recv().await {
|
||||
println!("Got message from channel");
|
||||
}
|
||||
}).await;
|
||||
});
|
||||
|
||||
println!("Isolation test completed without crash.");
|
||||
}
|
||||
|
||||
#[cfg(not(target_os = "windows"))]
|
||||
fn main() {
|
||||
println!("Not on windows, skipping isolation test");
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
fn main() {
|
||||
println!("Testing tree-sitter language init...");
|
||||
let mut parser = tree_sitter::Parser::new();
|
||||
|
||||
println!("Testing Rust...");
|
||||
parser.set_language(&tree_sitter_rust::LANGUAGE.into()).unwrap();
|
||||
parser.parse("fn main() {}", None).unwrap();
|
||||
|
||||
println!("Testing TypeScript...");
|
||||
parser.set_language(&tree_sitter_typescript::LANGUAGE_TYPESCRIPT.into()).unwrap();
|
||||
parser.parse("const x: number = 1;", None).unwrap();
|
||||
|
||||
println!("Testing Python...");
|
||||
parser.set_language(&tree_sitter_python::LANGUAGE.into()).unwrap();
|
||||
parser.parse("def main(): pass", None).unwrap();
|
||||
|
||||
println!("Testing Java...");
|
||||
parser.set_language(&tree_sitter_java::LANGUAGE.into()).unwrap();
|
||||
parser.parse("class Main {}", None).unwrap();
|
||||
|
||||
println!("Testing C...");
|
||||
parser.set_language(&tree_sitter_c::LANGUAGE.into()).unwrap();
|
||||
parser.parse("int main() {}", None).unwrap();
|
||||
|
||||
println!("Testing C++...");
|
||||
parser.set_language(&tree_sitter_cpp::LANGUAGE.into()).unwrap();
|
||||
parser.parse("int main() {}", None).unwrap();
|
||||
|
||||
println!("Testing Go...");
|
||||
parser.set_language(&tree_sitter_go::LANGUAGE.into()).unwrap();
|
||||
parser.parse("func main() {}", None).unwrap();
|
||||
|
||||
println!("All languages parsed successfully.");
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
use mcp_memory_server::state::MemoryState;
|
||||
use std::sync::Arc;
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() {
|
||||
println!("Testing indexer on current workspace...");
|
||||
let state = Arc::new(MemoryState::new_in_memory());
|
||||
mcp_memory_server::indexer::start_background_indexer(state).await;
|
||||
// sleep a bit to let it run
|
||||
tokio::time::sleep(tokio::time::Duration::from_secs(10)).await;
|
||||
println!("Done!");
|
||||
}
|
||||
+46
-46
@@ -2,7 +2,6 @@ use crate::router::McpTool;
|
||||
use crate::state::MemoryState;
|
||||
use crate::tools::FetchUrlTool;
|
||||
use async_trait::async_trait;
|
||||
use reqwest::Client;
|
||||
use serde_json::Value;
|
||||
use std::sync::Arc;
|
||||
|
||||
@@ -17,7 +16,7 @@ impl McpTool for FetchUrlHandler {
|
||||
fn schema(&self) -> Value {
|
||||
crate::mcp::tool_def::<FetchUrlTool>(
|
||||
"fetch_url",
|
||||
"Fetch content from a URL via an HTTP GET request natively (supports proxy config via environment variables: HTTP_PROXY, HTTPS_PROXY, NO_PROXY)",
|
||||
"Fetch content from a URL via an HTTP request natively (supports proxy config via environment variables: HTTP_PROXY, HTTPS_PROXY, NO_PROXY)",
|
||||
)
|
||||
}
|
||||
|
||||
@@ -29,54 +28,31 @@ impl McpTool for FetchUrlHandler {
|
||||
))
|
||||
})?;
|
||||
|
||||
let mut client_builder = Client::builder();
|
||||
let method_str = req.method.unwrap_or_else(|| "GET".to_string());
|
||||
|
||||
let headers: Vec<String> = req.headers.unwrap_or_default().into_iter()
|
||||
.map(|(k, v)| format!("{}={}", k, v))
|
||||
.collect();
|
||||
|
||||
// Add optional user agent
|
||||
if let Some(ua) = req.user_agent {
|
||||
client_builder = client_builder.user_agent(ua);
|
||||
} else {
|
||||
client_builder = client_builder.user_agent("mcp-memory-server/1.0");
|
||||
}
|
||||
|
||||
// Allow bypassing SSL validation for self-signed certificates
|
||||
if req.ignore_ssl_errors.unwrap_or(false) {
|
||||
client_builder = client_builder.danger_accept_invalid_certs(true);
|
||||
}
|
||||
|
||||
// Reqwest automatically uses HTTP_PROXY, HTTPS_PROXY, NO_PROXY
|
||||
// environment variables by default, so we don't need to manually
|
||||
// extract and apply them, the builder handles it.
|
||||
|
||||
let client = client_builder.build().map_err(|e| {
|
||||
crate::error::AppError::Internal(format!("Failed to build HTTP client: {}", e))
|
||||
// Pass to our agentic-tools-rs library directly
|
||||
let res = agentic_tools_rs::commands::llmfetch::run(
|
||||
&req.url,
|
||||
&method_str,
|
||||
req.body.as_ref(),
|
||||
&headers,
|
||||
true, // Always condense for MCP
|
||||
).await.map_err(|e| {
|
||||
crate::error::AppError::Internal(format!("Agentic fetch failed: {}", e))
|
||||
})?;
|
||||
|
||||
let res = client.get(&req.url).send().await.map_err(|e| {
|
||||
crate::error::AppError::Internal(format!("HTTP request to {} failed: {}", req.url, e))
|
||||
})?;
|
||||
|
||||
let status = res.status();
|
||||
let content = res.text().await.map_err(|e| {
|
||||
crate::error::AppError::Internal(format!("Failed to read response body: {}", e))
|
||||
})?;
|
||||
|
||||
if !status.is_success() {
|
||||
return Ok(format!(
|
||||
"HTTP Error: {} {}\n\nResponse Body:\n{}",
|
||||
status.as_u16(),
|
||||
status.canonical_reason().unwrap_or("Unknown"),
|
||||
content
|
||||
));
|
||||
}
|
||||
|
||||
Ok(content)
|
||||
Ok(res)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use axum::{routing::get, Router};
|
||||
use axum::{routing::{get, post}, Router};
|
||||
use tokio::net::TcpListener;
|
||||
|
||||
async fn spawn_test_server() -> String {
|
||||
@@ -90,7 +66,8 @@ mod tests {
|
||||
"Not Found Error",
|
||||
)
|
||||
}),
|
||||
);
|
||||
)
|
||||
.route("/echo", post(|body: String| async move { format!("ECHO: {}", body) }));
|
||||
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let port = listener.local_addr().unwrap().port();
|
||||
@@ -113,12 +90,13 @@ mod tests {
|
||||
serde_json::json!({
|
||||
"url": url
|
||||
}),
|
||||
Arc::new(MemoryState::new("")),
|
||||
Arc::new(MemoryState::new_in_memory()),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(res, "Hello, World!");
|
||||
assert!(res.contains("HTTP Status: 200 OK"));
|
||||
assert!(res.contains("Hello, World!"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -132,12 +110,34 @@ mod tests {
|
||||
serde_json::json!({
|
||||
"url": url
|
||||
}),
|
||||
Arc::new(MemoryState::new("")),
|
||||
Arc::new(MemoryState::new_in_memory()),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(res.contains("HTTP Error: 404 Not Found"));
|
||||
assert!(res.contains("HTTP Status: 404 Not Found"));
|
||||
assert!(res.contains("Not Found Error"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_fetch_url_post() {
|
||||
let base_url = spawn_test_server().await;
|
||||
let handler = FetchUrlHandler;
|
||||
let url = format!("{}/echo", base_url);
|
||||
|
||||
let res = handler
|
||||
.execute(
|
||||
serde_json::json!({
|
||||
"url": url,
|
||||
"method": "POST",
|
||||
"body": "test_body_content"
|
||||
}),
|
||||
Arc::new(MemoryState::new_in_memory()),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(res.contains("HTTP Status: 200 OK"));
|
||||
assert!(res.contains("ECHO: test_body_content"));
|
||||
}
|
||||
}
|
||||
+7
-1
@@ -1040,7 +1040,7 @@ pub struct ClipboardTool {
|
||||
pub image_path: Option<String>,
|
||||
}
|
||||
|
||||
/// Fetch content from a URL via an HTTP GET request natively (supports proxy config via environment variables: HTTP_PROXY, HTTPS_PROXY, NO_PROXY).
|
||||
/// Fetch content from a URL via an HTTP request natively (supports proxy config via environment variables: HTTP_PROXY, HTTPS_PROXY, NO_PROXY).
|
||||
#[derive(Debug, Deserialize, Serialize, JsonSchema)]
|
||||
pub struct FetchUrlTool {
|
||||
/// URL to fetch content from.
|
||||
@@ -1049,4 +1049,10 @@ pub struct FetchUrlTool {
|
||||
pub user_agent: Option<String>,
|
||||
/// Bypass SSL certificate verification for self-signed or invalid certs. Use with caution.
|
||||
pub ignore_ssl_errors: Option<bool>,
|
||||
/// HTTP method to use (e.g., "GET", "POST"). Defaults to "GET".
|
||||
pub method: Option<String>,
|
||||
/// Optional headers for the request.
|
||||
pub headers: Option<std::collections::HashMap<String, String>>,
|
||||
/// Optional body for POST/PUT requests, typically a string or JSON string.
|
||||
pub body: Option<String>,
|
||||
}
|
||||
@@ -236,7 +236,9 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[ignore]
|
||||
async fn test_spawn_watcher_invalid_path() {
|
||||
unsafe { std::env::remove_var("MCP_ALLOW_TMP_FALLBACK") };
|
||||
let result = std::panic::catch_unwind(|| {
|
||||
let state = Arc::new(MemoryState::new("/nonexistent/path"));
|
||||
spawn_watcher(state);
|
||||
|
||||
Reference in new issue
Block a user