@@ -129,10 +129,25 @@ int main(void) {
129129 "\"epochs\":50}')" ,
130130 1 );
131131
132+ /* fit() aggregate: train + register native students (gbt and tree) once
133+ * here (a second register would hit MODEL_EXISTS); the scalar predict()
134+ * inside the loop serves them per row. */
135+ if (run_discard (db ,
136+ "SELECT fit(f1, f2, label,"
137+ " '{\"kind\":\"gbt\",\"register\":\"soak_fit\"}') FROM tab" ,
138+ 1 ) ||
139+ run_discard (db ,
140+ "SELECT fit(f1, f2, label,"
141+ " '{\"kind\":\"tree\",\"register\":\"soak_fit_tree\"}')"
142+ " FROM tab" ,
143+ 1 ))
144+ goto done_fail ;
145+
132146 int fails = 0 ;
133147 for (int i = 0 ; i < 50 ; i ++ ) {
134148 /* serving loop: the aggregate forms (the one calling convention for
135- * forecast/detect_anomalies) + predict, across models and options */
149+ * forecast/detect_anomalies) + batched predict_batch, across models
150+ * and options */
136151 if (run_discard (db , "SELECT forecast(ts, value, 6) FROM series" , 1 ) ||
137152 run_discard (db ,
138153 "SELECT grp, forecast(ts, value, 4) FROM series"
@@ -145,36 +160,49 @@ int main(void) {
145160 " '{\"model\":\"sub-pca\"}') FROM series" ,
146161 1 ) ||
147162 run_discard (db ,
148- "SELECT * FROM predict ("
163+ "SELECT * FROM predict_batch ("
149164 "'SELECT f1, f2, label FROM tab WHERE id < 100',"
150165 "'SELECT id, f1, f2 FROM tab WHERE id >= 100',"
151166 " '{\"target\":\"label\"}')" ,
152167 1 ) ||
153- run_discard (db , "SELECT * FROM predict (NULL,'SELECT id, f1, f2 FROM tab',"
168+ run_discard (db , "SELECT * FROM predict_batch (NULL,'SELECT id, f1, f2 FROM tab',"
154169 " '{\"model\":\"soak_student\"}')" ,
155170 1 ))
156171 goto done_fail ;
157172
158- /* gbt-student (forest runtime) + tree/forest error paths */
159- fails += run_discard (db , "SELECT * FROM predict(NULL,'SELECT id, f1, f2 FROM tab',"
173+ /* scalar predict(): serve the fit()-registered students per row, plus
174+ * proba, a fit() blob via subquery, and the arity/unknown-model errors */
175+ fails += run_discard (db , "SELECT id, predict('soak_fit', f1, f2) FROM tab" , 1 );
176+ fails += run_discard (db , "SELECT id, predict('soak_fit_tree', f1, f2) FROM tab" , 1 );
177+ fails += run_discard (db , "SELECT id, predict('soak_fit', f1, f2,"
178+ " '{\"proba\":true}') FROM tab" ,
179+ 1 );
180+ fails += run_discard (db , "SELECT id, predict((SELECT fit(f1, f2, label)"
181+ " FROM tab), f1, f2) FROM tab" ,
182+ 1 );
183+ fails += run_discard (db , "SELECT predict('soak_fit', f1) FROM tab" , 0 );
184+ fails += run_discard (db , "SELECT predict('nope', f1, f2) FROM tab" , 0 );
185+
186+ /* gbt-student (forest runtime) + tree/forest error paths, served batched */
187+ fails += run_discard (db , "SELECT * FROM predict_batch(NULL,'SELECT id, f1, f2 FROM tab',"
160188 " '{\"model\":\"soak_gbt\"}')" ,
161189 1 );
162- fails += run_discard (db , "SELECT * FROM predict (NULL,'SELECT id, f1, f2 FROM tab',"
190+ fails += run_discard (db , "SELECT * FROM predict_batch (NULL,'SELECT id, f1, f2 FROM tab',"
163191 " '{\"model\":\"soak_knn\"}')" ,
164192 1 );
165- fails += run_discard (db , "SELECT * FROM predict (NULL,'SELECT id, f1, f2 FROM tab',"
193+ fails += run_discard (db , "SELECT * FROM predict_batch (NULL,'SELECT id, f1, f2 FROM tab',"
166194 " '{\"model\":\"soak_soft\"}')" ,
167195 1 );
168- fails += run_discard (db , "SELECT * FROM predict (NULL,'SELECT id, f1, f2 FROM tab',"
196+ fails += run_discard (db , "SELECT * FROM predict_batch (NULL,'SELECT id, f1, f2 FROM tab',"
169197 " '{\"model\":\"soak_mlp\"}')" ,
170198 1 );
171- fails += run_discard (db , "SELECT * FROM predict (NULL,'SELECT id, f1, f2 FROM tab',"
199+ fails += run_discard (db , "SELECT * FROM predict_batch (NULL,'SELECT id, f1, f2 FROM tab',"
172200 " '{\"model\":\"soak_bad\"}')" ,
173201 0 );
174- fails += run_discard (db , "SELECT * FROM predict (NULL,'SELECT id, f1, f2 FROM tab',"
202+ fails += run_discard (db , "SELECT * FROM predict_batch (NULL,'SELECT id, f1, f2 FROM tab',"
175203 " '{\"model\":\"soak_gbt_bad\"}')" ,
176204 0 );
177- fails += run_discard (db , "SELECT * FROM predict (NULL,'SELECT id, f1, f2 FROM tab',"
205+ fails += run_discard (db , "SELECT * FROM predict_batch (NULL,'SELECT id, f1, f2 FROM tab',"
178206 " '{\"model\":\"soak_mlp_bad\"}')" ,
179207 0 );
180208 fails += run_discard (db ,
@@ -185,8 +213,8 @@ int main(void) {
185213 " '{\"student_id\":\"nope\"}')" ,
186214 0 );
187215
188- /* auto selection, conformal intervals, backtest() (grouped, gapped) --
189- * exercises the rolling-origin backtest scratch allocations */
216+ /* auto selection, conformal intervals, backtest() aggregate (grouped,
217+ * gapped) -- exercises the rolling-origin backtest scratch allocations */
190218 fails += run_discard (db , "SELECT forecast(ts, value, 6,"
191219 " '{\"model\":\"auto\"}') FROM series" ,
192220 1 );
@@ -198,15 +226,19 @@ int main(void) {
198226 "\"conformal\",\"folds\":8,\"gap\":2}')"
199227 " FROM series GROUP BY grp" ,
200228 1 );
201- fails += run_discard (db , "SELECT * FROM backtest('SELECT ts, value FROM series' , 6,"
202- " '{\"folds\":10}') " ,
229+ fails += run_discard (db , "SELECT backtest(ts, value, 6, '{\"folds\":10}') "
230+ " FROM series " ,
203231 1 );
204- fails += run_discard (db , "SELECT * FROM backtest('SELECT ts, value FROM series' , 6,"
232+ fails += run_discard (db , "SELECT backtest(ts, value, 6,"
205233 " '{\"model\":\"auto\",\"interval_method\":\"conformal\","
206- "\"folds\":12}')" ,
234+ "\"folds\":12}') FROM series" ,
235+ 1 );
236+ fails += run_discard (db , "SELECT grp, backtest(ts, value, 5, '{\"gap\":3}')"
237+ " FROM series GROUP BY grp" ,
207238 1 );
208- fails += run_discard (db , "SELECT * FROM backtest('SELECT ts, value, grp FROM series',"
209- " 5, '{\"group_cols\":[\"grp\"],\"gap\":3}')" ,
239+ /* backtest_rows expansion round-trip */
240+ fails += run_discard (db , "SELECT * FROM backtest_rows((SELECT backtest(ts,"
241+ " value, 6, '{\"folds\":8}') FROM series))" ,
210242 1 );
211243 /* option error + edge paths */
212244 fails += run_discard (db , "SELECT forecast(ts, value, 6,"
@@ -215,11 +247,11 @@ int main(void) {
215247 fails += run_discard (db , "SELECT forecast(ts, value, 6, '{\"folds\":0}')"
216248 " FROM series" ,
217249 0 );
218- fails += run_discard (db , "SELECT * FROM backtest('SELECT ts, value FROM series' , 6,"
219- " '{\"model\":\"nope\"}') " ,
250+ fails += run_discard (db , "SELECT backtest(ts, value, 6, '{\"model\":\"nope\"}') "
251+ " FROM series " ,
220252 0 );
221- fails += run_discard (db , "SELECT * FROM backtest('SELECT ts, value FROM series' , 6,"
222- " '{\"gap\":100000}') " ,
253+ fails += run_discard (db , "SELECT backtest(ts, value, 6, '{\"gap\":100000}') "
254+ " FROM series " ,
223255 1 );
224256
225257 /* tsb + auto candidate sets (statistical models; the student-candidate
@@ -262,8 +294,8 @@ int main(void) {
262294 0 );
263295
264296 /* error paths every iteration too: aggregate misuse and option
265- * rejection, plus backtest's collect_series failure branches, so
266- * valgrind sees the partial-series cleanup under load */
297+ * rejection, plus the backtest aggregate 's partial-group cleanup under
298+ * load, so valgrind sees the failure branches free everything */
267299 fails += run_discard (db , "SELECT forecast('SELECT ts FROM series', value, 3)"
268300 " FROM series" ,
269301 0 );
@@ -287,35 +319,34 @@ int main(void) {
287319 0 );
288320 fails += run_discard (db , "SELECT forecast(ts, value, 0) FROM series" , 0 );
289321 fails += run_discard (db , "SELECT forecast(ts, value, 4) FROM series WHERE 0" , 1 );
290- fails += run_discard (db , "SELECT * FROM backtest('DELETE FROM series', 3)" , 0 );
291- fails += run_discard (db , "SELECT * FROM backtest('NOT SQL', 3)" , 0 );
292- fails += run_discard (db , "SELECT * FROM backtest('SELECT ts FROM series', 3)" , 0 );
293- fails += run_discard (db , "SELECT * FROM backtest('SELECT ts, value FROM series', 3,"
294- " '{\"time_col\":\"nope\"}')" ,
322+ /* backtest aggregate error + edge paths (no inner query to validate now) */
323+ fails += run_discard (db , "SELECT backtest(ts, value, 0) FROM series" , 0 );
324+ fails += run_discard (db , "SELECT backtest(ts, value, 3, '{\"bogus\":1}')"
325+ " FROM series" ,
295326 0 );
296- fails += run_discard (db , "SELECT * FROM backtest('SELECT ts, value, grp FROM series', "
297- " 3, '{\"group_cols\":[\" nope\"] }')" ,
327+ fails += run_discard (db , "SELECT backtest(ts, value, 3, '{\"time_col\": "
328+ "\" nope\"}') FROM series " ,
298329 0 );
330+ fails += run_discard (db , "SELECT backtest(ts, value, 3) FROM series WHERE 0" , 1 );
299331 /* duplicate option keys (the CI fuzzer's leak): last-wins, no leak */
300332 fails += run_discard (db , "SELECT forecast(ts, value, 3,"
301333 " '{\"model\":\"theta-classic\",\"model\":"
302334 "\"stub-seasonal-naive\"}') FROM series" ,
303335 1 );
304336 fails += run_discard (db ,
305- "SELECT * FROM backtest('SELECT ts, value, grp FROM series',"
306- " 3, '{\"model\":\"theta-classic\",\"model\":"
307- "\"stub-seasonal-naive\",\"time_col\":\"ts\",\"time_col\":"
308- "\"ts\",\"group_cols\":[\"grp\"],\"group_cols\":[\"grp\"]}')" ,
337+ "SELECT grp, backtest(ts, value, 3,"
338+ " '{\"model\":\"theta-classic\",\"model\":"
339+ "\"stub-seasonal-naive\"}') FROM series GROUP BY grp" ,
309340 1 );
310341 fails += run_discard (db ,
311- "SELECT * FROM predict ('SELECT f1, f2, label FROM tab',"
342+ "SELECT * FROM predict_batch ('SELECT f1, f2, label FROM tab',"
312343 " 'SELECT id, f1, f2 FROM tab',"
313344 " '{\"target\":\"label\",\"task\":\"classify\",\"task\":"
314345 "\"classify\",\"model\":\"knn5-incontext\",\"model\":"
315346 "\"knn5-incontext\"}')" ,
316347 1 );
317348 fails += run_discard (db ,
318- "SELECT * FROM predict ('SELECT f1, f2, label FROM tab',"
349+ "SELECT * FROM predict_batch ('SELECT f1, f2, label FROM tab',"
319350 " 'SELECT id, f1 FROM tab', '{\"target\":\"label\"}')" ,
320351 0 );
321352 fails += run_discard (db , "SELECT predict_ulid('not a time')" , 0 );
@@ -329,6 +360,7 @@ int main(void) {
329360 1 );
330361 fails += run_discard (db , "SELECT * FROM forecast_rows('not json')" , 0 );
331362 fails += run_discard (db , "SELECT * FROM anomaly_rows('[1,2]')" , 0 );
363+ fails += run_discard (db , "SELECT * FROM backtest_rows('not json')" , 0 );
332364 fails += run_discard (db , "SELECT * FROM forecast_rows(NULL)" , 1 );
333365 }
334366
0 commit comments