File size: 4,558 Bytes
5716a3a
 
 
 
95707b2
 
 
 
 
5716a3a
 
 
 
 
 
 
 
 
95707b2
 
 
5716a3a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a785d13
95707b2
 
 
a785d13
95707b2
5716a3a
 
 
 
 
 
95707b2
2112542
95707b2
 
5716a3a
 
 
 
 
 
95707b2
5716a3a
 
 
 
 
 
 
 
 
 
 
95707b2
5716a3a
 
 
95707b2
 
 
 
 
 
 
5716a3a
 
 
95707b2
5716a3a
 
 
 
2112542
5716a3a
 
 
 
95707b2
5716a3a
 
 
 
 
95707b2
 
5716a3a
 
95707b2
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
import re
from sql_env.models import SQLAction, SQLTask, SQLReward


def _clamp(value: float) -> float:
    """Ensure score is strictly between 0 and 1 (gt=0.0, lt=1.0)."""
    return max(0.001, min(0.999, value))


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:
    normed = _normalize(query)
    normed = re.sub(r"'[^']*'", '__STR__', normed)  # normalize string literals
    return set(re.findall(r"[A-Z0-9_'*.=><]+", normed))


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.
    0.99 β€” 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.01 β€” completely wrong
    All scores are clamped to be strictly between 0.0 and 1.0.
    """
    agent = _normalize(action.corrected_query)
    correct = _normalize(task.canonical_answer)

    # ── Level 1: Exact match ─────────────────────────────────
    if agent == correct:
        return SQLReward(
            value=_clamp(0.999),
            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=_clamp(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=_clamp(0.4),
            reason=f"Most keywords correct ({kw_overlap:.0%} keyword match, {token_overlap:.0%} token match)."
        )

    # ── Level 3.5: Partial keyword and structure match ────────
    if kw_overlap >= 0.7 and token_overlap >= 0.55:
        return SQLReward(
            value=_clamp(0.3),
            reason=f"Partial keyword and structure match ({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=_clamp(0.2),
            reason="Basic SELECT/FROM structure present but significant errors remain."
        )

    # ── Level 0: No recognizable SQL ─────────────────────────
    return SQLReward(value=_clamp(0.001), 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 >= 0.99:
        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.3:
        return "Partial match β€” right direction but several keywords or columns are off."
    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 ..."