Merge pull request 'MCP-setup' (#16) from MCP-setup into main
Reviewed-on: meulilivio/AISE1_Project#16
This commit is contained in:
commit
06c5c49092
@ -70,6 +70,7 @@ async def dispatch_tool(tool_name: str, arguments: dict) -> str:
|
|||||||
return f"DONE: {summary}"
|
return f"DONE: {summary}"
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
print(f"Trying to call tool '{tool_name}' with arguments: {arguments}")
|
||||||
print(f"Trying to call tool '{tool_name}' in dispatch_tool through MCPToolAdapter...")
|
print(f"Trying to call tool '{tool_name}' in dispatch_tool through MCPToolAdapter...")
|
||||||
result = await adapter.call_tool(tool_name, arguments)
|
result = await adapter.call_tool(tool_name, arguments)
|
||||||
|
|
||||||
|
|||||||
@ -11,7 +11,6 @@ class MCPToolAdapter:
|
|||||||
def __init__(self, config_path: str = "mcp_server_config.json"):
|
def __init__(self, config_path: str = "mcp_server_config.json"):
|
||||||
self.config_path = config_path
|
self.config_path = config_path
|
||||||
self.servers: Dict[str, Dict] = {}
|
self.servers: Dict[str, Dict] = {}
|
||||||
#self.exit_stack: Dict[str, Any] = {}
|
|
||||||
self.tool_registry: List[Dict[str, Any]] = []
|
self.tool_registry: List[Dict[str, Any]] = []
|
||||||
|
|
||||||
def _load_config(self) -> Dict[str, Any]:
|
def _load_config(self) -> Dict[str, Any]:
|
||||||
@ -57,10 +56,8 @@ class MCPToolAdapter:
|
|||||||
await session.initialize()
|
await session.initialize()
|
||||||
print(f"Session initialized for {server_name}. Requesting tools...")
|
print(f"Session initialized for {server_name}. Requesting tools...")
|
||||||
result = await session.list_tools()
|
result = await session.list_tools()
|
||||||
print(f"Tools received from {server_name}: {result}")
|
|
||||||
tools = result.tools
|
tools = result.tools
|
||||||
print(f"Tools received from {server_name}: {result}")
|
print(f"Tools received from {server_name}: {len(tools)} Tools")
|
||||||
#tools = getattr(result, 'tools', [])
|
|
||||||
|
|
||||||
for tool in tools:
|
for tool in tools:
|
||||||
t_params = tool.inputSchema.get("properties", {})
|
t_params = tool.inputSchema.get("properties", {})
|
||||||
@ -84,9 +81,6 @@ class MCPToolAdapter:
|
|||||||
})
|
})
|
||||||
|
|
||||||
print(f"Registered tool '{tool.name}' from {server_name}.")
|
print(f"Registered tool '{tool.name}' from {server_name}.")
|
||||||
|
|
||||||
print(f"Session for {server_name} ready. {len(tools)} tools found.")
|
|
||||||
|
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f"Failed to initialize {server_name}: {e}")
|
print(f"Failed to initialize {server_name}: {e}")
|
||||||
|
|||||||
@ -1,114 +0,0 @@
|
|||||||
import asyncio
|
|
||||||
import json
|
|
||||||
# import os
|
|
||||||
import numpy as np
|
|
||||||
from typing import List, Dict, Any
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
from sentence_transformers import SentenceTransformer # embedder
|
|
||||||
from mcp import ClientSession, StdioServerParameters
|
|
||||||
from mcp.client.stdio import stdio_client
|
|
||||||
|
|
||||||
class MCPToolRAGAdapter:
|
|
||||||
def __init__ (self, config_path: str = "mcp_server_config.json"):
|
|
||||||
self.config_path = config_path
|
|
||||||
self.tools = []
|
|
||||||
self.toolnames = []
|
|
||||||
self.embedder = SentenceTransformer('all-MiniLM-L6-v2') # for embedding tool descriptions
|
|
||||||
self.sessions = {}
|
|
||||||
self.exit_stack = {}
|
|
||||||
self.tool_registry = {}
|
|
||||||
self.tool_embeddings = None
|
|
||||||
|
|
||||||
def _load_config(self) -> Dict[str, Any]:
|
|
||||||
config_path = Path(__file__).parent / self.config_path
|
|
||||||
if not config_path.exists():
|
|
||||||
return {}
|
|
||||||
|
|
||||||
try:
|
|
||||||
with open(self.config_path, 'r') as f:
|
|
||||||
return json.load(f)
|
|
||||||
except json.JSONDecodeError as e:
|
|
||||||
print(f"Error decoding JSON config: {e}")
|
|
||||||
return {}
|
|
||||||
|
|
||||||
async def initialize_all_sessions(self):
|
|
||||||
"""Initialize all MCP sessions defined in the config file and index their tools."""
|
|
||||||
config = self._load_config()
|
|
||||||
for server_name, params in config.items():
|
|
||||||
print(f"initializing session for {server_name} with params: {params}")
|
|
||||||
server_params = StdioServerParameters(
|
|
||||||
commanf=params["command"],
|
|
||||||
args=params.get("args", []),
|
|
||||||
# env=params.get("env", {}),
|
|
||||||
)
|
|
||||||
|
|
||||||
# Verbindung aufbauen (Kontext-Manager manuell handhaben für Langzeit-Sessions)
|
|
||||||
transport_gen = stdio_client(server_params)
|
|
||||||
read, write = await transport_gen.__aenter__()
|
|
||||||
session = ClientSession(read, write)
|
|
||||||
await session.__aenter__()
|
|
||||||
await session.initialize()
|
|
||||||
|
|
||||||
self.sessions[server_name] = session
|
|
||||||
self.exit_stack[server_name] = (transport_gen, session) # Zum späteren sauberen Schließen speichern
|
|
||||||
print(f"Session for {server_name} initialized successfully.")
|
|
||||||
|
|
||||||
# call tools and index thme
|
|
||||||
result = await session.list_tools()
|
|
||||||
tools = result.get("tools", [])
|
|
||||||
|
|
||||||
for tool in tools:
|
|
||||||
self.tool_registry.append({
|
|
||||||
"server": server_name,
|
|
||||||
"tool_name": tool["name"],
|
|
||||||
"definition": tool,
|
|
||||||
"search_text": f"{tool['name']}: {tool.get('description', '')}",
|
|
||||||
})
|
|
||||||
self.tool_names.append(tool["name"])
|
|
||||||
|
|
||||||
# embeddings for all tools in this session
|
|
||||||
if self.tool_registry:
|
|
||||||
texts = [t["search_text"] for t in self.tool_registry]
|
|
||||||
self.tool_embeddings = self.embedder.encode(texts)
|
|
||||||
print(f"Indexing completed. {len(texts)} tools ready.")
|
|
||||||
|
|
||||||
def get_relevant_tools(self, query: str, top_k: int = 5) -> List[Dict[str, Any]]:
|
|
||||||
"""Given a user query, return the most relevant tools based on semantic similarity."""
|
|
||||||
if not self.tool_embeddings or not self.tool_registry:
|
|
||||||
print("No tools indexed yet.")
|
|
||||||
return []
|
|
||||||
|
|
||||||
query_embedding = self.embedder.encode([query])
|
|
||||||
similarities = np.dot(self.tool_embeddings, query_embedding.T).flatten()
|
|
||||||
top_indices = np.argsort(similarities)[-top_k:][::-1]
|
|
||||||
|
|
||||||
relevant_tools = [self.tool_registry[i] for i in top_indices]
|
|
||||||
|
|
||||||
return relevant_tools
|
|
||||||
|
|
||||||
async def call_tool(self, tool_name: str, arguments: Dict):
|
|
||||||
""" Finds the right server for the tool and calls it with the provided arguments. """
|
|
||||||
for item in self.tool_registry:
|
|
||||||
if item["definition"].name == tool_name:
|
|
||||||
server_name = item["server"]
|
|
||||||
session = self.sessions.get(server_name)
|
|
||||||
if session:
|
|
||||||
try:
|
|
||||||
result = await session.call_tool(tool_name, arguments)
|
|
||||||
return result
|
|
||||||
except Exception as e:
|
|
||||||
print(f"Error calling tool {tool_name} on server {server_name}: {e}")
|
|
||||||
return f"Error calling tool: {e}"
|
|
||||||
|
|
||||||
return f"Tool '{tool_name}' not found in registry."
|
|
||||||
|
|
||||||
async def shutdown_all_sessions(self):
|
|
||||||
"""Gracefully shutdown all MCP sessions."""
|
|
||||||
for server_name, (transport_gen, session) in self.exit_stack.items():
|
|
||||||
try:
|
|
||||||
await session.__aexit__(None, None, None)
|
|
||||||
await transport_gen.__aexit__(None, None, None)
|
|
||||||
print(f"Session for {server_name} shut down successfully.")
|
|
||||||
except Exception as e:
|
|
||||||
print(f"Error shutting down session for {server_name}: {e}")
|
|
||||||
@ -1,35 +1,15 @@
|
|||||||
import ast
|
import ast
|
||||||
from datetime import datetime
|
|
||||||
import subprocess
|
import subprocess
|
||||||
|
import sys
|
||||||
|
|
||||||
import io
|
import io
|
||||||
from pyflakes.api import check
|
from pyflakes.api import check # For linting Code
|
||||||
from pyflakes.reporter import Reporter
|
from pyflakes.reporter import Reporter # For linting Code
|
||||||
from mcp.server.fastmcp import FastMCP
|
from mcp.server.fastmcp import FastMCP
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
import venv
|
|
||||||
import shutil
|
|
||||||
|
|
||||||
# ── Sandbox venv ────────────────────────────────────────────────────────────
|
|
||||||
SERVER_BASE_DIR = Path(__file__).parent.resolve()
|
|
||||||
SANDBOX_DIR = SERVER_BASE_DIR / ".mcp_sandbox"
|
|
||||||
WORKSPACE_DIR = SERVER_BASE_DIR.parent.parent.parent.parent / "workspace"
|
|
||||||
|
|
||||||
def get_sandbox_paths():
|
|
||||||
"""Bestimmt die Executables innerhalb der Venv ohne os-Modul."""
|
|
||||||
if not SANDBOX_DIR.exists():
|
|
||||||
venv.create(SANDBOX_DIR, with_pip=True)
|
|
||||||
|
|
||||||
bin_folder = "Scripts" if Path("C:/").exists() else "bin" # Einfacher Check für Windows
|
|
||||||
|
|
||||||
python_exe = SANDBOX_DIR / bin_folder / "python"
|
|
||||||
pip_exe = SANDBOX_DIR / bin_folder / "pip"
|
|
||||||
|
|
||||||
return str(python_exe), str(pip_exe)
|
|
||||||
|
|
||||||
PYTHON_EXE, PIP_EXE = get_sandbox_paths()
|
|
||||||
|
|
||||||
# ── Configuration ────────────────────────────────────────────────────────────
|
# ── Configuration ────────────────────────────────────────────────────────────
|
||||||
EXEC_TIMEOUT = 10 # seconds before killing the subprocess
|
EXEC_TIMEOUT = 45 # seconds before killing the subprocess
|
||||||
MAX_OUTPUT_LENGTH = 3000 # max characters of stdout+stderr to return
|
MAX_OUTPUT_LENGTH = 3000 # max characters of stdout+stderr to return
|
||||||
|
|
||||||
# ── Create the MCP server ────────────────────────────────────────────────────
|
# ── Create the MCP server ────────────────────────────────────────────────────
|
||||||
@ -68,6 +48,8 @@ FORBIDDEN_SEQUENCES = ["../", "..\\", "/etc/", "/dev/",
|
|||||||
"C:\\Windows", "C:\\Program Files", "C:\\Users",
|
"C:\\Windows", "C:\\Program Files", "C:\\Users",
|
||||||
"compile(", "__import__", "os.", "sys.", "subprocess."]
|
"compile(", "__import__", "os.", "sys.", "subprocess."]
|
||||||
|
|
||||||
|
ALLOWED_PACKAGES = ["pygame", "numpy", "pandas"]
|
||||||
|
|
||||||
# ── Static Analysis ────────────────────────────────────────────────────
|
# ── Static Analysis ────────────────────────────────────────────────────
|
||||||
def check_code_safety(code: str) -> str | None:
|
def check_code_safety(code: str) -> str | None:
|
||||||
"""
|
"""
|
||||||
@ -229,86 +211,8 @@ def lint_code(code: str) -> str:
|
|||||||
|
|
||||||
return "\n".join(report)
|
return "\n".join(report)
|
||||||
|
|
||||||
|
|
||||||
@mcp.tool()
|
@mcp.tool()
|
||||||
def list_sandbox_packages() -> str:
|
def run_python_sandboxed(code: str) -> str:
|
||||||
"""
|
|
||||||
Lists all Python-Packages, that are installed in the Sandbox and their Version.
|
|
||||||
Helpful to determine if packages like 'pygame', 'numpy' or similair are already available
|
|
||||||
"""
|
|
||||||
try:
|
|
||||||
result = subprocess.run(
|
|
||||||
[PIP_EXE, "list"],
|
|
||||||
capture_output=True,
|
|
||||||
text=True,
|
|
||||||
timeout=10
|
|
||||||
)
|
|
||||||
|
|
||||||
if result.returncode != 0:
|
|
||||||
return f"Error while listing the packages: {result.stderr}"
|
|
||||||
|
|
||||||
if not result.stdout.strip():
|
|
||||||
return "The Sandbox environment is empty (only Standard-Libraries are available)."
|
|
||||||
|
|
||||||
return f"Installed Packages: {result.stdout}"
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
return f"Error trying to list packages from the Sandbox venv: {str(e)}"
|
|
||||||
|
|
||||||
|
|
||||||
@mcp.tool()
|
|
||||||
def install_package_into_sandbox(package_name: str) -> str:
|
|
||||||
"""
|
|
||||||
Install a Python package into the sandbox environment using pip.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
package_name: The name of the package to install (e.g., "requests").
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
A success message or an error message if installation fails.
|
|
||||||
"""
|
|
||||||
clean_name = "".join(e for e in package_name if e.isalnum() or e in "-_.")
|
|
||||||
|
|
||||||
if clean_name in BLOCKED_IMPORTS:
|
|
||||||
return f"Error: Installation of package '{clean_name}' is blocked due to security policies."
|
|
||||||
|
|
||||||
if clean_name in BLOCKED_BUILTINS:
|
|
||||||
return f"Error: Installation of package '{clean_name}' is blocked due to security policies."
|
|
||||||
|
|
||||||
if not clean_name:
|
|
||||||
return "Error: Invalid package name provided."
|
|
||||||
|
|
||||||
try:
|
|
||||||
result = subprocess.run(
|
|
||||||
[PIP_EXE, "install", clean_name],
|
|
||||||
capture_output=True,
|
|
||||||
text=True,
|
|
||||||
timeout=EXEC_TIMEOUT
|
|
||||||
)
|
|
||||||
|
|
||||||
if result.returncode == 0:
|
|
||||||
return f"Package '{clean_name}' installed successfully in the sandbox."
|
|
||||||
else:
|
|
||||||
return (f"Error installing package '{clean_name}':\n"
|
|
||||||
f"{result.stdout}\n{result.stderr}")
|
|
||||||
|
|
||||||
except subprocess.TimeoutExpired:
|
|
||||||
return f"Error: Package installation exceeded time limit of {EXEC_TIMEOUT} seconds and was terminated."
|
|
||||||
except Exception as e:
|
|
||||||
return f"Error during package installation: {e}"
|
|
||||||
|
|
||||||
|
|
||||||
@mcp.tool()
|
|
||||||
def reset_sandbox() -> str:
|
|
||||||
"""Löscht die gesamte Sandbox und erstellt sie neu (Full Reset)."""
|
|
||||||
if SANDBOX_DIR.exists():
|
|
||||||
shutil.rmtree(SANDBOX_DIR)
|
|
||||||
get_sandbox_paths()
|
|
||||||
return "Sandbox wurde komplett zurückgesetzt."
|
|
||||||
|
|
||||||
|
|
||||||
@mcp.tool()
|
|
||||||
def run_python_code_sandboxed(code: str) -> str:
|
|
||||||
"""
|
"""
|
||||||
Run Python code in a sandboxed environment.
|
Run Python code in a sandboxed environment.
|
||||||
|
|
||||||
@ -328,24 +232,11 @@ def run_python_code_sandboxed(code: str) -> str:
|
|||||||
static_safety = check_code_safety(code)
|
static_safety = check_code_safety(code)
|
||||||
if static_safety:
|
if static_safety:
|
||||||
return f"Code rejected:{static_safety}"
|
return f"Code rejected:{static_safety}"
|
||||||
|
|
||||||
run_id = datetime.now().strftime("%Y%m%d_%H%M%S")
|
|
||||||
jail_dir = WORKSPACE_DIR / f"sandbox_run_{run_id}"
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
jail_dir.mkdir(parents=True, exist_ok=True)
|
|
||||||
|
|
||||||
custom_env = {
|
|
||||||
"PYTHONPATH": str(WORKSPACE_DIR),
|
|
||||||
"PATH": str(Path(PYTHON_EXE).parent),
|
|
||||||
"HOME": str(jail_dir),
|
|
||||||
"TMPDIR": str(jail_dir)
|
|
||||||
}
|
|
||||||
|
|
||||||
result = subprocess.run(
|
result = subprocess.run(
|
||||||
[PYTHON_EXE, "-c", code],
|
[sys.executable, "-c", code],
|
||||||
cwd=str(WORKSPACE_DIR),
|
stdin=subprocess.DEVNULL,
|
||||||
env=custom_env,
|
|
||||||
capture_output=True,
|
capture_output=True,
|
||||||
text=True,
|
text=True,
|
||||||
timeout=EXEC_TIMEOUT)
|
timeout=EXEC_TIMEOUT)
|
||||||
@ -364,11 +255,6 @@ def run_python_code_sandboxed(code: str) -> str:
|
|||||||
return f"Error: Code execution exceeded time limit of {EXEC_TIMEOUT} seconds and was terminated."
|
return f"Error: Code execution exceeded time limit of {EXEC_TIMEOUT} seconds and was terminated."
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
return f"Error during code execution: {e}"
|
return f"Error during code execution: {e}"
|
||||||
|
|
||||||
finally:
|
|
||||||
if jail_dir.exists():
|
|
||||||
shutil.rmtree(jail_dir)
|
|
||||||
|
|
||||||
|
|
||||||
@mcp.tool()
|
@mcp.tool()
|
||||||
def python_code_validation(code: str) -> str:
|
def python_code_validation(code: str) -> str:
|
||||||
@ -393,6 +279,8 @@ def python_code_validation(code: str) -> str:
|
|||||||
return f"Valid Syntax, but with safety concerns: {static_analysis_result}; code execution is not allowed."
|
return f"Valid Syntax, but with safety concerns: {static_analysis_result}; code execution is not allowed."
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
return f"Error during code safety analysis: {e}"
|
return f"Error during code safety analysis: {e}"
|
||||||
|
|
||||||
|
return "Code is valid and can be executed in the sandbox"
|
||||||
|
|
||||||
|
|
||||||
# ── Run the server ───────────────────────────────────────────────────────────
|
# ── Run the server ───────────────────────────────────────────────────────────
|
||||||
|
|||||||
@ -240,10 +240,10 @@ def render_filetree_arborist(tree):
|
|||||||
selected = tree_view(
|
selected = tree_view(
|
||||||
data=data,
|
data=data,
|
||||||
icons={"open": "📂", "closed": "📁"},
|
icons={"open": "📂", "closed": "📁"},
|
||||||
height=200,
|
height=400,
|
||||||
selection=None,
|
selection=None,
|
||||||
select_internal_nodes=True, # allow clicking folder names, not just files
|
select_internal_nodes=True, # allow clicking folder names, not just files
|
||||||
open_by_default=False
|
open_by_default=True
|
||||||
)
|
)
|
||||||
|
|
||||||
return selected
|
return selected
|
||||||
|
|||||||
@ -24,3 +24,6 @@ python-dotenv>=1.0.0
|
|||||||
|
|
||||||
#For code editor functionality
|
#For code editor functionality
|
||||||
streamlit-ace>=0.1.0
|
streamlit-ace>=0.1.0
|
||||||
|
|
||||||
|
#Whitelisted Imports from Agent-Sandbox
|
||||||
|
pygame
|
||||||
@ -0,0 +1,278 @@
|
|||||||
|
import pytest
|
||||||
|
from pathlib import Path
|
||||||
|
from backend.agent.servers import mcp_server_file_search as server
|
||||||
|
|
||||||
|
|
||||||
|
# =========================================================
|
||||||
|
# FIXTURES
|
||||||
|
# =========================================================
|
||||||
|
|
||||||
|
@pytest.fixture()
|
||||||
|
def workspace(tmp_path, monkeypatch):
|
||||||
|
"""
|
||||||
|
Erstellt einen isolierten Workspace für jeden Test.
|
||||||
|
"""
|
||||||
|
ws = tmp_path / "workspace"
|
||||||
|
ws.mkdir()
|
||||||
|
|
||||||
|
monkeypatch.setattr(server, "ALLOWED_DIR", ws)
|
||||||
|
|
||||||
|
return ws
|
||||||
|
|
||||||
|
|
||||||
|
# =========================================================
|
||||||
|
# BASIC TESTS (1–10)
|
||||||
|
# =========================================================
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 1. _safe_path erlaubt gültige Pfade
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_safe_path_valid(workspace):
|
||||||
|
result = server._safe_path("test.txt")
|
||||||
|
|
||||||
|
assert result == workspace / "test.txt"
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 2. _safe_path blockiert Path Traversal
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_safe_path_blocks_traversal(workspace):
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
server._safe_path("../secret.txt")
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 3. list_files liefert leeren Hinweis
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_list_files_empty(workspace):
|
||||||
|
result = server.list_files()
|
||||||
|
|
||||||
|
assert result == "No files found in the project directory."
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 4. list_files findet Dateien rekursiv
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_list_files_recursive(workspace):
|
||||||
|
src = workspace / "src"
|
||||||
|
src.mkdir()
|
||||||
|
|
||||||
|
(src / "main.py").write_text("print('hello')")
|
||||||
|
|
||||||
|
result = server.list_files()
|
||||||
|
|
||||||
|
assert "src/main.py" in result
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 5. read_file liest Datei korrekt
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_read_file_success(workspace):
|
||||||
|
file = workspace / "hello.txt"
|
||||||
|
file.write_text("Hello World")
|
||||||
|
|
||||||
|
result = server.read_file("hello.txt")
|
||||||
|
|
||||||
|
assert result == "Hello World"
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 6. read_file erkennt fehlende Datei
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_read_file_missing(workspace):
|
||||||
|
result = server.read_file("missing.txt")
|
||||||
|
|
||||||
|
assert "does not exist" in result
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 7. write_new_file erstellt Datei
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_write_new_file_success(workspace):
|
||||||
|
result = server.write_new_file("new.txt", "content")
|
||||||
|
|
||||||
|
assert "OK:" in result
|
||||||
|
assert (workspace / "new.txt").exists()
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 8. write_new_file verhindert Überschreiben
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_write_new_file_existing(workspace):
|
||||||
|
file = workspace / "exists.txt"
|
||||||
|
file.write_text("old")
|
||||||
|
|
||||||
|
result = server.write_new_file("exists.txt", "new")
|
||||||
|
|
||||||
|
assert "already exists" in result
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 9. create_new_directory erstellt Verzeichnis
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_create_new_directory_success(workspace):
|
||||||
|
result = server.create_new_directory("mydir")
|
||||||
|
|
||||||
|
assert "OK:" in result
|
||||||
|
assert (workspace / "mydir").is_dir()
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 10. search_files findet Inhalte
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_search_files_content_match(workspace):
|
||||||
|
file = workspace / "notes.txt"
|
||||||
|
file.write_text("Python MCP Server")
|
||||||
|
|
||||||
|
result = server.search_files("mcp")
|
||||||
|
|
||||||
|
assert "[content]" in result
|
||||||
|
|
||||||
|
|
||||||
|
# =========================================================
|
||||||
|
# EDGE CASE TESTS (11–20)
|
||||||
|
# =========================================================
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 11. Mehrfaches Traversal blockieren
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_safe_path_double_traversal(workspace):
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
server._safe_path("../../../../etc/passwd")
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 12. Symlink Escape verhindern
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_safe_path_symlink_escape(workspace):
|
||||||
|
outside = workspace.parent / "outside"
|
||||||
|
outside.mkdir()
|
||||||
|
|
||||||
|
target = outside / "evil.txt"
|
||||||
|
target.write_text("bad")
|
||||||
|
|
||||||
|
link = workspace / "link"
|
||||||
|
link.symlink_to(outside)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
server._safe_path("link/evil.txt")
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 13. Dateien ohne Extension blockieren
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_write_file_without_extension(workspace):
|
||||||
|
result = server.write_new_file("README", "test")
|
||||||
|
|
||||||
|
assert "can only write" in result
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 14. Hidden Files blockieren
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_write_hidden_file(workspace):
|
||||||
|
result = server.write_new_file(".env", "SECRET=123")
|
||||||
|
|
||||||
|
assert "can only write" in result
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 15. Binary Files korrekt behandeln
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_read_binary_file(workspace):
|
||||||
|
binary = workspace / "data.bin"
|
||||||
|
binary.write_bytes(b"\xFF\xFE\xFD")
|
||||||
|
|
||||||
|
result = server.read_file("data.bin")
|
||||||
|
|
||||||
|
assert "not a text file" in result
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 16. Sehr große Zeilen durchsuchen
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_search_huge_line(workspace):
|
||||||
|
huge_text = "A" * 1_000_000 + "needle"
|
||||||
|
|
||||||
|
file = workspace / "huge.txt"
|
||||||
|
file.write_text(huge_text)
|
||||||
|
|
||||||
|
result = server.search_files("needle")
|
||||||
|
|
||||||
|
assert "[content]" in result
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 17. Leere Dateien lesen
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_read_empty_file(workspace):
|
||||||
|
file = workspace / "empty.txt"
|
||||||
|
file.write_text("")
|
||||||
|
|
||||||
|
result = server.read_file("empty.txt")
|
||||||
|
|
||||||
|
assert result == ""
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 18. Sonderzeichen im Query
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_search_special_characters(workspace):
|
||||||
|
file = workspace / "test.txt"
|
||||||
|
file.write_text("hello [world] (test)")
|
||||||
|
|
||||||
|
result = server.search_files("[world]")
|
||||||
|
|
||||||
|
assert "[content]" in result
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 19. Unicode-Dateinamen unterstützen
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_write_unicode_filename(workspace):
|
||||||
|
filename = "🔥_überraschung.txt"
|
||||||
|
|
||||||
|
result = server.write_new_file(filename, "unicode")
|
||||||
|
|
||||||
|
assert "OK:" in result
|
||||||
|
assert (workspace / filename).exists()
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 20. Tiefe Verzeichnisstrukturen
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_list_files_deep_nesting(workspace):
|
||||||
|
current = workspace
|
||||||
|
|
||||||
|
for i in range(50):
|
||||||
|
current = current / f"dir_{i}"
|
||||||
|
current.mkdir()
|
||||||
|
|
||||||
|
file = current / "deep.txt"
|
||||||
|
file.write_text("deep")
|
||||||
|
|
||||||
|
result = server.list_files()
|
||||||
|
|
||||||
|
assert "deep.txt" in result
|
||||||
@ -0,0 +1,297 @@
|
|||||||
|
import pytest
|
||||||
|
from unittest.mock import Mock, patch
|
||||||
|
from backend.agent.servers import mcp_server_web_search as server
|
||||||
|
|
||||||
|
|
||||||
|
# =========================================================
|
||||||
|
# BASIC TESTS (1–10)
|
||||||
|
# =========================================================
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 1. _validate_url erlaubt HTTPS
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_validate_url_https():
|
||||||
|
url = "https://example.com"
|
||||||
|
|
||||||
|
result = server._validate_url(url)
|
||||||
|
|
||||||
|
assert result == url
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 2. _validate_url erlaubt HTTP
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_validate_url_http():
|
||||||
|
url = "http://example.com"
|
||||||
|
|
||||||
|
result = server._validate_url(url)
|
||||||
|
|
||||||
|
assert result == url
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 3. _validate_url blockiert localhost
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_validate_url_localhost():
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
server._validate_url("http://localhost/admin")
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 4. _validate_url blockiert 127.0.0.1
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_validate_url_loopback():
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
server._validate_url("http://127.0.0.1")
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 5. _validate_url blockiert private IP
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_validate_url_private_ip():
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
server._validate_url("http://192.168.1.10")
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 6. web_search liefert Suchergebnisse
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
@patch("ddgs.DDGS")
|
||||||
|
def test_web_search_success(mock_ddgs):
|
||||||
|
mock_instance = Mock()
|
||||||
|
|
||||||
|
mock_instance.text.return_value = [
|
||||||
|
{
|
||||||
|
"title": "Example",
|
||||||
|
"href": "https://example.com",
|
||||||
|
"body": "Example snippet"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
|
||||||
|
mock_ddgs.return_value = mock_instance
|
||||||
|
|
||||||
|
result = server.web_search("example")
|
||||||
|
|
||||||
|
assert "Title: Example" in result
|
||||||
|
assert "https://example.com" in result
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 7. web_search ohne Ergebnisse
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
@patch("ddgs.DDGS")
|
||||||
|
def test_web_search_no_results(mock_ddgs):
|
||||||
|
mock_instance = Mock()
|
||||||
|
mock_instance.text.return_value = []
|
||||||
|
|
||||||
|
mock_ddgs.return_value = mock_instance
|
||||||
|
|
||||||
|
result = server.web_search("nothing")
|
||||||
|
|
||||||
|
assert "No results found" in result
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 8. fetch_page lädt HTML
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
@patch("requests.get")
|
||||||
|
def test_fetch_page_success(mock_get):
|
||||||
|
response = Mock()
|
||||||
|
|
||||||
|
response.status_code = 200
|
||||||
|
response.text = """
|
||||||
|
<html>
|
||||||
|
<body>
|
||||||
|
<h1>Hello World</h1>
|
||||||
|
</body>
|
||||||
|
</html>
|
||||||
|
"""
|
||||||
|
|
||||||
|
mock_get.return_value = response
|
||||||
|
|
||||||
|
result = server.fetch_page("https://example.com")
|
||||||
|
|
||||||
|
assert "Hello World" in result
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 9. fetch_page entfernt script Tags
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
@patch("requests.get")
|
||||||
|
def test_fetch_page_removes_script(mock_get):
|
||||||
|
response = Mock()
|
||||||
|
|
||||||
|
response.status_code = 200
|
||||||
|
response.text = """
|
||||||
|
<html>
|
||||||
|
<script>alert('xss')</script>
|
||||||
|
<body>Hello</body>
|
||||||
|
</html>
|
||||||
|
"""
|
||||||
|
|
||||||
|
mock_get.return_value = response
|
||||||
|
|
||||||
|
result = server.fetch_page("https://example.com")
|
||||||
|
|
||||||
|
assert "alert" not in result
|
||||||
|
assert "Hello" in result
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 10. fetch_page erkennt HTTP Fehler
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
@patch("requests.get")
|
||||||
|
def test_fetch_page_http_error(mock_get):
|
||||||
|
response = Mock()
|
||||||
|
|
||||||
|
response.status_code = 404
|
||||||
|
response.text = "Not Found"
|
||||||
|
|
||||||
|
mock_get.return_value = response
|
||||||
|
|
||||||
|
result = server.fetch_page("https://example.com")
|
||||||
|
|
||||||
|
assert "HTTP error 404" in result
|
||||||
|
|
||||||
|
|
||||||
|
# =========================================================
|
||||||
|
# EDGE CASE TESTS (11–20)
|
||||||
|
# =========================================================
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 11. Blockiere file:// SSRF
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_validate_url_blocks_file_scheme():
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
server._validate_url("file:///etc/passwd")
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 12. Blockiere ftp://
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_validate_url_blocks_ftp():
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
server._validate_url("ftp://example.com")
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 13. Blockiere AWS Metadata Endpoint
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_validate_url_blocks_metadata_ip():
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
server._validate_url("http://169.254.169.254")
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 14. Blockiere internes Docker Netzwerk
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_validate_url_blocks_docker_network():
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
server._validate_url("http://172.20.0.5")
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 15. Sehr lange URL
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
def test_validate_url_very_long():
|
||||||
|
long_url = "https://example.com/" + ("a" * 5000)
|
||||||
|
|
||||||
|
result = server._validate_url(long_url)
|
||||||
|
|
||||||
|
assert result == long_url
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 16. fetch_page behandelt Timeout
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
@patch("requests.get")
|
||||||
|
def test_fetch_page_timeout(mock_get):
|
||||||
|
import requests
|
||||||
|
|
||||||
|
mock_get.side_effect = requests.Timeout("timeout")
|
||||||
|
|
||||||
|
result = server.fetch_page("https://example.com")
|
||||||
|
|
||||||
|
assert "Error fetching page" in result
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 17. fetch_page behandelt Connection Error
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
@patch("requests.get")
|
||||||
|
def test_fetch_page_connection_error(mock_get):
|
||||||
|
import requests
|
||||||
|
|
||||||
|
mock_get.side_effect = requests.ConnectionError("connection failed")
|
||||||
|
|
||||||
|
result = server.fetch_page("https://example.com")
|
||||||
|
|
||||||
|
assert "Error fetching page" in result
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 18. fetch_page truncatet große Seiten
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
@patch("requests.get")
|
||||||
|
def test_fetch_page_truncates_large_content(mock_get):
|
||||||
|
response = Mock()
|
||||||
|
|
||||||
|
response.status_code = 200
|
||||||
|
response.text = "<html><body>" + ("A" * 10000) + "</body></html>"
|
||||||
|
|
||||||
|
mock_get.return_value = response
|
||||||
|
|
||||||
|
result = server.fetch_page("https://example.com")
|
||||||
|
|
||||||
|
assert "[... truncated ...]" in result
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 19. fetch_page bei leerem Body
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
@patch("requests.get")
|
||||||
|
def test_fetch_page_empty_content(mock_get):
|
||||||
|
response = Mock()
|
||||||
|
|
||||||
|
response.status_code = 200
|
||||||
|
response.text = "<html></html>"
|
||||||
|
|
||||||
|
mock_get.return_value = response
|
||||||
|
|
||||||
|
result = server.fetch_page("https://example.com")
|
||||||
|
|
||||||
|
assert "no text content found" in result.lower()
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
# 20. web_search behandelt Exception sauber
|
||||||
|
# ---------------------------------------------------------
|
||||||
|
|
||||||
|
@patch("ddgs.DDGS")
|
||||||
|
def test_web_search_exception(mock_ddgs):
|
||||||
|
mock_ddgs.side_effect = Exception("DDGS failed")
|
||||||
|
|
||||||
|
result = server.web_search("test")
|
||||||
|
|
||||||
|
assert "Search error" in result
|
||||||
Loading…
x
Reference in New Issue
Block a user