File size: 15,996 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 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 | """Build fictitious tool/SQL data with semantic labels and grouped held-out domains."""
import argparse
import hashlib
import json
import random
import sqlite3
from pathlib import Path
from tinyquery.recipes import RECIPES, LANGUAGES
DOMAINS = {
'train': ['orders','products','employees','books','payments','students','invoices','tickets',
'shipments','courses','devices','expenses','appointments','subscriptions','accounts','visits'],
'validation': ['rentals','repairs','enrollments','deliveries'],
'test': ['specimens','exhibits','reservations','inspections'],
}
TEXT_COLS=['name','title','description','customer_name','item_name','display_name']
NUM_COLS=['amount','price','salary','score','cost','quantity','balance','rating']
CAT_COLS=['category','city','status','department','region','kind']
DATE_COLS=['created_at','order_date','start_date','recorded_at']
VALUES=['Delhi','Mumbai','Pune','active','pending','Books','Electronics','North','South',"O'Reilly"]
SQL_OPS=list(RECIPES)[:33]
def compact(value): return json.dumps(value,ensure_ascii=False,separators=(',',':'))
def serialize(context,question):
return '<|context|>'+compact(context)+'\n<|user|>'+question+'\n<|assistant|>'
def obj_schema(properties,required=None):
return {'type':'object','properties':properties,'required':list(properties) if required is None else required,
'additionalProperties':False}
def tool(name,description,properties,required=None):
return {'name':name,'description':description,'inputSchema':obj_schema(properties,required)}
def make_tools(backend,rng,split,project_id):
prefix=rng.choice(['','db.','database.']) if split=='train' else ('warehouse.' if split=='test' else 'store.')
string={'type':'string'}
query_key='query' if backend=='supabase' else rng.choice(['sql','query'])
query_name=prefix+('execute_sql' if backend=='supabase' else rng.choice(['query','run_query','execute_sql']))
scoped=backend!='supabase' or rng.random()<0.5
project={} if scoped else {'project_id':string}
tools=[tool(query_name,'Execute a read-only '+('PostgreSQL' if backend=='supabase' else 'MySQL')+' SELECT query.',
{**project,query_key:string}),
tool(prefix+'list_tables','List database tables and their columns.',
{**project,'schemas':{'type':'array','items':string}},list(project))]
if backend=='mysql':
tools.append(tool(prefix+'describe_table','Inspect columns and types for one table.',
{**project,'table':string}))
return tools,query_key,project,scoped
def schema_for(domain,index):
h=int(hashlib.sha256(domain.encode()).hexdigest()[:8],16)
return {'table':domain,'column':TEXT_COLS[(h+index)%len(TEXT_COLS)],
'numeric':NUM_COLS[(h+index)%len(NUM_COLS)],'category':CAT_COLS[(h+index)%len(CAT_COLS)],
'date_column':DATE_COLS[(h+index)%len(DATE_COLS)],'parent':domain+'_groups',
'foreign':'group_id','label':'label'}
def ddls(s):
return [f"CREATE TABLE {s['parent']} (id INTEGER PRIMARY KEY, {s['label']} TEXT);",
f"CREATE TABLE {s['table']} (id INTEGER PRIMARY KEY, {s['column']} TEXT, {s['numeric']} REAL, "
f"{s['category']} TEXT, {s['date_column']} DATE, {s['foreign']} INTEGER REFERENCES {s['parent']}(id));"]
def quote(value): return "'"+str(value).replace("'","''")+"'"
def gold_sql(op,s,backend):
t,c,n,g,d=s['table'],s['column'],s['numeric'],s['category'],s['date_column']
v,o=quote(s['value']),quote(s['other'])
x,y,k=s['number'],s['upper'],s['limit']
base=f'SELECT * FROM {t}'
queries={
'all':base,'project':f'SELECT {c} FROM {t}',
'eq':f'{base} WHERE {g} = {v}','project_eq':f'SELECT {c} FROM {t} WHERE {g} = {v}',
'gt':f'{base} WHERE {n} > {x}','lt':f'{base} WHERE {n} < {x}',
'gte':f'{base} WHERE {n} >= {x}','lte':f'{base} WHERE {n} <= {x}',
'between':f'{base} WHERE {n} BETWEEN {x} AND {y}',
'and':f'{base} WHERE {g} = {v} AND {n} > {x}',
'or':f'{base} WHERE {g} IN ({v}, {o})',
'null':f'{base} WHERE {c} IS NULL','not_null':f'{base} WHERE {c} IS NOT NULL',
'count':f'SELECT COUNT(*) FROM {t}','count_eq':f'SELECT COUNT(*) FROM {t} WHERE {g} = {v}',
'sum':f'SELECT SUM({n}) FROM {t}','avg':f'SELECT AVG({n}) FROM {t}',
'max':f'SELECT MAX({n}) FROM {t}','min':f'SELECT MIN({n}) FROM {t}',
'distinct':f'SELECT DISTINCT {g} FROM {t}',
'sort_desc':f'{base} ORDER BY {n} DESC','sort_asc':f'{base} ORDER BY {n} ASC',
'top':f'{base} ORDER BY {n} DESC LIMIT {k}','bottom':f'{base} ORDER BY {n} ASC LIMIT {k}',
'group_count':f'SELECT {g}, COUNT(*) FROM {t} GROUP BY {g}',
'group_sum':f'SELECT {g}, SUM({n}) FROM {t} GROUP BY {g}',
'having':f'SELECT {g}, COUNT(*) FROM {t} GROUP BY {g} HAVING COUNT(*) > {x}',
'date_after':f'{base} WHERE {d} >= '+quote(s['date']),
'year':f'{base} WHERE '+(f'YEAR({d})' if backend=='mysql' else f'EXTRACT(YEAR FROM {d})')+f" = {s['year']}",
'month':f'{base} WHERE '+(f'MONTH({d})' if backend=='mysql' else f'EXTRACT(MONTH FROM {d})')+f" = {s['month']}",
'contains':f'{base} WHERE LOWER({c}) LIKE LOWER('+quote('%'+s['value']+'%')+')',
'join':f"SELECT {t}.{c}, {s['parent']}.{s['label']} FROM {t} JOIN {s['parent']} ON {t}.{s['foreign']} = {s['parent']}.id",
'join_filter':f"SELECT {t}.{c} FROM {t} JOIN {s['parent']} ON {t}.{s['foreign']} = {s['parent']}.id WHERE {s['parent']}.{s['label']} = {v}",
}
return queries[op]+';'
def fixture(s,seed=0):
rng=random.Random(seed)
db=sqlite3.connect(':memory:')
db.create_function('YEAR',1,lambda d: None if d is None else int(str(d)[:4]))
db.create_function('MONTH',1,lambda d: None if d is None else int(str(d)[5:7]))
for ddl in ddls(s): db.execute(ddl)
db.executemany(f"INSERT INTO {s['parent']} VALUES (?,?)",[(1,s['value']),(2,s['other']),(3,'West')])
vals=[s['number']-1,s['number'],s['number']+1,s['upper'],s['upper']+1]
rows=[]
for i in range(24):
rows.append((i+1, None if i%9==0 else rng.choice(['Asha','Ravi','item '+s['value'],'Other']),
(vals[i%len(vals)] if i<6 else rng.choice(vals)+i*0.001) if i%6 else None,
rng.choice([s['value'],s['other'],'Other',None]),
rng.choice([s['date'],f"{s['year']}-01-01",'2022-07-18',None]),rng.choice([1,2,3,None])))
db.executemany(f"INSERT INTO {s['table']} VALUES (?,?,?,?,?,?)",rows)
db.execute('PRAGMA query_only=ON')
return db
def validate_sql(sql,backend,s):
import sqlglot
from sqlglot import exp
parsed=sqlglot.parse(sql,read='postgres' if backend=='supabase' else 'mysql')
if len(parsed)!=1 or not isinstance(parsed[0],exp.Select): raise ValueError('Not one SELECT')
result=[]
for seed in (1,7):
db=fixture(s,seed)
translated=sqlite_sql(parsed[0])
try: result.append(db.execute(translated).fetchall())
except Exception as exc: raise ValueError(f'{sql} -> {translated}: {exc}') from exc
finally: db.close()
return result
def validate_context_sql(sql,backend,schema):
"""Compile against the supplied fictitious schema during offline evaluation/auditing."""
import sqlglot
db=sqlite3.connect(':memory:');db.create_function('YEAR',1,lambda x:0);db.create_function('MONTH',1,lambda x:0)
try:
for ddl in schema:db.execute(ddl)
db.execute('PRAGMA query_only=ON')
db.execute(sqlite_sql(sqlglot.parse_one(sql,read='postgres' if backend=='supabase' else 'mysql')))
finally:db.close()
def sqlite_sql(tree):
from sqlglot import exp
def rewrite(node):
if isinstance(node,exp.Extract):
unit=str(node.this).upper()
if unit in ('YEAR','MONTH'):
return exp.Cast(this=exp.Anonymous(this='STRFTIME',expressions=[
exp.Literal.string('%Y' if unit=='YEAR' else '%m'),node.expression.copy()]),
to=exp.DataType.build('INTEGER'))
return node
return tree.copy().transform(rewrite).sql(dialect='sqlite')
def load_templates(path):
templates={}
for line in Path(path).read_text().splitlines():
row=json.loads(line); templates[row['op']]=row['templates']
missing=set(RECIPES)-set(templates)
if missing: raise ValueError(f'Missing teacher recipes: {sorted(missing)}')
return templates
def choose_template(templates,op,lang,split,rng):
pool=templates[op][lang]
if len(pool)<4: raise ValueError(f'Need four templates for {op}/{lang}, have {len(pool)}')
indices=list(range(len(pool)-2)) if split=='train' else [len(pool)-(2 if split=='validation' else 1)]
idx=rng.choice(indices)
return pool[idx],idx
def build_record(op,domain,index,backend,split,templates,rng):
s=schema_for(domain,index)
number=rng.choice([0,1,2,3,5,10,20,50,100,200,500,1000,1500,5000,10000])
s.update(number=number,upper=number+rng.choice([10,50,100,500]),limit=rng.choice([1,3,5,10,20]),
value=rng.choice(VALUES),other=rng.choice(VALUES),date=f'202{rng.randrange(2,7)}-{rng.randrange(1,13):02d}-15',
year=rng.randrange(2021,2027),month=rng.randrange(1,13),city=rng.choice(['Delhi','Mumbai','Lucknow','Jaipur']),
unit=rng.choice(['celsius','fahrenheit']),query=rng.choice(['row level security','SQL joins','API authentication','database backups','indexes']),
path=rng.choice(['/docs/README.md','/reports/sales.csv','/notes/setup.txt','/config/app.json']),
ticket_id='TKT-'+str(rng.randrange(100,10000)),count=rng.randrange(0,150))
project_id='demo_'+domain
tools,key,project,scoped=make_tools(backend,rng,split,project_id)
project_args={k:project_id for k in project}
context={'backend':backend,'project_id':project_id,'schema':ddls(s), 'tools':tools,
'policy':'Read-only database access. Use provided tools. Ask when required information is missing.'}
if op in SQL_OPS:
sql=gold_sql(op,s,backend)
validate_sql(sql,backend,s)
target={'action':'call','name':tools[0]['name'],'arguments':{**project_args,key:sql}}
elif op in ('list_tables','missing_schema'):
if op=='missing_schema': context['schema']=[]
target={'action':'call','name':tools[1]['name'],'arguments':{**project_args,'schemas':['public']}}
elif op in ('describe','sql_error'):
if op=='sql_error':
context['history']=[{'role':'tool','isError':True,'content':'Unknown column in previous query. Inspect the schema.'}]
context['schema']=[]
target=({'action':'call','name':tools[2]['name'],'arguments':{**project_args,'table':domain}}
if backend=='mysql' else
{'action':'call','name':tools[1]['name'],'arguments':{**project_args,'schemas':['public']}})
elif op in ('ambiguous','missing_value','write'):
target={'action':'clarify','question':{'ambiguous':'What does best mean: which column and order?',
'missing_value':f"Which value should {s['category']} equal?",'write':'This connection is read-only. Would you like a SELECT query instead?'}[op]}
elif op=='final':
context['history']=[{'role':'tool','content':{'row_count':s['count']}}]
target={'action':'answer','text':f"The tool returned {s['count']} rows."}
else:
definitions={
'weather':('get_weather','Get current weather for a city.',{'city':{'type':'string'},'unit':{'type':'string','enum':['celsius','fahrenheit']}}, {'city':s['city'],'unit':s['unit']}),
'search':('search_docs','Search documentation for a query.',{'query':{'type':'string'}},{'query':s['query']}),
'read_file':('read_file','Read text from a file path.',{'path':{'type':'string'}},{'path':s['path']}),
'ticket':('get_ticket','Retrieve a support ticket by its ID.',{'ticket_id':{'type':'string'}},{'ticket_id':s['ticket_id']}),
}
generic=[]
for kind,(name,desc,properties,args) in definitions.items():
prefix=rng.choice(['','tools.','app.']) if split=='train' else ('external.' if split=='test' else 'service.')
generic.append(tool(prefix+name,desc,properties))
if kind==op: target={'action':'call','name':prefix+name,'arguments':args}
tools.extend(generic)
context['schema']=[]
rng.shuffle(tools)
semantic_id=hashlib.sha256(compact([split,backend,s,op,target]).encode()).hexdigest()[:20]
records=[]
for lang in LANGUAGES:
template,template_index=choose_template(templates,op,lang,split,rng)
question=template.format(**s)
answer=dict(target)
if lang=='hi' and answer['action']=='clarify':
answer['question']={'ambiguous':'सबसे अच्छा से आपका क्या मतलब है? कौन सा कॉलम और क्रम?',
'missing_value':f"{s['category']} का कौन सा मान चाहिए?",'write':'यह कनेक्शन केवल पढ़ने के लिए है। क्या आपको SELECT क्वेरी चाहिए?'}[op]
elif lang=='hinglish' and answer['action']=='clarify':
answer['question']={'ambiguous':'Best ka matlab kya hai? Kaunsa column aur order?',
'missing_value':f"{s['category']} ki kaunsi value chahiye?",'write':'Yeh connection read-only hai. Kya SELECT query chahiye?'}[op]
elif answer['action']=='answer':
if lang=='hi': answer['text']=f"टूल ने {s['count']} पंक्तियाँ लौटाईं।"
elif lang=='hinglish': answer['text']=f"Tool ne {s['count']} rows return ki."
records.append({'id':semantic_id+'_'+lang,'scenario_id':semantic_id,'split':split,'domain':domain,
'language':lang,'backend':backend,'operation':op,'slots':s,'context':context,
'question':question,'target':answer,'prompt':serialize(context,question),
'response':compact(answer),'template_index':template_index,
'provenance':'programmatic semantics + Qwen3.8-27B-FP8 language template'})
return records
def main():
p=argparse.ArgumentParser(); p.add_argument('--templates',required=True)
p.add_argument('--out',default='data/tinyquery'); p.add_argument('--scenarios',type=int,default=18000)
args=p.parse_args(); out=Path(args.out); out.mkdir(parents=True,exist_ok=True)
templates=load_templates(args.templates); stats={}; global_prompts={}
for split,count in [('train',args.scenarios),('validation',300),('test',300)]:
rng=random.Random({'train':42,'validation':43,'test':44}[split]); seen=set(); counts={}
with (out/(split+'.jsonl')).open('w') as stream:
for i in range(count):
op=list(RECIPES)[i%len(RECIPES)]; domain=rng.choice(DOMAINS[split])
for record in build_record(op,domain,i,rng.choice(['mysql','supabase']),split,templates,rng):
key=hashlib.sha256(record['prompt'].encode()).hexdigest()
if key in seen: continue
if key in global_prompts and global_prompts[key]!=split: raise ValueError('Cross-split prompt leakage')
seen.add(key); global_prompts[key]=split
stream.write(json.dumps(record,ensure_ascii=False)+'\n')
counts[op]=counts.get(op,0)+1
if i%1000==0: print(split,i,flush=True)
stats[split]={'examples':len(seen),'operations':counts,'domains':DOMAINS[split]}
stats['sql_validation_method']='Dialect parsing plus transpiled SQLite execution on two fixtures; not native MySQL/Postgres execution.'
(out/'dataset-stats.json').write_text(json.dumps(stats,indent=2))
print(json.dumps(stats),flush=True)
if __name__=='__main__': main()
|