"""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()