Jarian Cottingham 40454e385f fix: batch fix all 8 open issues
- #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
2026-07-05 22:04:57 +00:00

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())