Skip to content
Merged
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
34 changes: 17 additions & 17 deletions bbot/test/test_step_2/module_tests/test_module_zeromq.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,32 +15,32 @@ class TestZeroMQ(ModuleTestBase):
}

async def setup_before_prep(self, module_test):
# Setup ZeroMQ context and socket
self.context = zmq.asyncio.Context()
self.socket = self.context.socket(zmq.SUB)
self.socket.connect("tcp://localhost:5555")
self.socket.setsockopt_string(zmq.SUBSCRIBE, "")

async def check(self, module_test, events):
try:
events_json = [e.json() for e in events]
events_json.sort(key=lambda x: x["timestamp"])

# Collect events from ZeroMQ
zmq_events = []
while len(zmq_events) < len(events_json):
msg = await self.socket.recv()
event_data = json.loads(msg.decode("utf-8"))
zmq_events.append(event_data)
async def setup_after_prep(self, module_test):
self.zmq_events = []
zeromq_module = module_test.scan.modules["zeromq"]
original_send = zeromq_module.socket.send

zmq_events.sort(key=lambda x: x["timestamp"])
async def capturing_send(data, *args, **kwargs):
self.zmq_events.append(json.loads(data.decode("utf-8")))
return await original_send(data, *args, **kwargs)

assert len(events_json) == len(zmq_events), "Number of events does not match"
zeromq_module.socket.send = capturing_send

# Verify the events match
assert events_json == zmq_events, "Events do not match"
def check(self, module_test, events):
try:
events_json = [e.json() for e in events]
events_json.sort(key=lambda x: x["timestamp"])
self.zmq_events.sort(key=lambda x: x["timestamp"])

assert len(events_json) == len(self.zmq_events), (
f"Event count mismatch: expected {len(events_json)}, got {len(self.zmq_events)}"
)
assert events_json == self.zmq_events, "Events do not match"
finally:
# Clean up: Close the ZeroMQ socket
self.socket.close()
self.context.term()