File size: 5,910 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
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
"""Validate queries on disposable native MySQL 8/PostgreSQL 16 temporary tables."""
import argparse
import json
import re
import time
from pathlib import Path
import sqlglot
from sqlglot import exp
from tinyquery.data import ddls,fixture,SQL_OPS


class NativeDB:
    def __init__(self,backend):
        self.backend=backend
        if backend=='mysql':
            import pymysql
            self.conn=pymysql.connect(unix_socket='/run/mysqld/mysqld.sock',user='root',database='tinyquery',autocommit=True)
        else:
            import psycopg
            self.conn=psycopg.connect(host='/var/run/postgresql',user='root',dbname='postgres',autocommit=True)
        self.cursor=self.conn.cursor(); self.tables=[]
        self.cursor.execute('SET SESSION max_execution_time=2000' if backend=='mysql' else "SET statement_timeout='2s'")
    def setup(self,s,seed):
        for table in reversed(self.tables): self.cursor.execute(f'DROP TABLE IF EXISTS {table}')
        self.tables=[]
        for ddl in ddls(s):
            # Relationships are populated identically; MySQL temp tables cannot enforce foreign keys.
            ddl=re.sub(r' REFERENCES [a-zA-Z_]+\(id\)','',ddl)
            self.cursor.execute(ddl.replace('CREATE TABLE','CREATE TEMPORARY TABLE',1))
        self.tables=[s['parent'],s['table']]
        source=fixture(s,seed)
        for table in self.tables:
            rows=source.execute(f'SELECT * FROM {table}').fetchall()
            self.cursor.executemany(f"INSERT INTO {table} VALUES ({','.join(['%s']*len(rows[0]))})",rows)
        source.close()
    def execute(self,sql):
        statements=sqlglot.parse(sql,read='mysql' if self.backend=='mysql' else 'postgres')
        if len(statements)!=1 or not isinstance(statements[0],exp.Select): raise ValueError('Expected one SELECT')
        tree=statements[0]
        if tree.find(exp.Into): raise ValueError('SELECT INTO is not permitted')
        if any(isinstance(node,(exp.DML,exp.DDL)) for node in tree.walk()):
            raise ValueError('Data-changing statements are not permitted inside a query')
        if any(t.name not in self.tables for t in tree.find_all(exp.Table)): raise ValueError('Query references an unknown table')
        for f in tree.find_all(exp.Anonymous):
            if f.name.upper() not in ['YEAR','MONTH','LOWER','UPPER']:
                raise ValueError('Unsupported function '+f.name)
        self.cursor.execute(sql)
        rows=self.cursor.fetchmany(1001)
        if len(rows)>1000: raise ValueError('Result limit exceeded')
        return [tuple(str(x) if not isinstance(x,(str,int,float,type(None))) else x for x in row) for row in rows]
    def close(self): self.cursor.close(); self.conn.close()


def main():
    p=argparse.ArgumentParser(); p.add_argument('--data',required=True); p.add_argument('--predictions')
    p.add_argument('--out',required=True); p.add_argument('--limit',type=int,default=0)
    args=p.parse_args()
    rows=[json.loads(l) for l in Path(args.data).read_text().splitlines()]
    predictions={}
    if args.predictions: predictions={r['id']:r for r in map(json.loads,Path(args.predictions).read_text().splitlines())}
    dbs={b:NativeDB(b) for b in ['mysql','supabase']}; results=[]; seen=set(); start=time.time()
    for r in rows:
        if r['operation'] not in SQL_OPS: continue
        if not args.predictions:
            key=(r['operation'],r['backend'])
            if key in seen: continue
            seen.add(key)
        elif r['id'] not in predictions: continue
        if args.limit and len(results)>=args.limit: break
        db=dbs[r['backend']]
        item={'id':r['id'],'backend':r['backend'],'operation':r['operation'],'gold_executes':False,'prediction_equivalent':None}
        try:
            gold_key=next(k for k in ('query','sql') if k in r['target']['arguments'])
            gold=r['target']['arguments'][gold_key]
            prediction=None
            if args.predictions:
                item['prediction_equivalent']=False
                from tinyquery.evaluate import check_action,parse_action
                action=parse_action(predictions[r['id']]['output'])
                check_action(action,r['context'])
                if action.get('name')!=r['target']['name']: raise ValueError('Wrong tool')
                expected_other={k:v for k,v in r['target']['arguments'].items() if k!=gold_key}
                actual_other={k:v for k,v in action['arguments'].items() if k!=gold_key}
                if expected_other!=actual_other: raise ValueError('Wrong non-SQL arguments')
                prediction=action['arguments'][gold_key]
            equal=True
            for seed in [1,7]:
                db.setup(r['slots'],seed); expected=db.execute(gold)
                if prediction is not None:
                    actual=db.execute(prediction)
                    ordered=r['operation'] in ['sort_asc','sort_desc','top','bottom']
                    if not ordered: expected=sorted(expected,key=repr); actual=sorted(actual,key=repr)
                    equal=equal and actual==expected
            item['gold_executes']=True
            if prediction is not None: item['prediction_equivalent']=equal
        except Exception as exc: item['error']=str(exc)[:300]
        results.append(item)
    for db in dbs.values(): db.close()
    summary={'seconds':time.time()-start,'cases':len(results),'gold_executes':sum(r['gold_executes'] for r in results),
             'prediction_equivalent':sum(r['prediction_equivalent'] is True for r in results) if args.predictions else None,
             'engine_versions':{'mysql':'8.0.46','postgresql':'16.15'},'results':results}
    Path(args.out).write_text(json.dumps(summary,indent=2))
    print(json.dumps({k:v for k,v in summary.items() if k!='results'}),flush=True)
    for r in results:
        if 'error' in r: print(json.dumps(r),flush=True)


if __name__=='__main__': main()