Files
SystemSimulationApp/tests/test_native_result_transport.py
T

146 lines
8.7 KiB
Python

"""The HTTP fast path must preserve real native values and task semantics."""
from dataclasses import replace
import json
from pathlib import Path
import tempfile
from types import SimpleNamespace
import unittest
from unittest.mock import patch
from uuid import uuid4
import asyncio
from app.main import app, _register_simulation_task, _request_simulation_task_cancel, simulation_event_stream
from app.simulation.backends import simulation_config
from app.simulation.native_codegen.build import build_native
from app.simulation.native_codegen.compiler import compile_native_program
from app.simulation.native_codegen.input import load_input
from app.simulation.native_codegen.runner import execute_native
from app.simulation.native_codegen.transport import NativeSeriesJson, read_indexed_result, serialize_result_parts
from app.main import compile_system_xml_network
class AsgiClient:
"""Exercise real routing/response bodies without an optional HTTP client dependency."""
def __init__(self, application): self.application = application
def post(self, path, *, content=b'', headers=None): return self.request('POST', path, content, headers)
def get(self, path): return self.request('GET', path, b'', None)
def request(self, method, path, content, headers):
async def run():
messages = []
scope = {'type':'http', 'asgi':{'version':'3.0','spec_version':'2.4'},
'http_version':'1.1','method':method,'scheme':'http','path':path,
'raw_path':path.encode(),'query_string':b'', 'root_path':'',
'headers':[(k.lower().encode(),v.encode()) for k,v in (headers or {}).items()],
'server':('testserver',80),'client':('127.0.0.1',1234)}
async def receive(): return {'type':'http.request','body':content,'more_body':False}
async def send(message): messages.append(message)
await self.application(scope,receive,send)
status = next(m['status'] for m in messages if m['type']=='http.response.start')
body = b''.join(m.get('body',b'') for m in messages if m['type']=='http.response.body')
response_headers = next(m.get("headers", []) for m in messages if m["type"] == "http.response.start")
return SimpleNamespace(status_code=status,content=body,json=lambda:json.loads(body),
headers={k.decode():v.decode() for k,v in response_headers})
return asyncio.run(run())
class NativeResultTransportTests(unittest.TestCase):
@classmethod
def setUpClass(cls):
cls.temp = tempfile.TemporaryDirectory(prefix='test-native-transport-')
cls.root = Path(cls.temp.name)
cls.xml = Path('tests/fixtures/native-skill-test.xml').read_bytes().replace(b'tStop="10"', b'tStop="0.1"')
source = cls.root/'input.xml'; source.write_bytes(cls.xml)
_, document = load_input(source)
cls.config = simulation_config(document.simulation)
cls.build = build_native(compile_native_program(compile_system_xml_network(document)))
cls.normal = execute_native(cls.build, cls.config, .001, run_dir=cls.root/'ordinary')
cls.indexed = execute_native(cls.build, cls.config, .001, run_dir=cls.root/'indexed', raw_series=True)
@classmethod
def tearDownClass(cls):
cls.temp.cleanup()
def test_real_native_series_are_equal_without_large_python_parse(self):
self.assertIsInstance(self.indexed['series'], NativeSeriesJson)
self.assertEqual(self.indexed['series'].materialize(), self.normal['series'])
self.assertEqual(self.indexed['series'].sample_count, len(self.normal['series']['time']))
for key in ('final', 'finalState', 'nfev', 'acceptedSteps', 'rejectedSteps'):
self.assertEqual(self.indexed[key], self.normal[key])
self.assertGreater(len(self.indexed['series'].data), 100000)
original = json.loads
sizes = []
def small_only(value):
sizes.append(len(value))
self.assertLess(len(value), 50000, 'Raw series was decoded through Python')
return original(value)
with patch('app.simulation.native_codegen.transport.json', SimpleNamespace(loads=small_only)):
payload = read_indexed_result(self.root/'indexed/result.json', self.root/'indexed/result-index.json')
self.assertEqual(len(sizes), 2) # index and small metadata only
self.assertEqual(payload['series'].data, self.indexed['series'].data)
def test_real_http_stream_and_retained_task_get_keep_same_schema(self):
client = AsgiClient(app)
ident = 'transport-'+uuid4().hex
response = client.post('/api/system-xml/simulate-stream', content=self.xml, headers={'X-Simulation-Id':ident})
self.assertEqual(response.status_code, 200)
events = [json.loads(line) for line in response.content.splitlines()]
result = next(event['result'] for event in events if event['event'] == 'result')
self.assertTrue(result['success'])
self.assertEqual(result['simulatedUntil'], .1)
self.assertEqual(result['diagnostics']['sampleCount'], len(result['series']['time']))
# The bytes survive deletion of the worker directory and repeated task reads.
for _ in range(2):
retained = client.get('/api/system-xml/simulations/'+ident)
self.assertEqual(retained.status_code, 200)
self.assertEqual(retained.json()['result'], result)
synchronous = client.post('/api/system-xml/simulate', content=self.xml).json()
for key in ('series', 'final', 'variables', 'model', 'simulation'):
self.assertEqual(synchronous[key], result[key])
def test_cancelled_raw_stream_keeps_partial_result_and_public_status(self):
for reason, status in [('user', 'stopped'), ('stalled', 'stalled')]:
task = _register_simulation_task('transport-'+uuid4().hex)
_request_simulation_task_cancel(task, reason)
body = b''.join(part.encode() if isinstance(part, str) else part
for part in simulation_event_stream(self.xml, task=task, raw_series=True))
result = next(event['result'] for event in map(json.loads, body.splitlines()) if event['event']=='result')
self.assertEqual(result['status'], status)
self.assertTrue(result['partial'])
self.assertEqual(result['simulatedUntil'], 0.0)
self.assertEqual(result['series']['time'][-1], result['simulatedUntil'])
self.assertEqual(AsgiClient(app).get('/api/system-xml/simulations/'+task.simulation_id).json()['result'], result)
def test_index_corruption_or_truncated_output_is_rejected(self):
directory = self.root/'corrupt'; directory.mkdir(exist_ok=True)
data = (self.root/'indexed/result.json').read_bytes()
original = json.loads((self.root/'indexed/result-index.json').read_bytes())
output = directory/'result.json'; index = directory/'index.json'
output.write_bytes(data)
for change in ({'version':2}, {'version':True}, {'version':1.0}, {'seriesStart':True}, {'seriesStart':-1}, {'seriesEnd':len(data)+1},
{'resultBytes':len(data)-1}, {'sampleCount':-1}, {'seriesStart':original['seriesStart']+1}):
with self.subTest(change=change):
index.write_text(json.dumps(original | change))
with self.assertRaises(ValueError): read_indexed_result(output,index)
index.write_text(json.dumps(original)); output.write_bytes(data[:-2])
with self.assertRaises(ValueError): read_indexed_result(output,index)
def test_json_framing_and_escaping_cannot_confuse_raw_series(self):
values = {'time':[0,1], 'odd\\"},"final":{\n温度':[1e-300,-0.0]}
raw = NativeSeriesJson(json.dumps(values,ensure_ascii=False).encode(),2)
for payload in ({'result':{'series':raw}},
{'event':'result','message':'\n"series":{},"result":null',
'result':{'series':raw,'label':'\\"雪\n','success':True}}):
encoded = b''.join(serialize_result_parts(payload))
expected = dict(payload, result=dict(payload['result'], series=values))
self.assertEqual(json.loads(encoded), expected)
self.assertEqual(json.loads(b''.join(serialize_result_parts({'event':'progress'}))), {'event':'progress'})
def test_solve_only_raw_series_is_empty_object(self):
result = execute_native(self.build, replace(self.config,t_stop=.001),.001,
run_dir=self.root/'solve-only', raw_series=True,record_samples=False)
self.assertEqual(result['series'].data,b'{}')
self.assertEqual(result['series'].sample_count,0)
if __name__ == '__main__': unittest.main()