120 lines
4.1 KiB
Python
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()
|