import json import threading import unittest import tempfile import secrets import time import urllib.request import urllib.error from unittest.mock import patch from server import Server, HardwareProvider class Tests(unittest.TestCase): def setUp(self): self.state=tempfile.TemporaryDirectory() self.server=Server(0,state_dir=self.state.name) self.password=secrets.token_urlsafe(32) self.api_token=secrets.token_urlsafe(32) record=self.server.credentials.setup(self.password,self.api_token) self.server.sessions['test-session']=dict(expires=time.monotonic()+300,revision=record['password']['hash']) self.cookie='deck_session=test-session' self.thread=threading.Thread(target=self.server.serve_forever) self.thread.start() self.url=f'http://127.0.0.1:{self.server.server_port}' def tearDown(self): self.server.shutdown() self.server.demo.stop() self.server.server_close() self.thread.join() self.state.cleanup() def request(self,path,method='GET',headers=None): req=urllib.request.Request(self.url+'/api/v1/'+path,method=method,headers={"Cookie":self.cookie,**(headers or {})}) with urllib.request.urlopen(req) as r:return json.load(r) def test_lifecycle(self): self.assertEqual(self.request('status')['demo']['state'],'stopped') a=self.request('demo/start','POST',{'X-Athena-Deck':'1'}) self.assertTrue(a['reachable']) self.assertEqual(a['pid'],self.request('demo/start','POST',{'X-Athena-Deck':'1'})['pid']) child=self.server.demo.process self.assertEqual(self.request('demo/stop','POST',{'X-Athena-Deck':'1'})['state'],'stopped') self.assertIsNotNone(child.poll()) self.assertFalse(self.request('demo/stop','POST',{'X-Athena-Deck':'1'})['reachable']) def test_control_guard(self): for headers in ({},{'X-Athena-Deck':'1','Origin':'http://evil.invalid'}): with self.assertRaises(urllib.error.HTTPError) as e:self.request('demo/start','POST',headers) self.assertEqual(e.exception.code,403) e.exception.close() with self.assertRaises(urllib.error.HTTPError) as e:self.request('models/start','POST',{'X-Athena-Deck':'1'}) self.assertEqual(e.exception.code,404) e.exception.close() def test_unavailable_hardware(self): with patch('server.subprocess.run',side_effect=OSError()): result=HardwareProvider().snapshot() self.assertFalse(result['available']) self.assertEqual(result['gpus'],[]) self.assertIsNone(result['sampled_at']) if __name__=='__main__':unittest.main()