1#!/usr/bin/env python3
2import asyncio
3import base64
4import importlib.util
5import json
6from pathlib import Path
7import random
8import struct
9import threading
10import unittest
11import zlib
12from types import SimpleNamespace
13from unittest.mock import Mock
14
15
16spec = importlib.util.spec_from_file_location("vm_screen", Path(__file__).with_name("vm-screen.py"))
17screen = importlib.util.module_from_spec(spec)
18spec.loader.exec_module(screen)
19
20
21class 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
43class 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
108class 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
182class 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
245if __name__ == "__main__":
246 screen.Gst.init(None)
247 threading.Thread(target=screen.GLib.MainLoop().run, daemon=True).start()
248 unittest.main()