- #1 SearXNG auth: remove external port, internal-only network, enable limiter - #2 Health checks: add docker-compose healthcheck + depends_on condition - #3 Rate limiting: add per-IP RateLimitMiddleware (30 req/min) - #4 Hardcoded secret: replace with ${SEARXNG_SECRET} env var - #5 Debug disclosure: debug=False, generic error messages, no stack traces - #6 Health endpoint: add /health route returning JSON status - #7 asyncio deprecation: get_event_loop() -> get_running_loop() - #8 httpx reuse: module-level singleton AsyncClient with connection pool
282 lines
7.9 KiB
Python
282 lines
7.9 KiB
Python
import asyncio
|
|
import json
|
|
import logging
|
|
import os
|
|
import signal
|
|
import sys
|
|
from collections import defaultdict
|
|
from contextlib import asynccontextmanager
|
|
from datetime import datetime, timezone
|
|
|
|
import httpx
|
|
|
|
from mcp.server import Server
|
|
from mcp.server.sse import SseServerTransport
|
|
from starlette.applications import Starlette
|
|
from starlette.middleware.base import BaseHTTPMiddleware
|
|
from starlette.requests import Request
|
|
from starlette.responses import JSONResponse, Response
|
|
from starlette.routing import Mount, Route
|
|
from mcp.types import Tool, TextContent
|
|
|
|
logging.basicConfig(
|
|
level=logging.INFO,
|
|
format="%(asctime)s [%(levelname)s] %(message)s",
|
|
stream=sys.stderr,
|
|
)
|
|
logger = logging.getLogger("duckduckgo-mcp")
|
|
|
|
app = Server("duckduckgo-search")
|
|
PORT = int(os.environ.get("MCP_PORT", "3002"))
|
|
SEARXNG_URL = os.environ.get("SEARXNG_URL", "http://searxng:8080")
|
|
RATE_LIMIT_SECONDS = float(os.environ.get("RATE_LIMIT_SECONDS", "3"))
|
|
|
|
|
|
# --- HTTPX singleton (#8) ---
|
|
|
|
_httpx_client: httpx.AsyncClient | None = None
|
|
|
|
|
|
async def get_httpx_client() -> httpx.AsyncClient:
|
|
global _httpx_client
|
|
if _httpx_client is None:
|
|
_httpx_client = httpx.AsyncClient(
|
|
timeout=httpx.Timeout(15.0),
|
|
limits=httpx.Limits(max_connections=32, max_keepalive_connections=16),
|
|
)
|
|
return _httpx_client
|
|
|
|
|
|
async def close_httpx_client():
|
|
global _httpx_client
|
|
if _httpx_client is not None:
|
|
await _httpx_client.aclose()
|
|
_httpx_client = None
|
|
|
|
|
|
@asynccontextmanager
|
|
async def lifespan(app):
|
|
logger.info("HTTPX client initialized")
|
|
try:
|
|
yield
|
|
finally:
|
|
await close_httpx_client()
|
|
logger.info("HTTPX client closed")
|
|
|
|
|
|
# --- Rate limiter ---
|
|
|
|
class RateLimiter:
|
|
def __init__(self, min_interval: float):
|
|
self.min_interval = min_interval
|
|
self._last_request = None
|
|
self._lock = asyncio.Lock()
|
|
|
|
async def acquire(self):
|
|
async with self._lock:
|
|
now = asyncio.get_running_loop().time()
|
|
if self._last_request:
|
|
elapsed = now - self._last_request
|
|
wait = self.min_interval - elapsed
|
|
if wait > 0:
|
|
logger.info(f"Rate limit: waiting {wait:.1f}s")
|
|
await asyncio.sleep(wait)
|
|
self._last_request = asyncio.get_running_loop().time()
|
|
|
|
|
|
rate_limiter = RateLimiter(RATE_LIMIT_SECONDS)
|
|
|
|
|
|
# --- Per-IP rate limiting middleware (#3) ---
|
|
|
|
class RateLimitMiddleware(BaseHTTPMiddleware):
|
|
"""Per-IP rate limiting for HTTP endpoints."""
|
|
|
|
def __init__(self, app, max_requests: int = 30, window_seconds: int = 60):
|
|
super().__init__(app)
|
|
self.max_requests = max_requests
|
|
self.window_seconds = window_seconds
|
|
self._requests: dict[str, list[float]] = defaultdict(list)
|
|
self._lock = asyncio.Lock()
|
|
|
|
async def dispatch(self, request: Request, call_next):
|
|
if request.url.path == "/health":
|
|
return await call_next(request)
|
|
|
|
client_ip = request.client.host if request.client else "unknown"
|
|
now = datetime.now(timezone.utc).timestamp()
|
|
|
|
async with self._lock:
|
|
timestamps = self._requests[client_ip]
|
|
cutoff = now - self.window_seconds
|
|
self._requests[client_ip] = [t for t in timestamps if t > cutoff]
|
|
if len(self._requests[client_ip]) >= self.max_requests:
|
|
return JSONResponse(
|
|
{"error": "Rate limit exceeded. Try again later."},
|
|
status_code=429,
|
|
)
|
|
self._requests[client_ip].append(now)
|
|
|
|
response = await call_next(request)
|
|
return response
|
|
|
|
|
|
# --- Search logic ---
|
|
|
|
async def do_search(query: str, num_results: int, engine: str):
|
|
"""Shared search logic."""
|
|
await rate_limiter.acquire()
|
|
|
|
client = await get_httpx_client()
|
|
resp = await client.get(
|
|
f"{SEARXNG_URL}/search",
|
|
params={
|
|
"q": query,
|
|
"format": "json",
|
|
"engines": engine,
|
|
"categories": "general",
|
|
"language": "en",
|
|
},
|
|
)
|
|
resp.raise_for_status()
|
|
data = resp.json()
|
|
|
|
results = data.get("results", [])[:num_results]
|
|
return [
|
|
{
|
|
"title": r.get("title", ""),
|
|
"url": r.get("url", ""),
|
|
"snippet": r.get("content", "")[:200],
|
|
}
|
|
for r in results
|
|
]
|
|
|
|
|
|
# --- HTTP endpoint ---
|
|
|
|
async def search_http(request: Request):
|
|
query = request.query_params.get("q", "")
|
|
if not query:
|
|
return JSONResponse({"error": "'q' parameter required"}, status_code=400)
|
|
|
|
num = min(int(request.query_params.get("num", "10")), 20)
|
|
try:
|
|
results = await do_search(query, num, "duckduckgo")
|
|
return JSONResponse({"query": query, "results": results})
|
|
except Exception as e:
|
|
logger.exception("Search failed")
|
|
return JSONResponse({"error": "Search service unavailable"}, status_code=503)
|
|
|
|
|
|
# --- Health check (#6) ---
|
|
|
|
async def health_check(request: Request):
|
|
return JSONResponse({"status": "ok", "service": "duckduckgo-mcp"})
|
|
|
|
|
|
# --- MCP tool handlers ---
|
|
|
|
@app.list_tools()
|
|
async def list_tools() -> list[Tool]:
|
|
return [
|
|
Tool(
|
|
name="duckduckgo_search",
|
|
description="Search DuckDuckGo and return results with titles, URLs, and snippets.",
|
|
inputSchema={
|
|
"type": "object",
|
|
"properties": {
|
|
"query": {"type": "string", "description": "The search query."},
|
|
"num_results": {
|
|
"type": "integer",
|
|
"description": "Maximum number of results (default 10).",
|
|
"default": 10,
|
|
},
|
|
},
|
|
"required": ["query"],
|
|
},
|
|
)
|
|
]
|
|
|
|
|
|
@app.call_tool()
|
|
async def call_tool(name: str, arguments: dict) -> list[TextContent]:
|
|
if name != "duckduckgo_search":
|
|
raise ValueError(f"Unknown tool: {name}")
|
|
|
|
query = arguments.get("query", "")
|
|
if not query:
|
|
return [TextContent(type="text", text="Error: 'query' is required.")]
|
|
|
|
num_results = min(int(arguments.get("num_results", 10)), 20)
|
|
|
|
try:
|
|
results = await do_search(query, num_results, "duckduckgo")
|
|
except Exception as e:
|
|
logger.exception("Search failed")
|
|
return [TextContent(type="text", text="Search failed: service unavailable")]
|
|
|
|
if not results:
|
|
return [TextContent(type="text", text=f"No results found for: {query}")]
|
|
|
|
lines = [f"Search results for: {query}\n"]
|
|
for i, r in enumerate(results, 1):
|
|
lines.append(f"{i}. {r['title']}")
|
|
lines.append(f" URL: {r['url']}")
|
|
if r["snippet"]:
|
|
lines.append(f" {r['snippet']}")
|
|
lines.append("")
|
|
|
|
return [TextContent(type="text", text="\n".join(lines))]
|
|
|
|
|
|
# --- SSE MCP transport ---
|
|
|
|
sse = SseServerTransport("/messages/")
|
|
|
|
|
|
async def handle_sse(request):
|
|
async with sse.connect_sse(
|
|
request.scope, request.receive, request._send
|
|
) as (read_stream, write_stream):
|
|
await app.run(
|
|
read_stream,
|
|
write_stream,
|
|
app.create_initialization_options(),
|
|
)
|
|
return Response()
|
|
|
|
|
|
starlette_app = Starlette(
|
|
debug=False,
|
|
routes=[
|
|
Route("/health", endpoint=health_check),
|
|
Route("/search", endpoint=search_http),
|
|
Route("/sse", endpoint=handle_sse),
|
|
Mount("/messages/", app=sse.handle_post_message),
|
|
],
|
|
middleware=[RateLimitMiddleware],
|
|
)
|
|
|
|
|
|
async def main():
|
|
logger.info(f"DuckDuckGo MCP server ready on port {PORT}")
|
|
|
|
import uvicorn
|
|
|
|
config = uvicorn.Config(starlette_app, host="0.0.0.0", port=PORT, log_level="info")
|
|
server = uvicorn.Server(config)
|
|
|
|
loop = asyncio.get_running_loop()
|
|
|
|
def handle_signal():
|
|
server.should_exit = True
|
|
|
|
for sig in (signal.SIGINT, signal.SIGTERM):
|
|
loop.add_signal_handler(sig, handle_signal)
|
|
|
|
await server.serve()
|
|
|
|
|
|
if __name__ == "__main__":
|
|
asyncio.run(main())
|