1414
1515"""Verification Worker for testing and refining individual kernels."""
1616
17+ import ast
1718import json
1819import logging
1920import multiprocessing as mp
@@ -261,6 +262,27 @@ def _extract_code_from_response(
261262 self .logger .warning ("No code block found in LLM response" )
262263 return None
263264
265+ def _validate_kernel_candidate (self , kernel_code : str | None ) -> str | None :
266+ """Return a reason when extracted kernel code is structurally malformed."""
267+ if not kernel_code or not kernel_code .strip ():
268+ return "no Python kernel code was extracted from the model response"
269+
270+ try :
271+ module = ast .parse (kernel_code )
272+ except SyntaxError as exc :
273+ line = f"line { exc .lineno } " if exc .lineno else "unknown line"
274+ return f"invalid Python syntax ({ line } : { exc .msg } )"
275+
276+ has_kernel_function = any (
277+ isinstance (node , (ast .FunctionDef , ast .AsyncFunctionDef ))
278+ and node .name == "kernel_function"
279+ for node in module .body
280+ )
281+ if not has_kernel_function :
282+ return "missing required top-level kernel_function definition"
283+
284+ return None
285+
264286 def _write_kernel (self , kernel_code : str ):
265287 """Write only the kernel code to file."""
266288 self .kernel_file .write_text (kernel_code )
@@ -413,6 +435,12 @@ def _refine_kernel(
413435 response_text ,
414436 prefer_kernel_function = getattr (self , "_has_multiple_tests" , False ),
415437 )
438+ malformed_reason = self ._validate_kernel_candidate (refined_kernel )
439+
440+ if malformed_reason :
441+ raise ValueError (
442+ f"Malformed LLM kernel response: { malformed_reason } "
443+ )
416444
417445 if refined_kernel :
418446 self .logger .info (
@@ -425,6 +453,8 @@ def _refine_kernel(
425453 return kernel_code
426454
427455 except Exception as e :
456+ if "Malformed LLM kernel response" in str (e ):
457+ raise
428458 self .logger .error (f"Error refining kernel with LLM API: { e } " )
429459 # Fall back to mock refinement
430460
@@ -482,6 +512,17 @@ def run(
482512 self ._has_multiple_tests = len (test_code ) > 1
483513
484514 current_kernel = kernel_code
515+ malformed_reason = self ._validate_kernel_candidate (current_kernel )
516+ if malformed_reason :
517+ error_feedback = f"Malformed LLM kernel response: { malformed_reason } "
518+ self .logger .warning (f"❌ { error_feedback } " )
519+ return {
520+ "worker_id" : self .worker_id ,
521+ "success" : False ,
522+ "rounds" : 0 ,
523+ "error" : error_feedback ,
524+ "history" : list (self .history ),
525+ }
485526
486527 for round_num in range (self .max_rounds ):
487528 # Check if another worker has succeeded
@@ -516,12 +557,38 @@ def run(
516557 "stderr" : violation ,
517558 "history" : list (self .history ),
518559 }
519- current_kernel = self ._refine_kernel (
520- current_kernel ,
521- error_info ,
522- problem_description ,
523- format_test_code_for_llm (test_code ),
524- )
560+ try :
561+ current_kernel = self ._refine_kernel (
562+ current_kernel ,
563+ error_info ,
564+ problem_description ,
565+ format_test_code_for_llm (test_code ),
566+ )
567+ except ValueError as exc :
568+ if "Malformed LLM kernel response" not in str (exc ):
569+ raise
570+ error_feedback = str (exc )
571+ self .logger .warning (f"❌ { error_feedback } " )
572+ return {
573+ "worker_id" : self .worker_id ,
574+ "success" : False ,
575+ "rounds" : round_num + 1 ,
576+ "error" : error_feedback ,
577+ "history" : list (self .history ),
578+ }
579+ malformed_reason = self ._validate_kernel_candidate (current_kernel )
580+ if malformed_reason :
581+ error_feedback = (
582+ f"Malformed LLM kernel response: { malformed_reason } "
583+ )
584+ self .logger .warning (f"❌ { error_feedback } " )
585+ return {
586+ "worker_id" : self .worker_id ,
587+ "success" : False ,
588+ "rounds" : round_num + 1 ,
589+ "error" : error_feedback ,
590+ "history" : list (self .history ),
591+ }
525592 continue
526593
527594 # Log round
@@ -546,12 +613,36 @@ def run(
546613 "history" : list (self .history ),
547614 }
548615
549- current_kernel = self ._refine_kernel (
550- current_kernel ,
551- error_info ,
552- problem_description ,
553- format_test_code_for_llm (test_code ),
554- )
616+ try :
617+ current_kernel = self ._refine_kernel (
618+ current_kernel ,
619+ error_info ,
620+ problem_description ,
621+ format_test_code_for_llm (test_code ),
622+ )
623+ except ValueError as exc :
624+ if "Malformed LLM kernel response" not in str (exc ):
625+ raise
626+ error_feedback = str (exc )
627+ self .logger .warning (f"❌ { error_feedback } " )
628+ return {
629+ "worker_id" : self .worker_id ,
630+ "success" : False ,
631+ "rounds" : round_num + 1 ,
632+ "error" : error_feedback ,
633+ "history" : list (self .history ),
634+ }
635+ malformed_reason = self ._validate_kernel_candidate (current_kernel )
636+ if malformed_reason :
637+ error_feedback = f"Malformed LLM kernel response: { malformed_reason } "
638+ self .logger .warning (f"❌ { error_feedback } " )
639+ return {
640+ "worker_id" : self .worker_id ,
641+ "success" : False ,
642+ "rounds" : round_num + 1 ,
643+ "error" : error_feedback ,
644+ "history" : list (self .history ),
645+ }
555646
556647 # Max rounds reached without success
557648 self .logger .warning (f"Max rounds ({ self .max_rounds } ) reached without success" )
@@ -619,6 +710,12 @@ def verify_with_refinement(
619710 current_kernel = kernel_code
620711 self ._has_multiple_tests = len (test_code ) > 1
621712
713+ malformed_reason = self ._validate_kernel_candidate (current_kernel )
714+ if malformed_reason :
715+ error_feedback = f"Malformed LLM kernel response: { malformed_reason } "
716+ self .logger .warning (f"❌ { error_feedback } " )
717+ return False , current_kernel , error_feedback
718+
622719 # Write files for testing (primary + additional tests)
623720 self ._write_files (current_kernel , test_code )
624721
@@ -653,12 +750,24 @@ def verify_with_refinement(
653750 }
654751
655752 # Refine kernel
656- refined_kernel = self ._refine_kernel (
657- current_kernel ,
658- error_info ,
659- problem_description ,
660- format_test_code_for_llm (test_code ),
661- )
753+ try :
754+ refined_kernel = self ._refine_kernel (
755+ current_kernel ,
756+ error_info ,
757+ problem_description ,
758+ format_test_code_for_llm (test_code ),
759+ )
760+ except ValueError as exc :
761+ if "Malformed LLM kernel response" not in str (exc ):
762+ raise
763+ error_feedback = str (exc )
764+ self .logger .warning (f"❌ { error_feedback } " )
765+ return False , current_kernel , error_feedback
766+ malformed_reason = self ._validate_kernel_candidate (refined_kernel )
767+ if malformed_reason :
768+ error_feedback = f"Malformed LLM kernel response: { malformed_reason } "
769+ self .logger .warning (f"❌ { error_feedback } " )
770+ return False , current_kernel , error_feedback
662771
663772 # Write and test refined kernel
664773 self ._write_kernel (refined_kernel )
0 commit comments