File size: 3,874 Bytes
b296ad4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Round-trip direct paraphrases to SQL without exposing reference SQL to the verifier."""
import argparse
import concurrent.futures
import json
import time
import urllib.request
from pathlib import Path
from tinyquery.data import compact,validate_sql


def normalized(rows,ordered):
    return rows if ordered else sorted(rows,key=lambda row:repr(row))


def verify(item,base):
    records=item['accepted']
    if not records: return {'scenario_id':item['scenario_id'],'accepted':[],'rejected':[]}
    sample=records[0]
    payload={'model':'Qwen/Qwen3.8-27B-FP8','messages':[{'role':'user','content':
        'Translate each request to one read-only '+('PostgreSQL' if sample['backend']=='supabase' else 'MySQL')+
        ' SELECT query using only this schema. Return a JSON object mapping each language key to its SQL string. '
        'If a request is ambiguous or impossible, use an empty string. Do not add constraints. '+
        compact({'schema':sample['context']['schema'],'requests':{r['language']:r['question'] for r in records}})}],
        'max_tokens':700,'temperature':0,'response_format':{'type':'json_object'},
        'chat_template_kwargs':{'enable_thinking':False}}
    request=urllib.request.Request(base+'/chat/completions',data=compact(payload).encode(),headers={'Content-Type':'application/json'})
    with urllib.request.urlopen(request,timeout=180) as response: raw=json.load(response)
    answers=json.loads(raw['choices'][0]['message']['content'])
    gold=next(v for k,v in sample['target']['arguments'].items() if k in ['query','sql'])
    reference=validate_sql(gold,sample['backend'],sample['slots'])
    ordered=sample['operation'] in ['top','bottom','sort_asc','sort_desc']
    accepted=[]; rejected=[]
    for record in records:
        try:
            sql=answers[record['language']]
            results=validate_sql(sql,record['backend'],record['slots'])
            if any(normalized(a,ordered)!=normalized(b,ordered) for a,b in zip(reference,results)):
                raise ValueError('SQL results differ on fixtures')
            record['roundtrip_sql']=sql
            record['provenance']+='; teacher round-trip SQL equivalent on two generated SQLite fixtures'
            accepted.append(record)
        except Exception as exc:
            rejected.append({'id':record['id'],'question':record['question'],'reason':str(exc),'sql':answers.get(record['language'])})
    return {'scenario_id':item['scenario_id'],'accepted':accepted,'rejected':rejected,'verification_request':payload,'verification_response':raw}


def main():
    p=argparse.ArgumentParser(); p.add_argument('--input',required=True); p.add_argument('--out',required=True)
    p.add_argument('--workers',type=int,default=24); p.add_argument('--base',default='http://127.0.0.1:18000/v1')
    args=p.parse_args(); output=Path(args.out); done=set()
    if output.exists(): done={json.loads(l)['scenario_id'] for l in output.read_text().splitlines()}
    items=[json.loads(l) for l in Path(args.input).read_text().splitlines()]
    start=time.time(); completed=0; accepted=0; rejected=0
    with output.open('a') as stream,concurrent.futures.ThreadPoolExecutor(args.workers) as pool:
        jobs={pool.submit(verify,r,args.base):r['scenario_id'] for r in items if r['scenario_id'] not in done}
        for future in concurrent.futures.as_completed(jobs):
            completed+=1
            try:
                result=future.result(); accepted+=len(result['accepted']); rejected+=len(result['rejected'])
                stream.write(compact(result)+'\n'); stream.flush()
            except Exception as exc: print(compact({'error':str(exc),'scenario_id':jobs[future]}),flush=True)
            if completed%25==0: print(compact({'completed':completed,'accepted':accepted,'rejected':rejected,'seconds':time.time()-start}),flush=True)


if __name__=='__main__': main()