| 1 | #!/usr/bin/env python3 |
| 2 | import asyncio |
| 3 | import base64 |
| 4 | import importlib.util |
| 5 | import json |
| 6 | from pathlib import Path |
| 7 | import random |
| 8 | import struct |
| 9 | import threading |
| 10 | import unittest |
| 11 | import zlib |
| 12 | from types import SimpleNamespace |
| 13 | from unittest.mock import Mock |
| 14 | |
| 15 | |
| 16 | spec = importlib.util.spec_from_file_location("vm_screen", Path(__file__).with_name("vm-screen.py")) |
| 17 | screen = importlib.util.module_from_spec(spec) |
| 18 | spec.loader.exec_module(screen) |
| 19 | |
| 20 | |
| 21 | class Writer: |
| 22 | def __init__(self): |
| 23 | self.data = bytearray() |
| 24 | self.closed = False |
| 25 | self.transport = self |
| 26 | |
| 27 | def write(self, data): |
| 28 | self.data.extend(data) |
| 29 | |
| 30 | async def drain(self): |
| 31 | pass |
| 32 | |
| 33 | def get_write_buffer_size(self): |
| 34 | return len(self.data) |
| 35 | |
| 36 | def is_closing(self): |
| 37 | return self.closed |
| 38 | |
| 39 | def close(self): |
| 40 | self.closed = True |
| 41 | |
| 42 | |
| 43 | class DesktopTests(unittest.IsolatedAsyncioTestCase): |
| 44 | def setUp(self): |
| 45 | self.reader, self.writer = asyncio.StreamReader(), Writer() |
| 46 | self.messages = [] |
| 47 | self.desktop = screen.Desktop(self.reader, self.writer, self.messages.append) |
| 48 | self.desktop.resize(4, 2) |
| 49 | |
| 50 | async def update(self, rectangles): |
| 51 | self.reader.feed_data(struct.pack("!BBH", 0, 0, len(rectangles)) + b"".join(rectangles)) |
| 52 | self.reader.feed_eof() |
| 53 | self.presented = self.resized = 0 |
| 54 | |
| 55 | def present(): |
| 56 | self.presented += 1 |
| 57 | |
| 58 | def resize(): |
| 59 | self.resized += 1 |
| 60 | |
| 61 | with self.assertRaises(asyncio.IncompleteReadError): |
| 62 | await self.desktop.read(present, resize) |
| 63 | |
| 64 | async def test_handshake(self): |
| 65 | self.reader.feed_data(b"RFB 003.008\n\x01\x01\0\0\0\0" + struct.pack("!HH16sI", 4, 2, bytes(16), 4) + b"test") |
| 66 | await self.desktop.start() |
| 67 | self.assertEqual(self.writer.data[:14], b"RFB 003.008\n\x01\x01") |
| 68 | self.assertEqual(self.writer.data[-10:], struct.pack("!BBHHHH", 3, 0, 0, 0, 4, 2)) |
| 69 | |
| 70 | async def test_raw_and_overlapping_copy_present_once(self): |
| 71 | pixels = b"".join(bytes([value, 0, 0, 0]) for value in range(1, 9)) |
| 72 | await self.update([ |
| 73 | struct.pack("!HHHHi", 0, 0, 4, 2, 0) + pixels, |
| 74 | struct.pack("!HHHHiHH", 1, 0, 3, 2, 1, 0, 0), |
| 75 | ]) |
| 76 | self.assertEqual(list(self.desktop.pixels[::4]), [1, 1, 2, 3, 5, 5, 6, 7]) |
| 77 | self.assertEqual(self.presented, 1) |
| 78 | |
| 79 | async def test_idle_and_cursor_do_not_encode(self): |
| 80 | await self.update([struct.pack("!HHHHi", 0, 0, 2, 1, -239) + bytes([1, 2, 3, 0, 4, 5, 6, 0, 128])]) |
| 81 | self.assertEqual(self.presented, 0) |
| 82 | encoded = base64.b64decode(self.messages[0]["image"].split(",")[1]) |
| 83 | length = struct.unpack("!I", encoded[33:37])[0] |
| 84 | self.assertEqual(zlib.decompress(encoded[41:41 + length]), bytes([0, 3, 2, 1, 255, 6, 5, 4, 0])) |
| 85 | |
| 86 | async def test_resize_and_out_of_bounds(self): |
| 87 | await self.update([struct.pack("!HHHHi", 0, 0, 2, 2, -223)]) |
| 88 | self.assertEqual((self.desktop.width, self.desktop.height, self.resized), (2, 2, 1)) |
| 89 | self.reader = asyncio.StreamReader() |
| 90 | self.desktop.reader = self.reader |
| 91 | with self.assertRaisesRegex(ValueError, "outside"): |
| 92 | await self.update([struct.pack("!HHHHi", 1, 0, 2, 1, 0)]) |
| 93 | for width, height in [(0, 2), (4097, 1), (65535, 65535)]: |
| 94 | with self.assertRaises(ValueError): |
| 95 | self.desktop.resize(width, height) |
| 96 | |
| 97 | def test_input_validation_and_release(self): |
| 98 | self.desktop.input({"type": "key", "key": 65507, "down": True}) |
| 99 | self.desktop.input({"type": "pointer", "x": -1, "y": 10, "buttons": 1}) |
| 100 | self.desktop.input({"type": "release"}) |
| 101 | self.assertFalse(self.desktop.keys) |
| 102 | self.assertEqual(self.writer.data[-6:], struct.pack("!BBHH", 5, 0, 0, 1)) |
| 103 | for message in [{"type": "key", "key": True, "down": True}, {"type": "pointer", "x": 1, "y": 1, "buttons": 256}]: |
| 104 | with self.assertRaises(ValueError): |
| 105 | self.desktop.input(message) |
| 106 | |
| 107 | |
| 108 | class GuestTests(unittest.IsolatedAsyncioTestCase): |
| 109 | def setUp(self): |
| 110 | self.guest = screen.GuestDesktop("fixture") |
| 111 | self.messages = [] |
| 112 | self.presented = self.resized = 0 |
| 113 | self.viewer = Mock(desktop=SimpleNamespace(cursor={"type": "cursor", "source": "qemu"}), |
| 114 | writer=Writer(), send=self.messages.append, resize=self.resize, present=self.present) |
| 115 | self.guest.screens = {self.viewer} |
| 116 | |
| 117 | def present(self): |
| 118 | self.presented += 1 |
| 119 | |
| 120 | def resize(self): |
| 121 | self.resized += 1 |
| 122 | |
| 123 | async def packets(self, *packets): |
| 124 | reader = asyncio.StreamReader() |
| 125 | for name, words, data in packets: |
| 126 | kind = screen.GUEST_PROTOCOL["packets"][name]["id"] |
| 127 | payload = struct.pack("!" + "I" * len(words), *words) + data |
| 128 | reader.feed_data(struct.pack("!4sII", b"SGV1", kind, len(payload)) + payload) |
| 129 | reader.feed_eof() |
| 130 | with self.assertRaises(asyncio.IncompleteReadError): |
| 131 | await self.guest.read(reader) |
| 132 | |
| 133 | async def test_dirty_rectangles_and_cursor_do_not_reencode_video(self): |
| 134 | await self.packets(("screen", [4, 2, 0, 0, 4, 2], bytes(32)), |
| 135 | ("screen", [4, 2, 2, 1, 1, 1], b"\x01\x02\x03\0"), |
| 136 | ("cursor", [7, 1, 1, 0, 0, 2], struct.pack("!I", 17) + bytes([0,0,0,0,7]) + struct.pack("!I", 33) + bytes([1,2,3,255,0])), |
| 137 | ("pointer", [7, 0xffffffff, 9, 1], b"")) |
| 138 | self.assertEqual(self.guest.pixels[24:28], b"\x01\x02\x03\0") |
| 139 | self.assertEqual((self.presented, self.resized), (2, 1)) |
| 140 | self.assertEqual([message["duration"] for message in self.messages if message["type"] == "cursor"], [17, 33]) |
| 141 | self.assertIn("invert", self.messages[0]) |
| 142 | self.assertEqual(self.guest.pointer["x"], -1) |
| 143 | |
| 144 | async def test_suspend_restores_fallback_and_requires_full_refresh(self): |
| 145 | await self.packets(("screen", [2, 2, 0, 0, 2, 2], bytes(16)), ("suspend", [], b"")) |
| 146 | self.assertFalse(self.guest.pixels) |
| 147 | self.assertEqual(self.messages[-1]["source"], "qemu") |
| 148 | with self.assertRaisesRegex(ValueError, "complete"): |
| 149 | await self.packets(("screen", [2, 2, 0, 0, 1, 1], bytes(4))) |
| 150 | self.viewer.desktop.cursor = None |
| 151 | await self.packets(("screen", [1, 1, 0, 0, 1, 1], bytes(4)), ("suspend", [], b"")) |
| 152 | self.assertTrue(self.messages[-1]["fallback"]) |
| 153 | |
| 154 | async def test_untrusted_guest_packets_are_bounded(self): |
| 155 | for packet in [("screen", [4097, 1, 0, 0, 1, 1], bytes(4)), |
| 156 | ("screen", [2, 2, 1, 1, 2, 2], bytes(16)), |
| 157 | ("cursor", [1, 1, 1, 1, 0, 1], struct.pack("!I", 100) + bytes(5)), |
| 158 | ("cursor", [1, 1, 1, 0, 0, 1], struct.pack("!I", 0) + bytes(5)), |
| 159 | ("pointer", [1, 0, 0, 2], b""), ("suspend", [], b"x")]: |
| 160 | with self.subTest(packet=packet), self.assertRaises(ValueError): |
| 161 | await self.packets(packet) |
| 162 | reader = asyncio.StreamReader() |
| 163 | reader.feed_data(struct.pack("!4sII", b"SGV1", 1, 0xffffffff)) |
| 164 | with self.assertRaises(ValueError): |
| 165 | await self.guest.read(reader) |
| 166 | |
| 167 | async def test_viewer_close_during_cursor_fanout(self): |
| 168 | await self.packets(("screen", [1, 1, 0, 0, 1, 1], bytes(4))) |
| 169 | self.viewer.send = lambda _: self.guest.screens.discard(self.viewer) |
| 170 | await self.packets(("cursor", [1, 1, 1, 0, 0, 2], (struct.pack("!I", 100) + bytes(5)) * 2)) |
| 171 | self.assertFalse(self.guest.screens) |
| 172 | self.assertEqual(len(self.guest.cursor), 2) |
| 173 | |
| 174 | async def test_missing_heartbeat_times_out_but_does_not_encode(self): |
| 175 | await self.packets(("screen", [1, 1, 0, 0, 1, 1], bytes(4)), ("heartbeat", [], b"")) |
| 176 | self.assertEqual(self.presented, 1) |
| 177 | reader = asyncio.StreamReader() |
| 178 | with self.assertRaises(TimeoutError): |
| 179 | await self.guest.read(reader) |
| 180 | |
| 181 | |
| 182 | class StreamTests(unittest.IsolatedAsyncioTestCase): |
| 183 | async def asyncSetUp(self): |
| 184 | self.writer = Writer() |
| 185 | desktop = screen.Desktop(asyncio.StreamReader(), Writer(), lambda _: None) |
| 186 | desktop.resize(320, 200) |
| 187 | self.stream = screen.Screen(self.writer, desktop) |
| 188 | |
| 189 | async def asyncTearDown(self): |
| 190 | self.writer.close() |
| 191 | self.stream.close() |
| 192 | |
| 193 | async def test_browser_payload_and_sending_direction(self): |
| 194 | sdp = "\r\n".join([ |
| 195 | "v=0", "o=- 1 1 IN IP4 127.0.0.1", "s=-", "t=0 0", "a=group:BUNDLE 0", |
| 196 | "m=video 9 UDP/TLS/RTP/SAVPF 96 103", "c=IN IP4 0.0.0.0", "a=mid:0", "a=recvonly", |
| 197 | "a=rtcp-mux", "a=ice-ufrag:abcd", "a=ice-pwd:abcdefghijklmnopqrstuvwxyz", |
| 198 | "a=setup:actpass", "a=fingerprint:sha-256 " + ":".join(["00"] * 32), |
| 199 | "a=rtpmap:96 VP8/90000", "a=rtpmap:103 H264/90000", |
| 200 | "a=fmtp:103 packetization-mode=1;profile-level-id=42e01f;level-asymmetry-allowed=1", "", |
| 201 | ]) |
| 202 | self.stream.signal({"type": "offer", "sdp": sdp}) |
| 203 | async with asyncio.timeout(5): |
| 204 | while b'"type":"answer"' not in self.writer.data: |
| 205 | await asyncio.sleep(0.01) |
| 206 | answer = next(json.loads(line) for line in self.writer.data.splitlines() if json.loads(line)["type"] == "answer") |
| 207 | self.assertIn("a=sendonly", answer["sdp"]) |
| 208 | self.assertIn("a=rtpmap:103 H264/90000", answer["sdp"]) |
| 209 | self.assertEqual(self.stream.pipeline.get_by_name("pay").get_property("pt"), 103) |
| 210 | |
| 211 | async def test_stalled_guest_closes_and_ignores_queued_input(self): |
| 212 | message = json.dumps({"type": "clipboard", "text": "a" * 65536}) |
| 213 | for _ in range(64): |
| 214 | self.stream.input(message) |
| 215 | self.assertTrue(self.writer.closed) |
| 216 | self.assertLess(len(self.stream.desktop.writer.data), 327680) |
| 217 | |
| 218 | async def test_rtp_packets_fit_tunneled_network(self): |
| 219 | Gst = screen.Gst |
| 220 | pipeline = self.stream.pipeline |
| 221 | pipeline.set_state(Gst.State.NULL) |
| 222 | pipeline.remove(self.stream.rtc) |
| 223 | sink = Gst.ElementFactory.make("appsink") |
| 224 | sink.set_property("sync", False) |
| 225 | pipeline.add(sink) |
| 226 | pipeline.get_by_name("codec").link(sink) |
| 227 | pipeline.set_state(Gst.State.PLAYING) |
| 228 | pixels = random.Random(0).randbytes(len(self.stream.desktop.pixels)) |
| 229 | self.stream.frames.emit("push-buffer", Gst.Buffer.new_wrapped(pixels)) |
| 230 | packets = [] |
| 231 | async with asyncio.timeout(5): |
| 232 | while True: |
| 233 | sample = sink.emit("try-pull-sample", 0) |
| 234 | if sample: |
| 235 | buffer = sample.get_buffer() |
| 236 | packets.append(buffer.get_size()) |
| 237 | if buffer.extract_dup(1, 1)[0] & 0x80: |
| 238 | break |
| 239 | else: |
| 240 | await asyncio.sleep(0.01) |
| 241 | self.assertGreater(len(packets), 1) |
| 242 | self.assertLessEqual(max(packets), 1200) |
| 243 | |
| 244 | |
| 245 | if __name__ == "__main__": |
| 246 | screen.Gst.init(None) |
| 247 | threading.Thread(target=screen.GLib.MainLoop().run, daemon=True).start() |
| 248 | unittest.main() |