Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
79 changes: 58 additions & 21 deletions gateway/platforms/webhook.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,8 @@
import re
import subprocess
import time
from typing import Any, Dict, List, Optional
from collections import deque
from typing import Any, Deque, Dict, List, Optional

try:
from aiohttp import web
Expand Down Expand Up @@ -67,6 +68,7 @@
DEFAULT_PORT = 8644
_INSECURE_NO_AUTH = "INSECURE_NO_AUTH"
_DYNAMIC_ROUTES_FILENAME = "webhook_subscriptions.json"
_RATE_WINDOW_SECONDS = 60.0

# Hostnames/IP literals that only serve connections originating on the same
# machine. Anything else is treated as a public bind for safety-rail purposes.
Expand Down Expand Up @@ -122,6 +124,7 @@ def __init__(self, config: PlatformConfig):
# back to the "log" deliver type.
self._delivery_info: Dict[str, dict] = {}
self._delivery_info_created: Dict[str, float] = {}
self._delivery_info_order: Deque[tuple[float, str]] = deque()

# Reference to gateway runner for cross-platform delivery (set externally)
self.gateway_runner = None
Expand All @@ -130,9 +133,10 @@ def __init__(self, config: PlatformConfig):
# Prevents duplicate agent runs when webhook providers retry.
self._seen_deliveries: Dict[str, float] = {}
self._idempotency_ttl: int = 3600 # 1 hour
self._seen_deliveries_next_prune_at: float = 0.0

# Rate limiting: per-route timestamps in a fixed window.
self._rate_counts: Dict[str, List[float]] = {}
self._rate_counts: Dict[str, Deque[float]] = {}
self._rate_limit: int = int(config.extra.get("rate_limit", 30)) # per minute

