diff --git a/README.rst b/README.rst index ddc50f3d..8f2e611e 100644 --- a/README.rst +++ b/README.rst @@ -47,6 +47,26 @@ Connect to any S7 PLC:: No native libraries or platform-specific dependencies are required. +Protecting PLCs with request rate limits +---------------------------------------- + +Request limiting is opt-in and applies to every S7 PDU sent by both ``Client`` +and ``AsyncClient``. A multi-variable request counts once; an operation split +across several PDUs counts each PDU:: + + client = Client( + max_requests_per_second=10, + rate_limit_algorithm="fixed", # or "token_bucket" + rate_limit_behavior="block", # or "raise" / "drop" + ) + +``fixed`` spaces requests evenly. ``token_bucket`` permits a burst (one second +of requests by default, configurable with ``rate_limit_burst``) and then +refills at the configured rate. The default rate is ``0``, which disables the +limiter. ``raise`` and ``drop`` both raise ``S7RateLimitError`` immediately; +for ``drop``, its ``dropped`` attribute is true. This avoids waiting for a PLC +response to a request that was intentionally not sent. + .. note:: The ``s7`` package is the recommended import for the legacy S7 protocol. diff --git a/s7/__init__.py b/s7/__init__.py index a1fd2a7a..480b44f6 100644 --- a/s7/__init__.py +++ b/s7/__init__.py @@ -30,6 +30,7 @@ "logo", "optimizer", "partner", + "rate_limiter", "s7protocol", "server", "tags", diff --git a/snap7/async_client.py b/snap7/async_client.py index 5b629b6f..1d6825c1 100644 --- a/snap7/async_client.py +++ b/snap7/async_client.py @@ -23,6 +23,7 @@ from .client_base import ClientMixin from .szl import parse_cp_info_szl, parse_cpu_info_szl, parse_order_code_szl, parse_protection_szl from .client import _parse_force_szl +from .rate_limiter import RateLimitAlgorithm, RateLimitBehavior, RequestRateLimiter from .type import ( Area, Block, @@ -308,7 +309,14 @@ class AsyncClient(ClientMixin): MAX_VARS = 20 - def __init__(self) -> None: + def __init__( + self, + *, + max_requests_per_second: float = 0, + rate_limit_algorithm: RateLimitAlgorithm = "fixed", + rate_limit_behavior: RateLimitBehavior = "block", + rate_limit_burst: int | None = None, + ) -> None: self.connection: Optional[AsyncISOTCPConnection] = None self.protocol = S7Protocol() self.connected = False @@ -327,6 +335,12 @@ def __init__(self) -> None: self._last_error = 0 self._lock = asyncio.Lock() + self._rate_limiter = RequestRateLimiter( + max_requests_per_second, + algorithm=rate_limit_algorithm, + behavior=rate_limit_behavior, + burst_capacity=rate_limit_burst, + ) self._params = { Parameter.RemotePort: 102, @@ -346,6 +360,11 @@ def _get_connection(self) -> AsyncISOTCPConnection: raise S7ConnectionError("Not connected to PLC") return self.connection + async def _send_data(self, conn: AsyncISOTCPConnection, request: bytes) -> None: + """Apply the per-client rate limit and send one S7 request PDU.""" + await self._rate_limiter.acquire_async() + await conn.send_data(request) + async def _send_receive(self, request: bytes, max_stale_retries: int = 3) -> dict[str, Any]: """Send a request and receive/parse the response, holding the lock. @@ -365,7 +384,7 @@ async def _send_receive(self, request: bytes, max_stale_retries: int = 3) -> dic expected_seq = struct.unpack(">H", request[4:6])[0] async with self._lock: - await conn.send_data(request) + await self._send_data(conn, request) for attempt in range(max_stale_retries + 1): response_data = await conn.receive_data() @@ -710,7 +729,7 @@ async def list_blocks_of_type(self, block_type: Block, max_count: int) -> List[i async with self._lock: followup = self.protocol.build_userdata_followup_request(group, subfunction, sequence_number) - await conn.send_data(followup) + await self._send_data(conn, followup) response_data = await conn.receive_data() response = self.protocol.parse_response(response_data) @@ -834,7 +853,7 @@ async def download(self, data: bytearray, block_num: int = -1) -> int: ) async with self._lock: - await conn.send_data(header + param_data + data_section) + await self._send_data(conn, header + param_data + data_section) response_data = await conn.receive_data() self.protocol.parse_response(response_data) @@ -851,7 +870,7 @@ async def download(self, data: bytearray, block_num: int = -1) -> int: ) async with self._lock: - await conn.send_data(header + param_data) + await self._send_data(conn, header + param_data) response_data = await conn.receive_data() self.protocol.parse_response(response_data) @@ -1024,7 +1043,7 @@ async def read_szl(self, ssl_id: int, index: int = 0) -> S7SZL: async with self._lock: followup = self.protocol.build_userdata_followup_request(group, subfunction, sequence_number) - await conn.send_data(followup) + await self._send_data(conn, followup) response_data = await conn.receive_data() response = self.protocol.parse_response(response_data) @@ -1210,7 +1229,7 @@ async def iso_exchange_buffer(self, data: bytearray) -> bytearray: conn = self._get_connection() async with self._lock: - await conn.send_data(bytes(data)) + await self._send_data(conn, bytes(data)) response = await conn.receive_data() return bytearray(response) diff --git a/snap7/client.py b/snap7/client.py index 4da50dc9..28fc7b9a 100644 --- a/snap7/client.py +++ b/snap7/client.py @@ -28,6 +28,7 @@ from .client_base import ClientMixin from .log import PLCLoggerAdapter, OperationLogger from .optimizer import ReadItem, ReadPacket, sort_items, merge_items, packetize, extract_results +from .rate_limiter import RateLimitAlgorithm, RateLimitBehavior, RequestRateLimiter from .tags import Tag, _STRING_RE from . import util @@ -296,6 +297,10 @@ def __init__( backoff_factor: float = 2.0, max_delay: float = 30.0, heartbeat_interval: float = 0, + max_requests_per_second: float = 0, + rate_limit_algorithm: RateLimitAlgorithm = "fixed", + rate_limit_behavior: RateLimitBehavior = "block", + rate_limit_burst: int | None = None, on_disconnect: Optional[Callable[[], None]] = None, on_reconnect: Optional[Callable[[], None]] = None, **kwargs: Any, @@ -311,6 +316,10 @@ def __init__( backoff_factor: Multiplier for exponential backoff between retries. max_delay: Maximum delay between reconnection attempts in seconds. heartbeat_interval: Interval in seconds for heartbeat probes (0=disabled). + max_requests_per_second: Maximum outbound PLC requests per second (0=disabled). + rate_limit_algorithm: ``fixed`` for even spacing or ``token_bucket`` for bursts. + rate_limit_behavior: ``block`` to wait, or ``raise``/``drop`` to reject immediately. + rate_limit_burst: Token bucket capacity. Defaults to one second of requests. on_disconnect: Optional callback invoked when connection is lost. on_reconnect: Optional callback invoked after successful reconnection. **kwargs: Ignored. Kept for backwards compatibility. @@ -370,6 +379,12 @@ def __init__( self._max_delay = max_delay self._on_disconnect = on_disconnect self._on_reconnect = on_reconnect + self._rate_limiter = RequestRateLimiter( + max_requests_per_second, + algorithm=rate_limit_algorithm, + behavior=rate_limit_behavior, + burst_capacity=rate_limit_burst, + ) # Heartbeat settings self._heartbeat_interval = heartbeat_interval @@ -402,6 +417,11 @@ def _get_connection(self) -> ISOTCPConnection: raise S7ConnectionError("Not connected to PLC") return self.connection + def _send_data(self, conn: ISOTCPConnection, request: bytes) -> None: + """Apply the per-client rate limit and send one S7 request PDU.""" + self._rate_limiter.acquire() + conn.send_data(request) + def _send_receive(self, request: bytes, max_stale_retries: int = 3) -> dict[str, Any]: """Send a request and receive/parse the response with stale packet retry. @@ -424,7 +444,7 @@ def _send_receive(self, request: bytes, max_stale_retries: int = 3) -> dict[str, conn = self._get_connection() with self._reconnect_lock: - conn.send_data(request) + self._send_data(conn, request) for attempt in range(max_stale_retries + 1): response_data = conn.receive_data() @@ -1222,7 +1242,7 @@ def _send_receive_parallel(self, requests: list[Tuple[int, bytes]]) -> dict[int, # Send all requests back-to-back for _, pdu in requests: - conn.send_data(pdu) + self._send_data(conn, pdu) # Receive responses, matching by sequence number results: dict[int, dict[str, Any]] = {} @@ -1495,7 +1515,7 @@ def list_blocks_of_type(self, block_type: Block, max_count: int) -> List[int]: break followup = self.protocol.build_userdata_followup_request(group, subfunction, sequence_number) - conn.send_data(followup) + self._send_data(conn, followup) response_data = conn.receive_data() response = self.protocol.parse_response(response_data) @@ -1673,7 +1693,7 @@ def download(self, data: bytearray, block_num: int = -1) -> int: len(data_section), # Data length ) - conn.send_data(header + param_data + data_section) + self._send_data(conn, header + param_data + data_section) response_data = conn.receive_data() self.protocol.parse_response(response_data) @@ -1691,7 +1711,7 @@ def download(self, data: bytearray, block_num: int = -1) -> int: 0x0000, # Data length ) - conn.send_data(header + param_data) + self._send_data(conn, header + param_data) response_data = conn.receive_data() self.protocol.parse_response(response_data) @@ -2149,7 +2169,7 @@ def read_szl(self, ssl_id: int, index: int = 0) -> S7SZL: break followup = self.protocol.build_userdata_followup_request(group, subfunction, sequence_number) - conn.send_data(followup) + self._send_data(conn, followup) response_data = conn.receive_data() response = self.protocol.parse_response(response_data) @@ -2269,7 +2289,7 @@ def iso_exchange_buffer(self, data: bytearray) -> bytearray: """ conn = self._get_connection() - conn.send_data(bytes(data)) + self._send_data(conn, bytes(data)) response = conn.receive_data() return bytearray(response) diff --git a/snap7/error.py b/snap7/error.py index 5d71e483..61c61dfb 100644 --- a/snap7/error.py +++ b/snap7/error.py @@ -40,6 +40,14 @@ class S7AuthenticationError(S7Error): pass +class S7RateLimitError(S7Error): + """Raised when a non-blocking request rate limit is reached.""" + + def __init__(self, message: str, *, dropped: bool = False): + super().__init__(message) + self.dropped = dropped + + # S7 client error codes s7_client_errors = { 0x00100000: "errNegotiatingPDU", diff --git a/snap7/rate_limiter.py b/snap7/rate_limiter.py new file mode 100644 index 00000000..b331a73d --- /dev/null +++ b/snap7/rate_limiter.py @@ -0,0 +1,110 @@ +"""Request rate limiting for synchronous and asynchronous S7 clients.""" + +import asyncio +import math +import threading +import time +from collections.abc import Callable +from typing import Literal + +from .error import S7RateLimitError + +RateLimitAlgorithm = Literal["fixed", "token_bucket"] +RateLimitBehavior = Literal["block", "raise", "drop"] + + +class RequestRateLimiter: + """Thread-safe per-client request rate limiter. + + ``fixed`` spaces requests evenly. ``token_bucket`` permits bursts up to + ``burst_capacity`` and then refills continuously at the configured rate. + A rate of zero disables limiting. + """ + + def __init__( + self, + max_requests_per_second: float = 0, + *, + algorithm: RateLimitAlgorithm = "fixed", + behavior: RateLimitBehavior = "block", + burst_capacity: int | None = None, + _clock: Callable[[], float] = time.monotonic, + _sleep: Callable[[float], None] = time.sleep, + ) -> None: + if not math.isfinite(max_requests_per_second) or max_requests_per_second < 0: + raise ValueError("max_requests_per_second must be a finite non-negative number") + if algorithm not in ("fixed", "token_bucket"): + raise ValueError("rate_limit_algorithm must be 'fixed' or 'token_bucket'") + if behavior not in ("block", "raise", "drop"): + raise ValueError("rate_limit_behavior must be 'block', 'raise', or 'drop'") + if burst_capacity is not None and burst_capacity < 1: + raise ValueError("rate_limit_burst must be at least 1") + + self.rate = float(max_requests_per_second) + self.algorithm = algorithm + self.behavior = behavior + self.burst_capacity = burst_capacity or max(1, math.ceil(self.rate)) + self._clock = _clock + self._sleep = _sleep + self._lock = threading.Lock() + + now = self._clock() + self._next_request = now + self._tokens = float(self.burst_capacity) + self._last_refill = now + + @property + def enabled(self) -> bool: + """Whether rate limiting is active.""" + return self.rate > 0 + + def _reserve(self) -> float: + """Reserve one request and return the required delay in seconds.""" + if not self.enabled: + return 0.0 + + with self._lock: + now = self._clock() + if self.algorithm == "fixed": + delay = max(0.0, self._next_request - now) + if delay > 0 and self.behavior != "block": + self._reject() + slot = now + delay + self._next_request = slot + (1.0 / self.rate) + return delay + + # Refill only through the current time. A future _last_refill + # represents tokens already reserved by blocking callers. + if now > self._last_refill: + elapsed = now - self._last_refill + self._tokens = min(float(self.burst_capacity), self._tokens + elapsed * self.rate) + self._last_refill = now + + if self._tokens >= 1.0: + self._tokens -= 1.0 + return 0.0 + + queued_delay = max(0.0, self._last_refill - now) + delay = queued_delay + ((1.0 - self._tokens) / self.rate) + if self.behavior != "block": + self._reject() + self._tokens = 0.0 + self._last_refill = now + delay + return delay + + def _reject(self) -> None: + dropped = self.behavior == "drop" + action = "dropped" if dropped else "rejected" + raise S7RateLimitError(f"Request {action}: rate limit of {self.rate:g} requests/second exceeded", dropped=dropped) + + def acquire(self) -> None: + """Wait for or reserve permission to send one synchronous request.""" + delay = self._reserve() + if delay > 0: + self._sleep(delay) + + async def acquire_async(self) -> None: + """Wait for or reserve permission to send one asynchronous request.""" + delay = self._reserve() + if delay > 0: + await asyncio.sleep(delay) diff --git a/tests/test_rate_limiter.py b/tests/test_rate_limiter.py new file mode 100644 index 00000000..b4ba252a --- /dev/null +++ b/tests/test_rate_limiter.py @@ -0,0 +1,135 @@ +"""Tests for per-client request rate limiting.""" + +from unittest.mock import AsyncMock, Mock + +import pytest + +from snap7.async_client import AsyncClient +from snap7.client import Client +from snap7.error import S7RateLimitError +from snap7.rate_limiter import RequestRateLimiter + + +class FakeTime: + def __init__(self) -> None: + self.now = 0.0 + self.delays: list[float] = [] + + def monotonic(self) -> float: + return self.now + + def sleep(self, delay: float) -> None: + self.delays.append(delay) + self.now += delay + + +def test_fixed_rate_blocks_between_requests() -> None: + fake = FakeTime() + limiter = RequestRateLimiter(2, _clock=fake.monotonic, _sleep=fake.sleep) + + limiter.acquire() + limiter.acquire() + limiter.acquire() + + assert fake.delays == [0.5, 0.5] + + +def test_fixed_rate_can_reject_without_waiting() -> None: + fake = FakeTime() + limiter = RequestRateLimiter(10, behavior="raise", _clock=fake.monotonic) + + limiter.acquire() + with pytest.raises(S7RateLimitError) as exc_info: + limiter.acquire() + + assert not exc_info.value.dropped + + +def test_drop_marks_request_as_dropped() -> None: + fake = FakeTime() + limiter = RequestRateLimiter(10, behavior="drop", _clock=fake.monotonic) + + limiter.acquire() + with pytest.raises(S7RateLimitError) as exc_info: + limiter.acquire() + + assert exc_info.value.dropped + + +def test_token_bucket_allows_configured_burst() -> None: + fake = FakeTime() + limiter = RequestRateLimiter( + 2, + algorithm="token_bucket", + burst_capacity=2, + _clock=fake.monotonic, + _sleep=fake.sleep, + ) + + limiter.acquire() + limiter.acquire() + limiter.acquire() + + assert fake.delays == [0.5] + + +def test_disabled_limiter_never_waits() -> None: + fake = FakeTime() + limiter = RequestRateLimiter(0, _clock=fake.monotonic, _sleep=fake.sleep) + + for _ in range(100): + limiter.acquire() + + assert not fake.delays + + +@pytest.mark.parametrize( + ("kwargs", "message"), + [ + ({"max_requests_per_second": -1}, "finite non-negative"), + ({"algorithm": "sliding"}, "rate_limit_algorithm"), + ({"behavior": "wait"}, "rate_limit_behavior"), + ({"burst_capacity": 0}, "rate_limit_burst"), + ], +) +def test_invalid_configuration_rejected(kwargs: dict[str, object], message: str) -> None: + with pytest.raises(ValueError, match=message): + RequestRateLimiter(**kwargs) # type: ignore[arg-type] + + +def test_sync_client_limits_each_outbound_pdu() -> None: + client = Client(max_requests_per_second=1) + client._rate_limiter.acquire = Mock() + connection = Mock() + + client._send_data(connection, b"request") + + client._rate_limiter.acquire.assert_called_once_with() + connection.send_data.assert_called_once_with(b"request") + + +@pytest.mark.asyncio +async def test_async_limiter_waits_without_blocking(monkeypatch: pytest.MonkeyPatch) -> None: + fake = FakeTime() + limiter = RequestRateLimiter(4, _clock=fake.monotonic) + + async def advance(delay: float) -> None: + fake.sleep(delay) + + monkeypatch.setattr("snap7.rate_limiter.asyncio.sleep", advance) + await limiter.acquire_async() + await limiter.acquire_async() + + assert fake.delays == [0.25] + + +@pytest.mark.asyncio +async def test_async_client_limits_each_outbound_pdu() -> None: + client = AsyncClient(max_requests_per_second=1) + client._rate_limiter.acquire_async = AsyncMock() + connection = AsyncMock() + + await client._send_data(connection, b"request") + + client._rate_limiter.acquire_async.assert_awaited_once_with() + connection.send_data.assert_awaited_once_with(b"request")