Files
Athena-Deck/test_server.py
T

71 lines
3.8 KiB
Python

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_profile_and_download_routes(self):
from pathlib import Path
target=Path(self.state.name)/'models'/('a'*64);target.mkdir(parents=True)
(target/'model.gguf').write_bytes(b'GGUFtest')
(target/'entry.json').write_text(json.dumps(dict(repo='test/image',file='model.gguf',size=8,kind='image',revision='b'*40)))
payload=dict(id=None,revision=0,name='image-test',kind='image',model_id='a'*64,parameters=dict(width=512,height=512,steps=10,seed=-1,guidance=1))
req=urllib.request.Request(self.url+'/api/v1/profiles/save',data=json.dumps(payload).encode(),headers={'Cookie':self.cookie,'X-Athena-Deck':'1','Content-Type':'application/json'})
with urllib.request.urlopen(req) as r:self.assertEqual(r.status,200)
self.assertEqual(self.request('profiles')['profiles'][0]['name'],'image-test')
with self.assertRaises(urllib.error.HTTPError) as error:
urllib.request.urlopen(urllib.request.Request(self.url+'/api/v1/profiles',headers={'Authorization':'Bearer '+self.api_token}))
self.assertEqual(error.exception.code,401);error.exception.close()
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.dict('os.environ', {'DECK_LOCAL_HARDWARE':'0'}), 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()