import base64 import hashlib import json import os import sys import unittest sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', 'payload', 'user', 'remote_access', 'pager-webui')) import server def setUpModule(): __import__('importlib').reload(server) class FrameCodecTest(unittest.TestCase): def test_encode_small_text_frame(self): raw = server.ws_encode(b'hi', opcode=0x1) self.assertEqual(raw[0] & 0x80, 0x80) self.assertEqual(raw[0] & 0x0F, 0x1) self.assertEqual(raw[1], 2) def test_decode_unmasked_text(self): frame = server.ws_encode(b'hello', opcode=0x1) opcode, payload, consumed = server.ws_decode_frame(frame) self.assertEqual(opcode, 0x1) self.assertEqual(payload, b'hello') self.assertEqual(consumed, len(frame)) def test_decode_masked_text(self): import struct mask = b'\x01\x02\x03\x04' payload = b'abc' masked = bytes(payload[i] ^ mask[i % 4] for i in range(len(payload))) frame = bytes([0x81, 0x80 | len(payload)]) + mask + masked opcode, out, consumed = server.ws_decode_frame(frame) self.assertEqual(out, b'abc') def test_handshake_reply_uses_sha1(self): key = 'dGhlIHNhbXBsZSBub25jZQ==' reply = server.ws_handshake_reply(key).decode() self.assertIn('101 Switching Protocols', reply) self.assertIn('s3pPLMBiTxaQ9kYGzzhZRbK+xOo=', reply) def test_encode_masked_roundtrip(self): raw = server.ws_encode(b'hello', opcode=0x1, mask=True) self.assertTrue(raw[1] & 0x80) self.assertEqual(raw[1] & 0x7F, 5) opcode, payload, consumed = server.ws_decode_frame(raw) self.assertEqual(opcode, 0x1) self.assertEqual(payload, b'hello') self.assertEqual(consumed, len(raw)) class RelayDrainTest(unittest.TestCase): def test_complete_frame_plus_head_of_next(self): frame1 = server.ws_encode(b'one') head2 = server.ws_encode(b'two')[:3] remaining, frames, closed = server._relay_drain(b'', frame1 + head2) self.assertEqual(frames, [(0x1, b'one')]) self.assertFalse(closed) self.assertEqual(remaining, head2) def test_partial_frame_keeps_remainder(self): head = server.ws_encode(b'hello')[:4] remaining, frames, closed = server._relay_drain(b'', head) self.assertEqual(frames, []) self.assertFalse(closed) self.assertEqual(remaining, head) def test_split_frame_across_chunks(self): frame = server.ws_encode(b'abcdef') first, second = frame[:3], frame[3:] remaining, frames, closed = server._relay_drain(b'', first) self.assertEqual(frames, []) self.assertEqual(remaining, first) remaining, frames, closed = server._relay_drain(remaining, second) self.assertEqual(frames, [(0x1, b'abcdef')]) self.assertFalse(closed) self.assertEqual(remaining, b'') def test_binary_frame_collected_with_opcode(self): frame = server.ws_encode(b'\x00\x01\x02\x03', opcode=0x2) remaining, frames, closed = server._relay_drain(b'', frame) self.assertEqual(frames, [(0x2, b'\x00\x01\x02\x03')]) self.assertFalse(closed) self.assertEqual(remaining, b'') def test_close_frame_reported(self): close = bytes([0x88, 0x02]) + b'\x03\xe8' remaining, frames, closed = server._relay_drain(b'', close) self.assertTrue(closed) self.assertEqual(frames, []) self.assertEqual(remaining, b'') class WsPoolTest(unittest.TestCase): def test_broadcast_skips_dead(self): class FakeSock: def __init__(self, fail=False): self.fail = fail self.sent = [] def sendall(self, b): if self.fail: raise OSError('closed') self.sent.append(b) pool = server.WSPool() good = FakeSock() dead = FakeSock(fail=True) pool.add(good) pool.add(dead) pool.broadcast(b'x') self.assertEqual(len(good.sent), 1) self.assertNotIn(dead, pool.clients) if __name__ == '__main__': unittest.main()