Merge branch 'main' of https://gitea.fhgr.ch/meulilivio/AISE1_Project into func_improvments
This commit is contained in:
commit
d860b73710
@ -97,7 +97,7 @@ async def dispatch_tool(tool_name: str, arguments: dict) -> str:
|
||||
return f"DONE: {summary}"
|
||||
|
||||
try:
|
||||
# REVIEW: debug print — remove before shipping.
|
||||
print(f"Trying to call tool '{tool_name}' with arguments: {arguments}")
|
||||
print(f"Trying to call tool '{tool_name}' in dispatch_tool through MCPToolAdapter...")
|
||||
result = await adapter.call_tool(tool_name, arguments)
|
||||
|
||||
|
||||
@ -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}")
|
||||
@ -15,52 +15,17 @@ blocks dangerous imports and builtins before spawning any subprocess.
|
||||
"""
|
||||
|
||||
import ast
|
||||
# REVIEW: dead code — datetime is imported but only used to generate run_id in
|
||||
# run_python_code_sandboxed(). That is legitimate, but note the import is unused in all
|
||||
# other tools; it would be cleaner as a local import inside run_python_code_sandboxed().
|
||||
from datetime import datetime
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
import io
|
||||
from pyflakes.api import check
|
||||
from pyflakes.reporter import Reporter
|
||||
from pyflakes.api import check # For linting Code
|
||||
from pyflakes.reporter import Reporter # For linting Code
|
||||
from mcp.server.fastmcp import FastMCP
|
||||
from pathlib import Path
|
||||
import venv
|
||||
import shutil
|
||||
|
||||
# ── Sandbox venv ────────────────────────────────────────────────────────────
|
||||
SERVER_BASE_DIR = Path(__file__).parent.resolve()
|
||||
SANDBOX_DIR = SERVER_BASE_DIR / ".mcp_sandbox"
|
||||
# Navigate four levels up from servers/ to reach the project root, then into workspace/.
|
||||
WORKSPACE_DIR = SERVER_BASE_DIR.parent.parent.parent.parent / "workspace"
|
||||
|
||||
|
||||
def get_sandbox_paths():
|
||||
"""Locate (and lazily create) the sandbox venv and return its executable paths.
|
||||
|
||||
Creates the venv on first call if it does not yet exist. Uses a simple
|
||||
filesystem check to determine whether we are on Windows (``Scripts/``) or
|
||||
POSIX (``bin/``), avoiding the ``os`` module which is blocked inside the
|
||||
sandbox itself.
|
||||
|
||||
Returns:
|
||||
Tuple of (python_exe_path, pip_exe_path) as strings.
|
||||
"""
|
||||
if not SANDBOX_DIR.exists():
|
||||
venv.create(SANDBOX_DIR, with_pip=True)
|
||||
|
||||
# Path("C:/").exists() is True on Windows, False on Linux/macOS.
|
||||
bin_folder = "Scripts" if Path("C:/").exists() else "bin"
|
||||
|
||||
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 ────────────────────────────────────────────────────────────
|
||||
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
|
||||
|
||||
# ── Create the MCP server ────────────────────────────────────────────────────
|
||||
@ -99,6 +64,8 @@ FORBIDDEN_SEQUENCES = ["../", "..\\", "/etc/", "/dev/",
|
||||
"C:\\Windows", "C:\\Program Files", "C:\\Users",
|
||||
"compile(", "__import__", "os.", "sys.", "subprocess."]
|
||||
|
||||
ALLOWED_PACKAGES = ["pygame", "numpy", "pandas"]
|
||||
|
||||
# ── Static Analysis ────────────────────────────────────────────────────
|
||||
def check_code_safety(code: str) -> str | None:
|
||||
"""
|
||||
@ -260,94 +227,8 @@ def lint_code(code: str) -> str:
|
||||
|
||||
return "\n".join(report)
|
||||
|
||||
|
||||
@mcp.tool()
|
||||
def list_sandbox_packages() -> 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:
|
||||
"""Delete and recreate the sandbox virtual environment (full reset).
|
||||
|
||||
Useful when a package installation went wrong or the venv became corrupted.
|
||||
Calling get_sandbox_paths() after deletion triggers the lazy creation logic.
|
||||
|
||||
Returns:
|
||||
Confirmation string after the reset completes.
|
||||
"""
|
||||
if SANDBOX_DIR.exists():
|
||||
shutil.rmtree(SANDBOX_DIR)
|
||||
# Re-calling get_sandbox_paths() triggers venv creation for the fresh sandbox.
|
||||
get_sandbox_paths()
|
||||
return "Sandbox wurde komplett zurückgesetzt."
|
||||
|
||||
|
||||
@mcp.tool()
|
||||
def run_python_code_sandboxed(code: str) -> str:
|
||||
def run_python_sandboxed(code: str) -> str:
|
||||
"""
|
||||
Run Python code in a sandboxed environment.
|
||||
|
||||
@ -369,26 +250,10 @@ def run_python_code_sandboxed(code: str) -> str:
|
||||
if static_safety:
|
||||
return f"Code rejected:{static_safety}"
|
||||
|
||||
# Each run gets its own temporary directory so concurrent runs don't interfere.
|
||||
run_id = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
jail_dir = WORKSPACE_DIR / f"sandbox_run_{run_id}"
|
||||
|
||||
try:
|
||||
jail_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Minimal environment: only the sandbox Python is on PATH, HOME and TMPDIR
|
||||
# point to the per-run jail directory so the subprocess cannot access user files.
|
||||
custom_env = {
|
||||
"PYTHONPATH": str(WORKSPACE_DIR),
|
||||
"PATH": str(Path(PYTHON_EXE).parent),
|
||||
"HOME": str(jail_dir),
|
||||
"TMPDIR": str(jail_dir)
|
||||
}
|
||||
|
||||
result = subprocess.run(
|
||||
[PYTHON_EXE, "-c", code],
|
||||
cwd=str(WORKSPACE_DIR),
|
||||
env=custom_env,
|
||||
[sys.executable, "-c", code],
|
||||
stdin=subprocess.DEVNULL,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=EXEC_TIMEOUT)
|
||||
@ -409,12 +274,6 @@ def run_python_code_sandboxed(code: str) -> str:
|
||||
except Exception as e:
|
||||
return f"Error during code execution: {e}"
|
||||
|
||||
finally:
|
||||
# Always clean up the per-run jail directory, even if execution failed.
|
||||
if jail_dir.exists():
|
||||
shutil.rmtree(jail_dir)
|
||||
|
||||
|
||||
@mcp.tool()
|
||||
def python_code_validation(code: str) -> str:
|
||||
"""Validate Python code for syntax correctness and sandbox safety without executing it.
|
||||
@ -443,9 +302,8 @@ def python_code_validation(code: str) -> str:
|
||||
return f"Valid Syntax, but with safety concerns: {static_analysis_result}; code execution is not allowed."
|
||||
except Exception as e:
|
||||
return f"Error during code safety analysis: {e}"
|
||||
# REVIEW: unreachable code — python_code_validation() falls off the end of the function
|
||||
# without an explicit `return` when static_analysis_result is None (safe code); the function
|
||||
# implicitly returns None instead of returning a success message to the caller.
|
||||
|
||||
return "Code is valid and can be executed in the sandbox"
|
||||
|
||||
|
||||
# ── Run the server ───────────────────────────────────────────────────────────
|
||||
|
||||
@ -0,0 +1,122 @@
|
||||
"""Handles internet search requests and page fetching for use as AI chat context."""
|
||||
|
||||
import ipaddress
|
||||
import socket
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import requests
|
||||
from bs4 import BeautifulSoup
|
||||
from ddgs import DDGS
|
||||
|
||||
# Maximum characters extracted from a fetched page before truncating.
|
||||
MAX_PAGE_CHARS = 3000
|
||||
|
||||
_HEADERS = {"User-Agent": "Mozilla/5.0 (compatible; AICodeEditor/1.0)"}
|
||||
|
||||
|
||||
class SearchManager:
|
||||
"""Performs DuckDuckGo searches and fetches web pages for AI context injection.
|
||||
|
||||
All outbound requests are validated against an SSRF blocklist so that
|
||||
localhost and private network addresses can never be reached.
|
||||
"""
|
||||
|
||||
def perform_search(self, query: str, max_results: int = 5) -> list[dict]:
|
||||
"""Execute a DuckDuckGo text search and return normalised results.
|
||||
|
||||
Args:
|
||||
query: The search query string.
|
||||
max_results: Maximum number of results to return.
|
||||
|
||||
Returns:
|
||||
List of {"title": str, "url": str, "snippet": str} dicts,
|
||||
or an empty list if the search fails.
|
||||
"""
|
||||
try:
|
||||
with DDGS() as ddgs:
|
||||
raw = list(ddgs.text(query, max_results=max_results))
|
||||
return self.parse_results(raw)
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
def parse_results(self, raw_results: list[dict]) -> list[dict]:
|
||||
"""Normalise raw DDGS result dicts to a consistent {"title", "url", "snippet"} shape.
|
||||
|
||||
Args:
|
||||
raw_results: List of raw dicts returned by ddgs.text().
|
||||
|
||||
Returns:
|
||||
Normalised list of result dicts.
|
||||
"""
|
||||
results = []
|
||||
for r in raw_results:
|
||||
results.append({
|
||||
"title": r.get("title", ""),
|
||||
"url": r.get("href", r.get("url", "")),
|
||||
"snippet": r.get("body", r.get("snippet", "")),
|
||||
})
|
||||
return results
|
||||
|
||||
def fetch_page(self, url: str) -> str:
|
||||
"""Fetch a web page and return its plain text content, truncated to MAX_PAGE_CHARS.
|
||||
|
||||
Args:
|
||||
url: The URL to fetch.
|
||||
|
||||
Returns:
|
||||
Plain text extracted from the page, or an error message string.
|
||||
|
||||
Raises:
|
||||
ValueError: if the URL fails the SSRF safety check.
|
||||
"""
|
||||
self._validate_url(url)
|
||||
try:
|
||||
response = requests.get(url, timeout=10, headers=_HEADERS)
|
||||
response.raise_for_status()
|
||||
|
||||
soup = BeautifulSoup(response.text, "html.parser")
|
||||
|
||||
# Remove non-content elements before extracting text.
|
||||
for tag in soup(["script", "style", "nav", "footer"]):
|
||||
tag.decompose()
|
||||
|
||||
text = soup.get_text(separator="\n", strip=True)
|
||||
|
||||
if len(text) > MAX_PAGE_CHARS:
|
||||
text = text[:MAX_PAGE_CHARS] + "\n... [truncated]"
|
||||
|
||||
return text
|
||||
|
||||
except ValueError:
|
||||
raise
|
||||
except Exception as e:
|
||||
return f"Error fetching page: {e}"
|
||||
|
||||
def _validate_url(self, url: str) -> None:
|
||||
"""Block localhost, private IPs, and non-http(s) schemes to prevent SSRF attacks.
|
||||
|
||||
Args:
|
||||
url: The URL to validate.
|
||||
|
||||
Raises:
|
||||
ValueError: if the URL is considered unsafe.
|
||||
"""
|
||||
parsed = urlparse(url)
|
||||
|
||||
if parsed.scheme not in ("http", "https"):
|
||||
raise ValueError(f"Blocked: only http/https allowed, got '{parsed.scheme}'")
|
||||
|
||||
hostname = parsed.hostname or ""
|
||||
|
||||
if hostname.lower() in ("localhost", "127.0.0.1", "::1"):
|
||||
raise ValueError("Blocked: localhost access denied")
|
||||
|
||||
try:
|
||||
ip = ipaddress.ip_address(socket.gethostbyname(hostname))
|
||||
if ip.is_private or ip.is_loopback or ip.is_link_local:
|
||||
raise ValueError(f"Blocked: private/loopback IP denied ({ip})")
|
||||
except (socket.gaierror, ValueError) as e:
|
||||
# Re-raise our own ValueError; ignore DNS resolution failures
|
||||
# (let requests handle unknown hostnames naturally).
|
||||
if isinstance(e, ValueError):
|
||||
raise
|
||||
@ -12,16 +12,17 @@ class SystemPrompter:
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
# REVIEW: unused parameter (in production) — `file_context` is never passed by the only
|
||||
# production call site (frontend/chat.py line 242 calls generate_prompt() with no args),
|
||||
# so the file-embedding branch (lines 31-46) is dead in production. It is tested in
|
||||
# tests/test_system_prompter.py but the feature is not wired up in the UI.
|
||||
def generate_prompt(file_context: dict | None = None) -> str:
|
||||
"""Build a system prompt, optionally embedding a file's content.
|
||||
def generate_prompt(
|
||||
file_context: dict | None = None,
|
||||
search_context: list[dict] | None = None,
|
||||
) -> str:
|
||||
"""Build a system prompt, optionally embedding a file and/or web search results.
|
||||
|
||||
Args:
|
||||
file_context: dict with keys 'name' (filename) and 'content' (raw text),
|
||||
or None if no file should be included.
|
||||
search_context: list of {"title", "url", "snippet"} dicts from SearchManager,
|
||||
or None if no search results should be included.
|
||||
|
||||
Returns:
|
||||
A ready-to-use system prompt string.
|
||||
@ -32,6 +33,8 @@ class SystemPrompter:
|
||||
"Be concise and precise. Use markdown and fenced code blocks where appropriate."
|
||||
)
|
||||
|
||||
prompt = base
|
||||
|
||||
if file_context:
|
||||
name = file_context.get("name", "unknown")
|
||||
content = file_context.get("content", "")
|
||||
@ -40,13 +43,23 @@ class SystemPrompter:
|
||||
if len(content) > MAX_FILE_CHARS:
|
||||
content = content[:MAX_FILE_CHARS] + "\n... [truncated]"
|
||||
|
||||
file_section = (
|
||||
prompt += (
|
||||
f"\n\nThe user currently has the following file open in the editor:\n"
|
||||
f"<file name=\"{name}\">\n"
|
||||
f"<code>\n{content}\n</code>\n"
|
||||
f"</file>\n"
|
||||
f"Refer to this file when answering questions about the code."
|
||||
)
|
||||
return base + file_section
|
||||
|
||||
return base
|
||||
if search_context:
|
||||
search_section = "\n\nThe user has performed a web search. Use the results below as additional context if relevant:\n<search_results>\n"
|
||||
for i, r in enumerate(search_context, 1):
|
||||
search_section += (
|
||||
f"[{i}] {r.get('title', '')}\n"
|
||||
f"URL: {r.get('url', '')}\n"
|
||||
f"{r.get('snippet', '')}\n\n"
|
||||
)
|
||||
search_section += "</search_results>"
|
||||
prompt += search_section
|
||||
|
||||
return prompt
|
||||
|
||||
166
frontend/chat.py
166
frontend/chat.py
@ -6,6 +6,9 @@ from pathlib import Path
|
||||
import streamlit as st
|
||||
from backend.managers.chat_manager import ChatManager
|
||||
from backend.managers.system_prompter import SystemPrompter
|
||||
from backend.managers.search_manager import SearchManager
|
||||
|
||||
import asyncio
|
||||
|
||||
|
||||
# ── Agent Mode helpers ────────────────────────────────────────────────────────
|
||||
@ -268,46 +271,58 @@ def _clear_chat_dialog():
|
||||
|
||||
# ── Normal Chat ───────────────────────────────────────────────────────────────
|
||||
|
||||
def _render_search_panel():
|
||||
"""Render the collapsible web search panel above the chat history.
|
||||
|
||||
Stores results in session_state.search_results so they are automatically
|
||||
injected as context into the next message the user sends.
|
||||
"""
|
||||
search_results = st.session_state.get("search_results", [])
|
||||
label = f"🔍 Web Search ({len(search_results)} result{'s' if len(search_results) != 1 else ''} active)" if search_results else "🔍 Web Search"
|
||||
|
||||
with st.expander(label, expanded=False):
|
||||
col_input, col_btn = st.columns([5, 1])
|
||||
with col_input:
|
||||
query = st.text_input(
|
||||
"Search query",
|
||||
key="search_query_input",
|
||||
placeholder="e.g. Python asyncio best practices",
|
||||
label_visibility="collapsed",
|
||||
)
|
||||
with col_btn:
|
||||
search_clicked = st.button("Search", use_container_width=True)
|
||||
|
||||
if search_clicked and query.strip():
|
||||
with st.spinner("Searching..."):
|
||||
sm = SearchManager()
|
||||
results = sm.perform_search(query.strip())
|
||||
if results:
|
||||
st.session_state.search_results = results
|
||||
st.rerun()
|
||||
else:
|
||||
st.warning("No results found.")
|
||||
|
||||
# Display active results with a clear button.
|
||||
if search_results:
|
||||
st.caption("Results will be injected as context into your next message.")
|
||||
for r in search_results:
|
||||
st.markdown(f"**{r['title']}** \n{r['snippet']} \n[{r['url']}]({r['url']})")
|
||||
st.divider()
|
||||
if st.button("Clear search results", use_container_width=True):
|
||||
st.session_state.search_results = []
|
||||
st.rerun()
|
||||
|
||||
|
||||
def render_normal_chat():
|
||||
"""Render the standard multi-turn chat interface.
|
||||
|
||||
Execution order on every rerun:
|
||||
1. Apply model/token settings from the Settings panel (5e)
|
||||
2. Consume any pending debug message from the editor (5f)
|
||||
3. Replay chat history
|
||||
4. Handle chat input with updated system-prompt logic (5g)
|
||||
5. Render Clear Chat button and Settings expander (5d, 5h)
|
||||
On the first message the system prompt is injected into the history,
|
||||
including any active search results as context.
|
||||
Each subsequent message appends to the same conversation so the AI retains
|
||||
full context throughout the session. If search results are active when the
|
||||
user sends a message, they are prepended to that message as a context block.
|
||||
"""
|
||||
chat_manager: ChatManager = st.session_state.chat_manager
|
||||
|
||||
# 5e — Apply model/token overrides from the Settings panel before any API call.
|
||||
if st.session_state.get("selected_model"):
|
||||
chat_manager.model = st.session_state.selected_model
|
||||
if "chat_max_tokens" in st.session_state:
|
||||
chat_manager.max_tokens = st.session_state.chat_max_tokens
|
||||
|
||||
# 5f — Consume a debug message forwarded from the editor's "Debug with AI" button.
|
||||
pending_debug = st.session_state.pop("pending_debug_message", None)
|
||||
if pending_debug:
|
||||
if not chat_manager.get_history():
|
||||
custom_prompt = st.session_state.get("custom_system_prompt", "").strip()
|
||||
if custom_prompt:
|
||||
chat_manager.add_message("system", custom_prompt)
|
||||
else:
|
||||
file_ctx = _build_file_context()
|
||||
system_prompt = SystemPrompter.generate_prompt(file_ctx)
|
||||
chat_manager.add_message("system", system_prompt)
|
||||
|
||||
with st.spinner("Sending debug info to AI..."):
|
||||
try:
|
||||
ai_response = chat_manager.send_message(pending_debug)
|
||||
except Exception as e:
|
||||
ai_response = f"Error: {e}"
|
||||
|
||||
st.session_state.chat_history.append({"role": "user", "content": pending_debug})
|
||||
st.session_state.chat_history.append({"role": "assistant", "content": ai_response})
|
||||
st.rerun()
|
||||
return
|
||||
_render_search_panel()
|
||||
|
||||
# Replay the conversation history as chat bubbles (skip system messages).
|
||||
for message in st.session_state.chat_history:
|
||||
@ -317,33 +332,82 @@ def render_normal_chat():
|
||||
st.markdown(message["content"])
|
||||
|
||||
# Chat input — Enter to send, no extra button needed.
|
||||
user_input = st.chat_input("Type your message here...")
|
||||
# Supports /search <query> and /search clear as special commands.
|
||||
user_input = st.chat_input("Type a message or /search <query>...")
|
||||
if user_input:
|
||||
stripped = user_input.strip()
|
||||
|
||||
# ── /search command ───────────────────────────────────────────────────
|
||||
if stripped.lower().startswith("/search"):
|
||||
arg = stripped[len("/search"):].strip()
|
||||
|
||||
with st.chat_message("user"):
|
||||
st.markdown(stripped)
|
||||
|
||||
if arg.lower() == "clear" or arg == "":
|
||||
# /search clear (or bare /search) — remove active results.
|
||||
st.session_state.search_results = []
|
||||
with st.chat_message("assistant"):
|
||||
st.markdown("Search context cleared.")
|
||||
st.session_state.chat_history.append({"role": "user", "content": stripped})
|
||||
st.session_state.chat_history.append({"role": "assistant", "content": "Search context cleared."})
|
||||
else:
|
||||
# /search <query> — run search and store results in context.
|
||||
with st.chat_message("assistant"):
|
||||
with st.spinner(f'Searching for "{arg}"...'):
|
||||
sm = SearchManager()
|
||||
results = sm.perform_search(arg)
|
||||
|
||||
if results:
|
||||
st.session_state.search_results = results
|
||||
summary = f"Found {len(results)} result(s) for **{arg}**. They are now in context for this chat session.\n\n"
|
||||
for i, r in enumerate(results, 1):
|
||||
summary += f"**{i}. [{r['title']}]({r['url']})** \n{r['snippet']}\n\n"
|
||||
st.markdown(summary)
|
||||
response_text = summary
|
||||
else:
|
||||
msg = f'No results found for "{arg}".'
|
||||
st.warning(msg)
|
||||
response_text = msg
|
||||
|
||||
st.session_state.chat_history.append({"role": "user", "content": stripped})
|
||||
st.session_state.chat_history.append({"role": "assistant", "content": response_text})
|
||||
|
||||
st.rerun()
|
||||
return
|
||||
|
||||
# ── Normal chat message ───────────────────────────────────────────────
|
||||
chat_manager = st.session_state.chat_manager
|
||||
search_results = st.session_state.get("search_results", [])
|
||||
|
||||
# 5g — System-prompt logic: inject on first message, update on file change.
|
||||
if not chat_manager.get_history():
|
||||
custom_prompt = st.session_state.get("custom_system_prompt", "").strip()
|
||||
if custom_prompt:
|
||||
chat_manager.add_message("system", custom_prompt)
|
||||
else:
|
||||
file_ctx = _build_file_context()
|
||||
system_prompt = SystemPrompter.generate_prompt(file_ctx)
|
||||
system_prompt = SystemPrompter.generate_prompt()
|
||||
chat_manager.add_message("system", system_prompt)
|
||||
elif st.session_state.get("active_file") and st.session_state.get("include_file_context", True):
|
||||
# Follow-up messages: refresh the system prompt when the active file changes.
|
||||
history = chat_manager.get_history()
|
||||
if history and history[0]["role"] == "system":
|
||||
file_ctx = _build_file_context()
|
||||
if file_ctx:
|
||||
history[0]["content"] = SystemPrompter.generate_prompt(file_ctx)
|
||||
|
||||
# If search results are active, prepend them as a context block so the
|
||||
# AI can reference them regardless of where in the conversation we are.
|
||||
if search_results:
|
||||
context_block = "<search_context>\n"
|
||||
for r in search_results:
|
||||
context_block += (
|
||||
f"Title: {r['title']}\n"
|
||||
f"URL: {r['url']}\n"
|
||||
f"Snippet: {r['snippet']}\n\n"
|
||||
)
|
||||
context_block += "</search_context>\n\n"
|
||||
message_to_send = context_block + user_input
|
||||
else:
|
||||
message_to_send = user_input
|
||||
|
||||
# Show the original user text in the UI (not the context-enriched version).
|
||||
with st.chat_message("user"):
|
||||
st.markdown(user_input)
|
||||
|
||||
with st.chat_message("assistant"):
|
||||
with st.spinner("Thinking..."):
|
||||
try:
|
||||
ai_response = chat_manager.send_message(user_input)
|
||||
ai_response = chat_manager.send_message(message_to_send)
|
||||
except Exception as e:
|
||||
ai_response = f"Error: {e}"
|
||||
st.markdown(ai_response)
|
||||
|
||||
@ -259,10 +259,10 @@ def render_filetree_arborist(tree):
|
||||
selected = tree_view(
|
||||
data=data,
|
||||
icons={"open": "📂", "closed": "📁"},
|
||||
height=200,
|
||||
height=400,
|
||||
selection=None,
|
||||
select_internal_nodes=True, # allow clicking folder names, not just files
|
||||
open_by_default=False
|
||||
open_by_default=True
|
||||
)
|
||||
|
||||
return selected
|
||||
|
||||
@ -111,24 +111,12 @@ def init_state():
|
||||
if "agent_pending_action" not in st.session_state:
|
||||
st.session_state.agent_pending_action = None
|
||||
|
||||
# Whether to inject the currently open file as context into the system prompt
|
||||
if "include_file_context" not in st.session_state:
|
||||
st.session_state.include_file_context = True
|
||||
|
||||
# Optional custom system prompt entered by the user in Settings (overrides default)
|
||||
if "custom_system_prompt" not in st.session_state:
|
||||
st.session_state.custom_system_prompt = ""
|
||||
|
||||
# Holds a pre-built debug message to be sent to the AI on the next chat render
|
||||
if "pending_debug_message" not in st.session_state:
|
||||
st.session_state.pending_debug_message = None
|
||||
|
||||
# Per-file execution results: {file_path: {stdout, stderr, return_code, ast_error}}
|
||||
if "exec_results" not in st.session_state:
|
||||
st.session_state.exec_results = {}
|
||||
# Web search results to be injected as context into the next AI message.
|
||||
# List of {"title": str, "url": str, "snippet": str} dicts, or empty list.
|
||||
if "search_results" not in st.session_state:
|
||||
st.session_state.search_results = []
|
||||
|
||||
|
||||
|
||||
# REVIEW: dead code — state.py is never run as a script; this guard is useless here because
|
||||
# init_state() requires a running Streamlit session (st.session_state) to work.
|
||||
if __name__ == "__main__":
|
||||
init_state()
|
||||
|
||||
@ -24,3 +24,6 @@ python-dotenv>=1.0.0
|
||||
|
||||
#For code editor functionality
|
||||
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