Merge pull request #3 from kacperkwapisz/feat/rate-limiting-and-middleware

Add rate limiting, route filtering, and webhook improvements
This commit is contained in:
0xK
2026-03-24 15:57:00 +01:00
committed by GitHub
2 changed files with 126 additions and 12 deletions
+3 -1
View File
@@ -216,4 +216,6 @@ __marimo__/
.streamlit/secrets.toml .streamlit/secrets.toml
# poke-mail config (contains credentials) # poke-mail config (contains credentials)
config.yml config.yml
# macOS
.DS_Store
+123 -11
View File
@@ -4,6 +4,7 @@ import logging
import os import os
import smtplib import smtplib
import ssl import ssl
import time
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from datetime import date from datetime import date
from email import policy from email import policy
@@ -15,10 +16,14 @@ from typing import Optional
import hmac import hmac
import httpx import httpx
import uvicorn
import yaml import yaml
from imapclient import IMAPClient from imapclient import IMAPClient
from fastmcp import FastMCP, Context from fastmcp import FastMCP, Context
from fastmcp.server.auth import TokenVerifier, AccessToken from fastmcp.server.auth import TokenVerifier, AccessToken
from starlette.middleware import Middleware
from starlette.responses import JSONResponse, Response
from starlette.types import ASGIApp, Receive, Scope, Send
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
logger = logging.getLogger("poke-mail") logger = logging.getLogger("poke-mail")
@@ -43,6 +48,108 @@ class ApiKeyAuth(TokenVerifier):
return None return None
# ---------------------------------------------------------------------------
# Route filter — silently drop non-MCP requests
# ---------------------------------------------------------------------------
class DropNonMCPRoutes:
"""Return empty 404 for any path outside /mcp — reveals nothing to scanners."""
def __init__(self, app: ASGIApp):
self.app = app
async def __call__(self, scope: Scope, receive: Receive, send: Send):
if scope["type"] == "http" and not scope["path"].startswith("/mcp"):
response = Response(status_code=404)
await response(scope, receive, send)
return
await self.app(scope, receive, send)
class RateLimitMiddleware:
"""Per-IP sliding window rate limiter with separate buckets for GET and POST.
GET /mcp (health/polling) gets a tighter limit to curb excessive polling.
POST /mcp (tool calls) gets a higher limit so real work isn't blocked.
"""
MAX_TRACKED_IPS = 1024
def __init__(self, app: ASGIApp):
self.app = app
self.get_rpm = int(os.environ.get("RATE_LIMIT_GET_RPM", "30"))
self.post_rpm = int(os.environ.get("RATE_LIMIT_POST_RPM", "120"))
self.window = 60 # seconds
self._hits: dict[str, list[float]] = {}
self._last_cleanup = time.monotonic()
def _client_ip(self, scope: Scope) -> str:
# Only trust X-Forwarded-For from known proxy — take the rightmost
# entry (closest to our server) to resist spoofing via prepended IPs.
for header_name, header_val in scope.get("headers", []):
if header_name == b"x-forwarded-for":
parts = header_val.decode().split(",")
return parts[-1].strip()
client = scope.get("client")
return client[0] if client else "unknown"
def _cleanup_stale(self, now: float) -> None:
"""Periodically evict stale IPs to bound memory usage."""
if now - self._last_cleanup < self.window:
return
self._last_cleanup = now
cutoff = now - self.window
stale = [ip for ip, ts in self._hits.items() if not ts or ts[-1] <= cutoff]
for ip in stale:
del self._hits[ip]
# Hard cap: if still too many, drop the oldest entries
if len(self._hits) > self.MAX_TRACKED_IPS:
by_recency = sorted(self._hits, key=lambda ip: self._hits[ip][-1])
for ip in by_recency[: len(self._hits) - self.MAX_TRACKED_IPS]:
del self._hits[ip]
def _is_limited(self, bucket: str, rpm: int) -> tuple[bool, int]:
now = time.monotonic()
self._cleanup_stale(now)
timestamps = self._hits.get(bucket, [])
cutoff = now - self.window
timestamps = [t for t in timestamps if t > cutoff]
self._hits[bucket] = timestamps
if len(timestamps) >= rpm:
oldest = timestamps[0]
retry_after = int(oldest + self.window - now) + 1
return True, max(retry_after, 1)
timestamps.append(now)
return False, 0
async def __call__(self, scope: Scope, receive: Receive, send: Send):
if scope["type"] != "http":
await self.app(scope, receive, send)
return
ip = self._client_ip(scope)
method = scope.get("method", "GET")
if method == "POST":
bucket, rpm = f"{ip}:post", self.post_rpm
else:
bucket, rpm = f"{ip}:get", self.get_rpm
limited, retry_after = self._is_limited(bucket, rpm)
if limited:
response = JSONResponse(
{"error": "rate_limited", "retry_after": retry_after},
status_code=429,
headers={"Retry-After": str(retry_after)},
)
await response(scope, receive, send)
return
await self.app(scope, receive, send)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Config # Config
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -262,8 +369,12 @@ def detect_drafts_folder(client: IMAPClient) -> str:
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
async def forward_to_poke(email_data: dict, webhook_url: str, api_key: str) -> bool: async def forward_to_poke(
email_data: dict, account: dict, webhook_url: str, api_key: str
) -> bool:
payload = { payload = {
"account_id": account["id"],
"from_address": account["from_address"],
"from": email_data["from"], "from": email_data["from"],
"to": email_data["to"], "to": email_data["to"],
"subject": email_data["subject"], "subject": email_data["subject"],
@@ -361,12 +472,9 @@ async def watch_folder(
break break
await asyncio.to_thread(client.idle_done) await asyncio.to_thread(client.idle_done)
has_new = any( logger.debug(
isinstance(r, tuple) and len(r) >= 2 and r[1] == b"EXISTS" "[%s/%s] IDLE responses: %s", account["id"], folder, responses
for r in responses
) )
if not has_new:
continue
uids = await asyncio.to_thread(client.search, ["UNSEEN"]) uids = await asyncio.to_thread(client.search, ["UNSEEN"])
# Only forward emails that arrived after we started watching # Only forward emails that arrived after we started watching
@@ -382,9 +490,10 @@ async def watch_folder(
if not raw: if not raw:
continue continue
email_data = parse_email_message(raw) email_data = parse_email_message(raw)
await forward_to_poke(email_data, webhook_url, api_key) await forward_to_poke(email_data, account, webhook_url, api_key)
await asyncio.to_thread(client.set_flags, new_uids, [b"\\Seen"]) await asyncio.to_thread(client.set_flags, new_uids, [b"\\Seen"])
existing_unseen.update(new_uids)
except asyncio.CancelledError: except asyncio.CancelledError:
break break
@@ -426,7 +535,7 @@ async def _poll_folder(
if not raw: if not raw:
continue continue
email_data = parse_email_message(raw) email_data = parse_email_message(raw)
await forward_to_poke(email_data, webhook_url, api_key) await forward_to_poke(email_data, account, webhook_url, api_key)
await asyncio.to_thread(client.set_flags, uids, [b"\\Seen"]) await asyncio.to_thread(client.set_flags, uids, [b"\\Seen"])
except Exception as e: except Exception as e:
logger.warning("[%s/%s] Poll error: %s", account["id"], folder, e) logger.warning("[%s/%s] Poll error: %s", account["id"], folder, e)
@@ -503,8 +612,6 @@ mcp = FastMCP("poke-mail", lifespan=lifespan, auth=auth)
@mcp.custom_route("/mcp", methods=["GET"]) @mcp.custom_route("/mcp", methods=["GET"])
async def health(request): async def health(request):
from starlette.responses import JSONResponse
return JSONResponse({"status": "ok"}) return JSONResponse({"status": "ok"})
@@ -942,6 +1049,7 @@ async def get_server_info(ctx: Context) -> dict:
account_info.append( account_info.append(
{ {
"id": acc["id"], "id": acc["id"],
"from_address": acc["from_address"],
"imap_host": acc["imap_host"], "imap_host": acc["imap_host"],
"smtp_host": acc["smtp_host"], "smtp_host": acc["smtp_host"],
"watch_folders": acc.get("watch_folders", []), "watch_folders": acc.get("watch_folders", []),
@@ -965,4 +1073,8 @@ if __name__ == "__main__":
port = int(os.environ.get("PORT", 3000)) port = int(os.environ.get("PORT", 3000))
host = "0.0.0.0" host = "0.0.0.0"
logger.info("Starting poke-mail on %s:%d", host, port) logger.info("Starting poke-mail on %s:%d", host, port)
mcp.run(transport="http", host=host, port=port, stateless_http=True) app = mcp.http_app(
middleware=[Middleware(DropNonMCPRoutes), Middleware(RateLimitMiddleware)],
stateless_http=True,
)
uvicorn.run(app, host=host, port=port)