# Body size limit (auth-before-body pattern)
Expand Down Expand Up @@ -271,15 +275,57 @@ def _prune_delivery_info(self, now: float) -> None:
on each POST so the dict size is bounded by ``rate_limit * TTL``
even if many webhooks fire and never receive a final response.
"""
if len(self._delivery_info_order) < len(self._delivery_info_created):
self._delivery_info_order = deque(
(created_at, key)
for key, created_at in sorted(
self._delivery_info_created.items(), key=lambda item: item[1]
)
)
cutoff = now - self._idempotency_ttl
stale = [
k
for k, t in self._delivery_info_created.items()
if t < cutoff
]
while self._delivery_info_order and self._delivery_info_order[0][0] < cutoff:
created_at, key = self._delivery_info_order.popleft()
if self._delivery_info_created.get(key) != created_at:
continue
self._delivery_info.pop(key, None)
self._delivery_info_created.pop(key, None)

def _prune_seen_deliveries(self, now: float) -> None:
"""Occasionally prune expired delivery IDs without scanning every POST."""
if now < self._seen_deliveries_next_prune_at:
return
cutoff = now - self._idempotency_ttl
stale = [k for k, t in self._seen_deliveries.items() if t < cutoff]
for k in stale:
self._delivery_info.pop(k, None)
self._delivery_info_created.pop(k, None)
self._seen_deliveries.pop(k, None)
self._seen_deliveries_next_prune_at = now + min(60.0, max(1.0, self._idempotency_ttl / 10))

def _record_rate_limit_hit(self, route_name: str, now: float) -> bool:
"""Return True if route is still within limit after recording this hit."""
window = self._rate_counts.get(route_name)
if not isinstance(window, deque):
new_window: Deque[float] = deque(window or ())
self._rate_counts[route_name] = new_window
window = new_window
cutoff = now - _RATE_WINDOW_SECONDS
while window and window[0] < cutoff:
window.popleft()
if len(window) >= self._rate_limit:
return False
window.append(now)
return True

def _record_delivery_id(self, delivery_id: str, now: float) -> bool:
"""Return True when this delivery should be processed."""
seen_at = self._seen_deliveries.get(delivery_id)
if seen_at is not None and now - seen_at < self._idempotency_ttl:
return False
if seen_at is not None:
self._seen_deliveries.pop(delivery_id, None)
self._seen_deliveries[delivery_id] = now
if len(self._seen_deliveries) > max(self._rate_limit * 2, 128):
self._prune_seen_deliveries(now)
return True

async def get_chat_info(self, chat_id: str) -> Dict[str, Any]:
return {"name": chat_id, "type": "webhook"}
Expand Down Expand Up @@ -413,13 +459,10 @@ async def _handle_webhook(self, request: "web.Request") -> "web.Response":

# ── Rate limiting (after auth) ───────────────────────────
now = time.time()
window = self._rate_counts.setdefault(route_name, [])
window[:] = [t for t in window if now - t < 60]
if len(window) >= self._rate_limit:
if not self._record_rate_limit_hit(route_name, now):
return web.json_response(
{"error": "Rate limit exceeded"}, status=429
)
window.append(now)

# Parse payload
try:
Expand Down Expand Up @@ -504,21 +547,14 @@ async def _handle_webhook(self, request: "web.Request") -> "web.Response":
# ── Idempotency ─────────────────────────────────────────
# Skip duplicate deliveries (webhook retries).
now = time.time()
# Prune expired entries
self._seen_deliveries = {
k: v
for k, v in self._seen_deliveries.items()
if now - v < self._idempotency_ttl
}
if delivery_id in self._seen_deliveries:
if not self._record_delivery_id(delivery_id, now):
logger.info(
"[webhook] Skipping duplicate delivery %s", delivery_id
)
return web.json_response(
{"status": "duplicate", "delivery_id": delivery_id},
status=200,
)
self._seen_deliveries[delivery_id] = now

# ── Direct delivery mode (deliver_only) ─────────────────
# Skip the agent entirely — the rendered prompt IS the message we
Expand Down Expand Up @@ -594,6 +630,7 @@ async def _handle_webhook(self, request: "web.Request") -> "web.Response":
}
self._delivery_info[session_chat_id] = deliver_config
self._delivery_info_created[session_chat_id] = now
self._delivery_info_order.append((now, session_chat_id))
self._prune_delivery_info(now)

# Build source and event
Expand Down
55 changes: 54 additions & 1 deletion tests/gateway/test_webhook_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
import hmac
import json
import time
from collections import deque
from unittest.mock import AsyncMock, MagicMock, patch

import pytest
Expand Down Expand Up @@ -680,7 +681,7 @@ async def test_rate_limit_window_resets(self):
assert resp.status == 202

# Backdate all rate-limit timestamps to > 60 seconds ago
adapter._rate_counts["limited"] = [time.time() - 120]
adapter._rate_counts["limited"] = deque([time.time() - 120])

resp = await cli.post(
"/webhooks/limited",
Expand All @@ -689,6 +690,33 @@ async def test_rate_limit_window_resets(self):
)
assert resp.status == 202 # allowed again

def test_rate_limit_prunes_incrementally_from_left(self):
"""Expired rate-limit entries are pruned without rebuilding the window."""
adapter = _make_adapter(rate_limit=2)
adapter._rate_counts["limited"] = deque([100.0, 220.0])

assert adapter._record_rate_limit_hit("limited", 221.0) is True

window = adapter._rate_counts["limited"]
assert list(window) == [220.0, 221.0]

def test_seen_delivery_ttl_is_checked_per_delivery_without_full_prune(self):
"""Expired delivery IDs can reprocess even when stale siblings remain."""
adapter = _make_adapter(rate_limit=1)
adapter._idempotency_ttl = 60
adapter._seen_deliveries = {
"expired-target": 100.0,
"expired-sibling": 101.0,
"fresh-sibling": 155.0,
}

now = 200.0
assert adapter._record_delivery_id("expired-target", now) is True

assert adapter._seen_deliveries["expired-target"] == now
assert "expired-sibling" in adapter._seen_deliveries
assert "fresh-sibling" in adapter._seen_deliveries


# ===================================================================
# Body size limit
Expand Down Expand Up @@ -838,6 +866,31 @@ async def test_delivery_info_pruned_via_ttl(self):
assert "webhook:test:new" in adapter._delivery_info
assert "webhook:test:new" in adapter._delivery_info_created

@pytest.mark.asyncio
async def test_delivery_info_prune_uses_ordered_incremental_queue(self):
"""Delivery-info TTL pruning stops at the first fresh queued entry."""
adapter = _make_adapter()
adapter._idempotency_ttl = 60
now = 1000.0
for key, created_at in (
("webhook:test:old", now - 120),
("webhook:test:new", now - 5),
("webhook:test:newer", now),
):
adapter._delivery_info[key] = {"deliver": "log"}
adapter._delivery_info_created[key] = created_at
adapter._delivery_info_order.append((created_at, key))

adapter._prune_delivery_info(now)

assert "webhook:test:old" not in adapter._delivery_info
assert "webhook:test:new" in adapter._delivery_info
assert "webhook:test:newer" in adapter._delivery_info
assert list(adapter._delivery_info_order) == [
(now - 5, "webhook:test:new"),
(now, "webhook:test:newer"),
]


# ===================================================================
# check_webhook_requirements
Expand Down
Loading