Files
SystemSimulationApp/tests/test_native_solver_control.py
T

124 lines
6.2 KiB
Python

"""Exercise the production wrapper with real CVODE plus injected return cases.
The injection lives only in a temporary copy of the runtime. It uses the public
SUNDIALS ABI and the normal builder on both Windows and Linux.
"""
import json
import os
from pathlib import Path
import shutil
import subprocess
import tempfile
import unittest
from unittest.mock import patch
from app.simulation.native_codegen import build as builder
from app.simulation.native_codegen.compiler import compile_native_program
from tests.test_native_catalog import Circuit
ROOT=Path(__file__).resolve().parents[1]
INJECTION=r'''
#include <stdlib.h>
static int test_step(NativeRun *r,void *solver,double end,N_Vector y,double *next,int task) {
static unsigned calls=0;
const char *mode=getenv("NATIVE_WRAPPER_TEST");
calls++;
if(mode && !strcmp(mode,"recover") && calls<=1000) return CV_SUCCESS;
if(mode && !strcmp(mode,"solver-error")) return CV_CONV_FAILURE;
if(mode && !strcmp(mode,"backward")) {*next-=1;return CV_SUCCESS;}
if(mode && !strcmp(mode,"nan-time")) {*next=NAN;return CV_SUCCESS;}
if(mode && !strcmp(mode,"nan-state")) {N_VGetArrayPointer(y)[0]=NAN;return CV_SUCCESS;}
if(mode && !strcmp(mode,"stall")) return CV_SUCCESS;
if(mode && !strcmp(mode,"timeout")) {
r->wall_start-=r->options.timeout+1;
native_poll(r,*next);return CV_RHSFUNC_FAIL;
}
if(mode && !strcmp(mode,"cancel")) {
if(calls==32) {FILE *f=fopen(r->options.cancel_path,"wb");if(f){fputs("cancel",f);fclose(f);}}
return CV_SUCCESS;
}
if(mode && !strcmp(mode,"step-limit")) r->accepted=10000001;
if(mode && !strcmp(mode,"event-limit")) r->events=10001;
return CVode(solver,end,y,next,task);
}
'''
class NativeSolverControlTests(unittest.TestCase):
@classmethod
def setUpClass(cls):
cls.tmp=tempfile.TemporaryDirectory(prefix='native-wrapper-')
cls.addClassCleanup(cls.tmp.cleanup)
cls.root=Path(cls.tmp.name)
native=cls.root/'native'
shutil.copytree(ROOT/'native',native)
source=native/'runtime/cvode_solver.c'
text=source.read_text(encoding='utf-8')
needle='int flag=CVode(solver,end,y,&next,CV_ONE_STEP);'
assert text.count(needle)==1
text=text.replace(needle,'int flag=test_step(r,solver,end,y,&next,CV_ONE_STEP);')
text=text.replace('int native_bdf(NativeRun *r) {',INJECTION+'\nint native_bdf(NativeRun *r) {')
source.write_text(text,encoding='utf-8')
b=Circuit()
m=b.add('amesim_mecmas21','mass',mass=1,v0=1,x0=0,useFriction=1,stoptype=4)
for i in (1,2):
f=b.add('amesim_f000','free'+str(i));b.connect(m,'port_'+str(i),f,'port_1')
with patch.object(builder,'NATIVE',native):
cls.build=builder.build_native(compile_native_program(b.net),cache_dir=cls.root/'cache')
cls.addClassCleanup(cls.build.close)
def run_case(self,mode,timeout=1):
with tempfile.TemporaryDirectory(dir=self.root) as tmp:
path=Path(tmp)/'result.json'
env=dict(os.environ,NATIVE_WRAPPER_TEST=mode)
run=subprocess.run([str(self.build.executable),'--method','BDF','--stop','.2','--sample-step','.01',
'--max-step','.01','--rtol','1e-8','--timeout',str(timeout),'--cancel-file',str(Path(tmp)/'cancel'),
'--sample-file',str(Path(tmp)/'states.bin'),'--output-block-file',str(Path(tmp)/'outputs.bin'),
'--output',str(path)],capture_output=True,env=env,timeout=10)
result=json.loads(path.read_text(encoding='utf-8'))
self.assertEqual(run.returncode,2 if result['status']=='failed' else 0)
self.assertTrue(all(a<b for a,b in zip(result['series']['time'],result['series']['time'][1:])))
return result
def test_successful_equal_times_have_no_retry_count_limit_or_zero_length_events(self):
normal=self.run_case('normal');recovered=self.run_case('recover')
self.assertTrue(recovered['success'])
self.assertEqual(recovered['simulatedUntil'],.2)
self.assertEqual(recovered['series'],normal['series'])
self.assertEqual(recovered['solverStarts'],normal['solverStarts'])
self.assertEqual(recovered['stateTransitions'],0)
self.assertEqual(recovered['acceptedSteps'],normal['acceptedSteps'])
self.assertEqual(recovered['solverControl']['sameTimeReturns'],1000)
self.assertEqual(recovered['solverControl']['maxSameTimeStreak'],1000)
def test_distinct_errors_preserve_last_valid_output_and_cvode_code(self):
for mode,reason in [('solver-error','solver-error'),('backward','time-regression'),
('nan-time','nonfinite-time'),('nan-state','nonfinite-state')]:
with self.subTest(mode=mode):
result=self.run_case(mode)
self.assertFalse(result['success'])
self.assertEqual(result['solverControl']['reason'],reason)
self.assertEqual(result['simulatedUntil'],0)
self.assertEqual(result['final']['mass.v'],1)
if mode=='solver-error':self.assertEqual(result['solverControl']['returnCode'],-4)
if mode=='nan-time':self.assertIsNone(result['solverControl']['returnTime'])
def test_persistent_stagnation_expires_by_wall_time_and_remains_cancellable(self):
result=self.run_case('stall',timeout=.03)
self.assertEqual(result['solverControl']['reason'],'time-stagnation')
self.assertGreaterEqual(result['solveSeconds'],.03)
self.assertLess(result['solveSeconds'],1)
cancelled=self.run_case('cancel')
self.assertEqual(cancelled['status'],'cancelled')
self.assertEqual(cancelled['solverControl']['reason'],'cancelled')
self.assertGreaterEqual(cancelled['solverControl']['sameTimeReturns'],32)
def test_timeout_and_resource_guards_have_separate_reasons(self):
result=self.run_case('timeout')
self.assertEqual(result['solverControl']['reason'],'timeout')
for mode in ('step-limit','event-limit'):
with self.subTest(mode=mode):
result=self.run_case(mode)
self.assertEqual(result['solverControl']['reason'],'resource-limit')
if __name__=='__main__':unittest.main()