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