4242 "params" : [],
4343}
4444
45+
4546class OrionExtensionTest :
4647 """Base orion extension interface you need to implement"""
48+
4749 def __init__ (self ) -> None :
4850 self .calls = defaultdict (int )
4951
5052 def on_experiment_error (self , * args , ** kwargs ):
51- self .calls [' on_experiment_error' ] += 1
53+ self .calls [" on_experiment_error" ] += 1
5254
5355 def on_trial_error (self , * args , ** kwargs ):
54- self .calls [' on_trial_error' ] += 1
56+ self .calls [" on_trial_error" ] += 1
5557
5658 def start_experiment (self , * args , ** kwargs ):
57- self .calls [' start_experiment' ] += 1
59+ self .calls [" start_experiment" ] += 1
5860
5961 def new_trial (self , * args , ** kwargs ):
60- self .calls [' new_trial' ] += 1
62+ self .calls [" new_trial" ] += 1
6163
6264 def end_trial (self , * args , ** kwargs ):
63- self .calls [' end_trial' ] += 1
65+ self .calls [" end_trial" ] += 1
6466
6567 def end_experiment (self , * args , ** kwargs ):
66- self .calls [' end_experiment' ] += 1
68+ self .calls [" end_experiment" ] += 1
6769
6870
6971def test_client_extension ():
@@ -88,31 +90,38 @@ def foo(x):
8890 n_broken = len (experiment .fetch_trials_by_status ("broken" ))
8991 n_reserved = len (experiment .fetch_trials_by_status ("reserved" ))
9092
91- assert ext .calls ['new_trial' ] == n_trials + n_broken - n_reserved , 'all trials should have triggered callbacks'
92- assert ext .calls ['end_trial' ] == n_trials + n_broken - n_reserved , 'all trials should have triggered callbacks'
93- assert ext .calls ['on_trial_error' ] == n_broken , 'failed trial should be reported '
93+ assert (
94+ ext .calls ["new_trial" ] == n_trials + n_broken - n_reserved
95+ ), "all trials should have triggered callbacks"
96+ assert (
97+ ext .calls ["end_trial" ] == n_trials + n_broken - n_reserved
98+ ), "all trials should have triggered callbacks"
99+ assert (
100+ ext .calls ["on_trial_error" ] == n_broken
101+ ), "failed trial should be reported "
94102
95- assert ext .calls [' start_experiment' ] == 1 , ' experiment should have started'
96- assert ext .calls [' end_experiment' ] == 1 , ' experiment should have ended'
97- assert ext .calls [' on_experiment_error' ] == 1 , ' failed experiment '
103+ assert ext .calls [" start_experiment" ] == 1 , " experiment should have started"
104+ assert ext .calls [" end_experiment" ] == 1 , " experiment should have ended"
105+ assert ext .calls [" on_experiment_error" ] == 1 , " failed experiment "
98106
99107 unregistered_callback = client .extensions .unregister (ext )
100108 assert unregistered_callback == 6 , "All ext callbacks got unregistered"
101109
102110
103111class BadOrionExtensionTest :
104112 """Base orion extension interface you need to implement"""
113+
105114 def __init__ (self ) -> None :
106115 self .calls = defaultdict (int )
107116
108117 def on_extension_error (self , name , fun , exception , args ):
109- self .calls [' on_extension_error' ] += 1
118+ self .calls [" on_extension_error" ] += 1
110119
111120 def on_experiment_error (self , * args , ** kwargs ):
112- self .calls [' on_experiment_error' ] += 1
121+ self .calls [" on_experiment_error" ] += 1
113122
114123 def on_trial_error (self , * args , ** kwargs ):
115- self .calls [' on_trial_error' ] += 1
124+ self .calls [" on_trial_error" ] += 1
116125
117126 def new_trial (self , * args , ** kwargs ):
118127 raise RuntimeError ()
@@ -132,9 +141,9 @@ def foo(x):
132141 assert client .max_trials == MAX_TRIALS
133142 client .workon (foo , max_trials = MAX_TRIALS , max_broken = MAX_BROKEN )
134143
135- assert ext .calls [' on_trial_error' ] == 0 , ' Orion worked as expected'
136- assert ext .calls [' on_experiment_error' ] == 0 , ' Orion worked as expected'
137- assert ext .calls [' on_extension_error' ] == 9 , ' Extension error got reported'
144+ assert ext .calls [" on_trial_error" ] == 0 , " Orion worked as expected"
145+ assert ext .calls [" on_experiment_error" ] == 0 , " Orion worked as expected"
146+ assert ext .calls [" on_extension_error" ] == 9 , " Extension error got reported"
138147
139148 unregistered_callback = client .extensions .unregister (ext )
140149 assert unregistered_callback == 4 , "All ext callbacks got unregistered"
0 commit comments