Kalletlamadhav/sql-optimization-env
0
1# server/reward.py2 3import math4from .models import SQLOptReward, SQLOptAction5from .graders.speedup_grader import SpeedupGrader6from .graders.equivalence_grader import EquivalenceGrader7from .graders.antipattern_grader import AntiPatternGrader8from .graders.index_grader import IndexGrader9 10 11class RewardComposer:12 13 WEIGHTS = {14 'speedup': 0.35,15 'equivalence': 0.25,16 'pattern': 0.20,17 'index': 0.10,18 'simplicity': 0.10,19 }20 21 PENALTIES = {22 'syntax_error': -0.30,23 'wrong_results': -0.20,24 'slower_query': -0.10,25 'timeout_per_second': -0.05,26 'hack_detected': -0.40,27 }28 29 def __init__(self):30 self.speedup_grader = SpeedupGrader()31 self.equiv_grader = EquivalenceGrader()32 self.pattern_grader = AntiPatternGrader()33 self.index_grader = IndexGrader()34 35 def compute(36 self,37 task,38 action: SQLOptAction,39 orig_time: float,40 opt_time: float,41 opt_rows,42 opt_plan,43 hack,44 query_error45 ) -> SQLOptReward:46 47 penalty = 0.048 49 # ❗ Handle query failure immediately50 if query_error:51 return SQLOptReward(52 total=0.0,53 speedup_score=0.0,54 equivalence_score=0.0,55 pattern_score=0.0,56 index_score=0.0,57 simplicity_score=0.0,58 penalties=self.PENALTIES['syntax_error'],59 speedup_ratio=0.0,60 hack_detected=False,61 hack_type=None62 )63 64 # 1️⃣ Speedup Score65 speedup_ratio = orig_time / max(opt_time, 0.1)66 67 # Log scaling68 speedup_score = min(69 0.35,70 0.35 * (math.log10(max(speedup_ratio, 1)) / 2)71 )72 73 if speedup_ratio < 1.0:74 penalty += self.PENALTIES['slower_query']75 76 # 2️⃣ Equivalence Score77 equiv_score = self.equiv_grader.grade(task, action.optimized_query)78 equiv_score *= 0.2579 80 if equiv_score < 0.125:81 penalty += self.PENALTIES['wrong_results']82 83 # 3️⃣ Anti-pattern Score84 pattern_score = self.pattern_grader.grade(task, action) * 0.2085 86 # 4️⃣ Index Score87 index_score = self.index_grader.grade(task, action, opt_plan) * 0.1088 89 # 5️⃣ Simplicity Score90 simplicity_score = (91 self._simplicity_score(task.slow_query, action.optimized_query)92 * 0.1093 )94 95 # ⏱ Timeout penalty (>5s)96 if opt_time > 5000:97 penalty += self.PENALTIES['timeout_per_second'] * ((opt_time / 1000) - 5)98 99 # 🚨 Hack detection100 if hack:101 penalty += self.PENALTIES['hack_detected']102 103 # ✅ FIXED LINE BREAK104 raw_total = (105 speedup_score106 + equiv_score107 + pattern_score108 + index_score109 + simplicity_score110 + penalty111 )112 113 total = round(max(0.0, min(1.0, raw_total)), 4)114 115 return SQLOptReward(116 total=total,117 speedup_score=round(speedup_score, 4),118 equivalence_score=round(equiv_score, 4),119 pattern_score=round(pattern_score, 4),120 index_score=round(index_score, 4),121 simplicity_score=round(simplicity_score, 4),122 penalties=round(penalty, 4),123 speedup_ratio=round(speedup_ratio, 2),124 hack_detected=bool(hack),125 hack_type=hack126 )127 128 def _simplicity_score(self, original: str, optimized: str) -> float:129 """130 Reward simpler queries (only if logically correct)131 """132 orig_len = len(original.strip())133 opt_len = len(optimized.strip())134 135 if opt_len <= orig_len:136 return 1.0137 138 ratio = orig_len / max(opt_len, 1)139 return max(0.0, ratio)