| """Transparent TF-IDF nearest-question baseline; copies the retrieved action unchanged.""" |
| import argparse |
| from collections import Counter,defaultdict |
| import json |
| import math |
| from pathlib import Path |
| import re |
| import time |
| import numpy as np |
| from tinyquery.evaluate import score,aggregate |
|
|
|
|
| def bind_context(text,source,destination): |
| """Bind by tool descriptions and DDL positions, using no expected answer or slots.""" |
| import sqlglot |
| from sqlglot import exp |
| action=json.loads(text) |
| if action.get('action')!='call':return text |
| old=next(t for t in source['tools'] if t['name']==action['name']) |
| normalize=lambda s:s.replace('PostgreSQL','SQL').replace('MySQL','SQL') |
| candidates=[t for t in destination['tools'] if normalize(t['description'])==normalize(old['description'])] |
| if not candidates:return text |
| tool=candidates[0];props=tool['inputSchema']['properties'];args=dict(action['arguments']) |
| sqlkey=next((k for k in ['sql','query'] if isinstance(args.get(k),str) and args[k].startswith('SELECT ')),None) |
| if sqlkey: |
| targetkey=next((k for k in ['sql','query'] if k in props),sqlkey) |
| tables={};columns={} |
| for old_ddl,new_ddl in zip(source['schema'],destination['schema']): |
| a=sqlglot.parse_one(old_ddl).this;b=sqlglot.parse_one(new_ddl).this |
| tables[a.this.name]=b.this.name |
| for c,d in zip(a.expressions,b.expressions): |
| if isinstance(c,exp.ColumnDef) and isinstance(d,exp.ColumnDef):columns[c.name]=d.name |
| query=sqlglot.parse_one(args.pop(sqlkey),read='postgres' if source['backend']=='supabase' else 'mysql') |
| def rename(node): |
| if isinstance(node,exp.Table) and node.name in tables:node.set('this',exp.to_identifier(tables[node.name])) |
| if isinstance(node,exp.Column): |
| if node.name in columns:node.set('this',exp.to_identifier(columns[node.name])) |
| if node.table in tables:node.set('table',exp.to_identifier(tables[node.table])) |
| return node |
| args[targetkey]=query.transform(rename).sql(dialect='postgres' if destination['backend']=='supabase' else 'mysql')+';' |
| args={k:v for k,v in args.items() if k in props} |
| if 'project_id' in props:args['project_id']=destination['project_id'] |
| return json.dumps({'action':'call','name':tool['name'],'arguments':args},ensure_ascii=False,separators=(',',':')) |
|
|
|
|
| def words(text): |
| return re.findall(r'[^\W_]+|_',text.lower(),re.UNICODE) |
|
|
|
|
| def main(): |
| p=argparse.ArgumentParser(); p.add_argument('--train',required=True); p.add_argument('--data',required=True) |
| p.add_argument('--out',required=True);p.add_argument('--bind-context',action='store_true') |
| args=p.parse_args(); start=time.time() |
| training=[]; seen=set(); counts=[]; df=Counter() |
| for line in Path(args.train).open(): |
| row=json.loads(line) |
| if row['question'] in seen: continue |
| seen.add(row['question']); terms=Counter(words(row['question'])) |
| training.append((row['id'],row['response'],row['context'] if args.bind_context else None)); counts.append(terms); df.update(terms.keys()) |
| n=len(training); idf={word:math.log((1+n)/(1+freq))+1 for word,freq in df.items()} |
| postings=defaultdict(list) |
| for i,terms in enumerate(counts): |
| weighted={word:(1+math.log(freq))*idf[word] for word,freq in terms.items()} |
| norm=math.sqrt(sum(v*v for v in weighted.values())) or 1 |
| for word,value in weighted.items(): postings[word].append((i,value/norm)) |
| postings={word:(np.array([i for i,_ in pairs]),np.array([v for _,v in pairs],dtype=np.float32)) |
| for word,pairs in postings.items()} |
| results=[]; groups=defaultdict(list); out=Path(args.out); out.parent.mkdir(parents=True,exist_ok=True) |
| with out.open('w') as stream: |
| for line in Path(args.data).open(): |
| row=json.loads(line); similarities=np.zeros(n,dtype=np.float32) |
| for word,freq in Counter(words(row['question'])).items(): |
| if word in postings: |
| indices,weights=postings[word]; similarities[indices]+=(1+math.log(freq))*idf[word]*weights |
| index=int(similarities.argmax()); source,text,context=training[index] |
| if args.bind_context: |
| try:text=bind_context(text,context,row['context']) |
| except (ValueError,KeyError,StopIteration,AttributeError):pass |
| metrics=score(row,text) |
| results.append(metrics) |
| for field in ['language','backend','operation']: groups[field+':'+row[field]].append(metrics) |
| stream.write(json.dumps({'id':row['id'],'retrieved_id':source,'output':text,'metrics':metrics},ensure_ascii=False)+'\n') |
| report={'baseline':'TF-IDF nearest training question; copy action unchanged, no tool/schema adaptation', |
| 'training_questions':n,'examples':len(results),'seconds':time.time()-start,'metrics':aggregate(results), |
| 'groups':{key:aggregate(value) for key,value in groups.items()}} |
| if args.bind_context:report['baseline']='TF-IDF nearest question plus tool-description and DDL-position binding; no expected answer or gold slots used' |
| out.with_suffix('.summary.json').write_text(json.dumps(report,indent=2)); print(json.dumps(report['metrics'])) |
|
|
|
|
| if __name__=='__main__': main() |
|
|