Spaces:
Sleeping
Sleeping
| import re | |
| from sql_env.models import SQLAction, SQLTask, SQLReward | |
| def _normalize(query: str) -> str: | |
| """Uppercase, collapse whitespace, strip trailing semicolons.""" | |
| q = query.strip().upper() | |
| q = re.sub(r'\s+', ' ', q) | |
| q = q.rstrip(';').strip() | |
| return q | |
| def _tokenize(query: str) -> set: | |
| return set(re.findall(r"[A-Z0-9_'*.=><]+", _normalize(query))) | |
| def _sql_keywords_present(query: str) -> set: | |
| keywords = { | |
| 'SELECT', 'FROM', 'WHERE', 'GROUP', 'BY', 'HAVING', | |
| 'ORDER', 'JOIN', 'INNER', 'LEFT', 'RIGHT', 'OUTER', | |
| 'BETWEEN', 'DESC', 'ASC', 'LIMIT', 'COUNT', 'SUM', | |
| 'AVG', 'MAX', 'MIN', 'AS', 'ON', 'AND', 'OR', 'NOT', | |
| 'IN', 'LIKE', 'IS', 'NULL', 'DISTINCT' | |
| } | |
| normed = _normalize(query) | |
| found = set() | |
| for kw in keywords: | |
| if re.search(r'\b' + kw + r'\b', normed): | |
| found.add(kw) | |
| return found | |
| def grade(action: SQLAction, task: SQLTask) -> SQLReward: | |
| """ | |
| 4-level grader with partial progress signals. | |
| 1.0 β exact normalized match | |
| 0.7 β same tokens, minor whitespace/alias differences | |
| 0.4 β key SQL keywords all present and correct table/column names | |
| 0.2 β basic SELECT/FROM structure present | |
| 0.0 β completely wrong | |
| """ | |
| agent = _normalize(action.corrected_query) | |
| correct = _normalize(task.canonical_answer) | |
| # ββ Level 1: Exact match βββββββββββββββββββββββββββββββββ | |
| if agent == correct: | |
| return SQLReward(value=1.0, reason="Exact match β perfect correction.") | |
| # ββ Level 2: Same token set (right words, minor ordering) β | |
| agent_tokens = _tokenize(action.corrected_query) | |
| correct_tokens = _tokenize(task.canonical_answer) | |
| if agent_tokens == correct_tokens: | |
| return SQLReward( | |
| value=0.7, | |
| reason="All correct tokens present but structure differs slightly." | |
| ) | |
| # ββ Level 3: Most keywords correct + high token overlap βββ | |
| correct_kws = _sql_keywords_present(task.canonical_answer) | |
| agent_kws = _sql_keywords_present(action.corrected_query) | |
| kw_overlap = len(correct_kws & agent_kws) / max(len(correct_kws), 1) | |
| token_overlap = len(agent_tokens & correct_tokens) / max(len(correct_tokens), 1) | |
| if kw_overlap >= 0.85 and token_overlap >= 0.75: | |
| return SQLReward( | |
| value=0.4, | |
| reason=f"Most keywords correct ({kw_overlap:.0%} keyword match, {token_overlap:.0%} token match)." | |
| ) | |
| # ββ Level 4: Basic structure present βββββββββββββββββββββ | |
| if 'SELECT' in agent and 'FROM' in agent: | |
| return SQLReward( | |
| value=0.2, | |
| reason="Basic SELECT/FROM structure present but significant errors remain." | |
| ) | |
| # ββ Level 0: No recognizable SQL βββββββββββββββββββββββββ | |
| return SQLReward(value=0.0, reason="Response is not valid SQL.") | |
| def generate_feedback(action: SQLAction, task: SQLTask, reward: SQLReward) -> str: | |
| """Human-readable feedback shown in next observation.""" | |
| if reward.value == 1.0: | |
| return "Correct! Query matches perfectly." | |
| if reward.value >= 0.7: | |
| return "Very close β check spacing or minor clause differences." | |
| if reward.value >= 0.4: | |
| return "Good progress β most keywords are right, but check for typos in keywords or column names." | |
| if reward.value >= 0.2: | |
| return "Basic structure is there β look carefully at every SQL keyword for typos." | |
| return "The response doesn't look like valid SQL. Start with SELECT ... FROM ..." | |