#!/usr/bin/env python3 """ 自动集成few-shot功能到 Text2SQL 系统 用法: python scripts/patch_fewshot.py --enable --top-k 3 --min-rating 7 """ import argparse import re from pathlib import Path def patch_orchestrator( enable: bool = True, top_k: int = 3, min_rating: int = 7, data_path: str = "data/experiences/all_samples.jsonl" ): """ 修改 agents/orchestrator.py 添加few-shot支持 """ orch_path = Path(__file__).parent.parent / "agents" / "orchestrator.py" if not orch_path.exists(): print(f"❌ 文件不存在: {orch_path}") return False content = orch_path.read_text(encoding='utf-8') # 1. 在文件头部添加import import_line = "from utils.fewshot_selector import FewShotSelector" if import_line not in content: # 找到合适的插入点(在其他import之后) last_import = content.rfind("import ") if last_import != -1: insert_pos = content.find("\n", last_import) + 1 content = content[:insert_pos] + import_line + "\n" + content[insert_pos:] print("✅ 添加 import: from utils.fewshot_selector import FewShotSelector") # 2. 修改 __init__ 方法 init_pattern = r'def __init__\([\s\S]*?\):' init_match = re.search(init_pattern, content) if init_match: init_text = content[init_match.start():init_match.end()] # 检查是否已添加few-shot参数 if "fewshot_enabled" not in init_text: # 在参数列表末尾添加(在max_retry之后) new_params = """, # Few-shot 配置 fewshot_enabled: bool = True, fewshot_samples_path: Optional[str] = None, fewshot_top_k: int = 3, fewshot_min_rating: int = 7,""" # 找到max_retry参数后插入 max_retry_pos = init_text.find("max_retry:") if max_retry_pos != -1: comma_pos = init_text.find(",", max_retry_pos) + 1 new_init = ( init_text[:comma_pos] + new_params + init_text[comma_pos:] ) content = content[:init_match.start()] + new_init + content[init_match.end():] print("✅ 修改 __init__: 添加few-shot参数") # 在__init__方法体内添加few-shot初始化代码 # 找到"self.max_retry = max_retry"这一行 max_retry_assign = "self.max_retry = max_retry" if max_retry_assign in content[init_match.start():init_match.end()]: init_body_start = init_match.end() init_body = content[init_body_start:init_body_start + 2000] # 读取前2000字符 if "self.fewshot_selector" not in init_body: fewshot_init = ''' # Few-shot 初始化 self.fewshot_enabled = fewshot_enabled self.fewshot_top_k = fewshot_top_k self.fewshot_min_rating = fewshot_min_rating self.fewshot_selector = None if self.fewshot_enabled: try: from utils.fewshot_selector import FewShotSelector path = fewshot_samples_path or os.getenv( "FEWSHOT_DATA_PATH", "./data/experiences/all_samples.jsonl" ) self.fewshot_selector = FewShotSelector(path) logger.info( f"Few-shot已启用: top_k={fewshot_top_k}, " f"min_rating={fewshot_min_rating}" ) except Exception as e: logger.warning(f"Few-shot加载失败: {e},将使用标准生成") self.fewshot_enabled = False ''' # 插入到self.max_retry赋值之后 insert_pos = content[init_match.start():].find(max_retry_assign) if insert_pos != -1: absolute_pos = init_match.start() + insert_pos + len(max_retry_assign) + 1 content = content[:absolute_pos] + fewshot_init + content[absolute_pos:] print("✅ 修改 __init__: 添加few-shot初始化") # 3. 修改 _generate_sql 方法 generate_pattern = r'def _generate_sql\([\s\S]*?\) -> str:' generate_match = re.search(generate_pattern, content) if generate_match: generate_text = content[generate_match.start():generate_match.end()] if "if self.fewshot_enabled" not in generate_text: # 在方法开头注入few-shot逻辑 fewshot_logic = ''' # Few-shot 增强 if self.fewshot_enabled and self.fewshot_selector: try: examples = self.fewshot_selector.select( question=question, top_k=self.fewshot_top_k, min_rating=self.fewshot_min_rating ) if examples: examples_prompt = "\\n".join([ f"示例 {i+1}:\\n问题:{ex.question_zh}\\nSQL:\\n{ex.sql}" for i, ex in enumerate(examples) ]) examples_prompt = ( f"参考以下相似问题的SQL示例:\\n\\n{examples_prompt}\\n\\n" f"请学习上述示例的SQL编写风格和模式,为当前问题生成SQL。" ) # 注入到schema_str前或user_content中 # 这里选择在schema_str前添加 schema_str = f"{examples_prompt}Schema信息:\\n{schema_str}" logger.debug(f"已注入 {len(examples)} 个few-shot示例") except Exception as e: logger.warning(f"Few-shot检索失败: {e}") ''' # 插入到方法开头("dialect_label = dialect" 之前) dialect_pos = content[generate_match.start():].find("dialect_label = dialect") if dialect_pos != -1: absolute_pos = generate_match.start() + dialect_pos content = content[:absolute_pos] + fewshot_logic + content[absolute_pos:] print("✅ 修改 _generate_sql: 注入few-shot逻辑") # 保存修改 backup_path = orch_path.with_suffix('.py.bak') orch_path.write_text(content, encoding='utf-8') print(f"✅ 已修改 {orch_path}") print(f" 备份文件: {backup_path}") return True def create_env_config(top_k: int, min_rating: int): """更新.env文件添加few-shot配置""" env_path = Path(__file__).parent.parent / ".env" if not env_path.exists(): print(f"⚠️ .env文件不存在: {env_path}") return content = env_path.read_text(encoding='utf-8') additions = f""" # Few-shot 配置 FEWSHOT_ENABLED=true FEWSHOT_TOP_K={top_k} FEWSHOT_MIN_RATING={min_rating} FEWSHOT_DATA_PATH=./data/experiences/all_samples.jsonl """ # 检查是否已存在 if "FEWSHOT_ENABLED" not in content: with env_path.open('a', encoding='utf-8') as f: f.write(additions) print(f"✅ 已添加Few-shot配置到 .env") else: print(f"ℹ️ .env 中已存在FEWSHOT配置,跳过") def main(): parser = argparse.ArgumentParser(description="集成Few-shot功能到Text2SQL系统") parser.add_argument("--enable", action="store_true", help="启用few-shot") parser.add_argument("--top-k", type=int, default=3, help="每次使用的示例数量") parser.add_argument("--min-rating", type=int, default=7, help="示例最低评分") parser.add_argument("--env-only", action="store_true", help="仅修改.env") args = parser.parse_args() if not args.enable and not args.env_only: print("请使用 --enable 启用few-shot,或 --env-only 仅修改环境变量") return print("=" * 60) print("Text2SQL Few-shot 集成工具") print("=" * 60) if args.env_only: create_env_config(args.top_k, args.min_rating) return if args.enable: print(f"\n配置: top_k={args.top_k}, min_rating={args.min_rating}") # 1. 修改orchestrator.py success = patch_orchestrator( enable=True, top_k=args.top_k, min_rating=args.min_rating ) if success: # 2. 更新.env create_env_config(args.top_k, args.min_rating) print("\n" + "=" * 60) print("✅ Few-shot集成完成!") print("=" * 60) print(f"\n下一步:") print(f"1. 确保经验数据集已生成:") print(f" python scripts/parse_examples.py") print(f"\n2. 测试系统:") print(f" python main.py \"查询2024年1月的销售额\" --verbose") print(f"\n3. 查看日志中的few-shot选择:") print(f" Few-shot已选择: Q31, Q6, Q44") print(f"\n4. 调整参数:") print(f" 编辑 .env 修改 FEWSHOT_TOP_K 和 FEWSHOT_MIN_RATING") else: print("❌ 集成失败") if __name__ == "__main__": main()