Files
2026-08-11 20:24:24 -07:00

120 lines
4.1 KiB
Python

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()