Files
Athena-Deck/test_ltx_original.py
T

86 lines
5.2 KiB
Python

import json
import tempfile
import unittest
from pathlib import Path
from unittest.mock import Mock,patch
from ltx_original_runtime import LTXOriginalRuntime,DESKTOP_REV
from video_original import VideoRuntimes
from ltx_original_entry import install_device_policy
from test_video_comfy import VideoComfyTests
class OriginalSelectionTests(VideoComfyTests):
def setUp(self):
super().setUp();self.original=Mock();self.original.status.return_value={'installed':True}
self.video=VideoRuntimes(self.scheduler,Mock(),self.video.helper,self.root/'video',self.catalog,self.profiles,self.runtime,original_runtime=self.original)
def test_missing_comfy_does_not_block_installed_original(self):
self.runtime.paths.return_value=(self.root/'missing',self.root/'missing-comfy')
self.video.select_runtime('original');self.video.select(self.model['id'])
self.assertTrue(self.video.models()[0]['runnable']);self.assertEqual(self.video.status()['runtime_kind'],'original')
saved=json.loads(self.video.path.read_text());self.assertEqual(saved['runtime'],'original');self.assertEqual(saved['model_id'],self.model['id'])
def test_missing_original_is_explicit(self):
self.original.status.return_value={'installed':False};self.video.select_runtime('original')
self.assertFalse(self.video.models()[0]['runnable']);self.assertIn('LTX Original fehlt',self.video.models()[0]['blockers'][-1])
def test_no_runtime_switch_while_video_running(self):
self.scheduler.gpu_mode='video'
with self.assertRaises(ValueError):self.video.select_runtime('original')
self.assertEqual(self.video.runtime_kind(),'comfy')
def test_reject_unknown_runtime(self):
with self.assertRaises(ValueError):self.video.select_runtime('../backend')
class InstallerTests(unittest.TestCase):
def test_no_implicit_install_and_no_external_active_path(self):
with tempfile.TemporaryDirectory() as directory:
root=Path(directory);image=Mock();image.paths.return_value=(root/'missing',root/'comfy')
runtime=LTXOriginalRuntime(root,image);self.assertFalse(runtime.status()['installed'])
(root/'active.json').write_text(json.dumps({'id':'../../external'}));self.assertFalse(runtime.status()['installed'])
self.assertEqual(runtime.status()['required_disk_gib'],25)
def test_existing_ready_install_is_reused(self):
with tempfile.TemporaryDirectory() as directory:
root=Path(directory);base=root/('a'*32);(base/'python/bin').mkdir(parents=True);(base/'desktop/backend').mkdir(parents=True)
(base/'python/bin/python').touch();(base/'desktop/backend/ltx2_server.py').touch();(base/'ready').write_text(DESKTOP_REV)
(root/'active.json').write_text(json.dumps({'id':base.name}));image=Mock();image.paths.return_value=(root/'missing',root/'comfy')
runtime=LTXOriginalRuntime(root,image);self.assertTrue(runtime.status()['installed'])
with self.assertRaises(ValueError):runtime.start()
class DevicePolicyTests(unittest.TestCase):
def test_encoder_gpu_and_conditioning_transfer_preserve_native_contract(self):
from collections import namedtuple
Row=namedtuple('Row','video audio mask')
class Tensor:
def __init__(self,device):self.device=device
def to(self,device):return Tensor(device)
class Encoder:
def __init__(self,model_paths,dtype,device,offload_mode='none'):self.device=device;self.offload=offload_mode
def __call__(self,prompts,**kwargs):return [Row(Tensor(self.device),None,Tensor(self.device))]
torch=Mock();torch.device.side_effect=lambda x:x
modes=Mock();modes.DISK='disk'
install_device_policy(torch,Encoder,modes)
encoder=Encoder('paths','bf16','cuda:0');self.assertEqual(encoder.device,'cuda:1');self.assertEqual(encoder.offload,'disk')
result=encoder(['synthetic prompt'])[0];self.assertEqual(result.video.device,'cuda:0');self.assertIsNone(result.audio);self.assertEqual(result.mask.device,'cuda:0')
class ProxyContractTests(unittest.TestCase):
def test_native_paths_and_errors_are_forwarded_without_schema_adapter(self):
import http.client
import threading
from http.server import BaseHTTPRequestHandler,ThreadingHTTPServer
from video_comfy_proxy import relay
class Native(BaseHTTPRequestHandler):
def log_message(self,*args):pass
def do_POST(self):
raw=self.rfile.read(int(self.headers.get('Content-Length',0)))
self.server.received=(self.path,raw)
self.send_response(422);self.send_header('Content-Type','application/json');self.end_headers();self.wfile.write(b'{"detail":"native validation"}')
native=ThreadingHTTPServer(('127.0.0.1',0),Native)
class Forward(BaseHTTPRequestHandler):
def log_message(self,*args):pass
def do_POST(self):relay(self,native.server_port)
forward=ThreadingHTTPServer(('127.0.0.1',0),Forward)
for server in (native,forward):threading.Thread(target=server.serve_forever,daemon=True).start()
try:
client=http.client.HTTPConnection('127.0.0.1',forward.server_port);body=b'{"prompt":"synthetic","imagePath":"/synthetic/input.jpg"}'
client.request('POST','/api/generate?contract=1',body,{'Content-Type':'application/json'})
response=client.getresponse();self.assertEqual(response.status,422);self.assertEqual(response.read(),b'{"detail":"native validation"}');client.close()
self.assertEqual(native.received,('/api/generate?contract=1',body))
finally:
for server in (forward,native):server.shutdown();server.server_close()