Skip to content

Commit b68a345

Browse files
authored
Refactor ReplayWindow for improved type safety
1 parent 54c2f42 commit b68a345

1 file changed

Lines changed: 57 additions & 5 deletions

File tree

‎src/cyphersyntax/replay.py‎

Lines changed: 57 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,23 +1,75 @@
11
from __future__ import annotations
22

3+
from collections.abc import Callable
34
from dataclasses import dataclass, field
5+
from threading import RLock
6+
from types import TracebackType
7+
from typing import Protocol, TypeVar
48

59
from .errors import ReplayDetectedError
10+
from .protocol import validate_message_sequence
11+
12+
13+
_Result = TypeVar("_Result")
14+
15+
16+
class _ContextLock(Protocol):
17+
def __enter__(self) -> object: ...
18+
19+
def __exit__(
20+
self,
21+
exc_type: type[BaseException] | None,
22+
exc_value: BaseException | None,
23+
traceback: TracebackType | None,
24+
) -> bool | None: ...
625

726

827
@dataclass(slots=True)
928
class ReplayWindow:
1029
window_size: int = 1024
11-
highest_seen: int = -1
12-
seen: set[int] = field(default_factory=set)
30+
highest_seen: int = field(default=-1, init=False)
31+
seen: set[int] = field(default_factory=set, init=False)
32+
_lock: _ContextLock = field(
33+
default_factory=RLock,
34+
init=False,
35+
repr=False,
36+
compare=False,
37+
)
1338

14-
def observe(self, sequence: int) -> None:
39+
def __post_init__(self) -> None:
40+
if isinstance(self.window_size, bool) or not isinstance(self.window_size, int):
41+
raise TypeError("replay window size must be an integer")
42+
if self.window_size <= 0:
43+
raise ValueError("replay window size must be positive")
44+
45+
def _check_unlocked(self, sequence: int) -> None:
46+
validate_message_sequence(sequence)
1547
if sequence in self.seen:
1648
raise ReplayDetectedError(f"replayed sequence number: {sequence}")
1749
if sequence <= self.highest_seen - self.window_size:
18-
raise ReplayDetectedError(f"stale sequence number outside replay window: {sequence}")
50+
raise ReplayDetectedError(
51+
f"stale sequence number outside replay window: {sequence}"
52+
)
53+
54+
def _record_unlocked(self, sequence: int) -> None:
1955
self.seen.add(sequence)
2056
if sequence > self.highest_seen:
2157
self.highest_seen = sequence
2258
floor = self.highest_seen - self.window_size
23-
self.seen = {n for n in self.seen if n > floor}
59+
self.seen = {number for number in self.seen if number > floor}
60+
61+
def observe(self, sequence: int) -> None:
62+
with self._lock:
63+
self._check_unlocked(sequence)
64+
self._record_unlocked(sequence)
65+
66+
def authenticate_and_record(
67+
self,
68+
sequence: int,
69+
operation: Callable[[], _Result],
70+
) -> _Result:
71+
with self._lock:
72+
self._check_unlocked(sequence)
73+
result = operation()
74+
self._record_unlocked(sequence)
75+
return result

0 commit comments

Comments
 (0)