import asyncio import yaml import sys import os sys.path.append(os.path.dirname(os.path.abspath(__file__))) from database.connection import get_db_manager from database.models import Role, RoleLanguage from utils.logger import session_logger from sqlalchemy import select, delete async def import_role_from_yaml(yaml_file_path): """从YAML文件导入角色到数据库""" try: with open(yaml_file_path, 'r', encoding='utf-8') as file: role_data = yaml.safe_load(file) db_manager = await get_db_manager() session = await db_manager.get_session() try: role_key = os.path.basename(yaml_file_path).split('.')[0].lower() existing_role = await session.execute(select(Role).where(Role.role_key == role_key)) existing_role = existing_role.scalars().first() if existing_role: session_logger.info("system", "import", f"角色 {role_key} 已存在,将更新现有角色") existing_role.name = role_data.get('name', '') existing_role.default_language = role_data.get('default_language', 'zh') existing_role.asr_provider = role_data.get('asr_provider') existing_role.llm_provider = role_data.get('llm_provider') existing_role.tts_provider = role_data.get('tts_provider') existing_role.competitive_llm_mode = role_data.get('competitive_llm_mode') existing_role.volcano_model_id = role_data.get('volcano_model_id') existing_role.volcano_voice_type = role_data.get('volcano_voice_type') existing_role.tencent_voice_type = role_data.get('tencent_voice_type') existing_role.aliyun_voice_name = role_data.get('aliyun_voice_name') existing_role.minimax_voice_id = role_data.get('minimax_voice_id') existing_role.homophones = role_data.get('homophones') default_lang = role_data.get('default_language', 'zh') if default_lang in role_data.get('multilingual', {}): existing_role.content = role_data['multilingual'][default_lang].get('content', '') if 'description' in role_data['multilingual'][default_lang]: existing_role.description = role_data['multilingual'][default_lang]['description'] if 'url' in role_data['multilingual'][default_lang]: existing_role.url = role_data['multilingual'][default_lang]['url'] await session.commit() role_id = existing_role.id session_logger.info("system", "import", f"已更新角色 {role_key} 的基本信息") await session.execute(delete(RoleLanguage).where(RoleLanguage.role_id == role_id)) await session.commit() session_logger.info("system", "import", f"已删除角色 {role_key} 的现有语言配置") else: default_lang = role_data.get('default_language', 'zh') content = "" description = None url = None if default_lang in role_data.get('multilingual', {}): content = role_data['multilingual'][default_lang].get('content', '') description = role_data['multilingual'][default_lang].get('description') url = role_data['multilingual'][default_lang].get('url') new_role = Role( role_key=role_key, name=role_data.get('name', ''), description=description, content=content, default_language=default_lang, asr_provider=role_data.get('asr_provider'), llm_provider=role_data.get('llm_provider'), tts_provider=role_data.get('tts_provider'), competitive_llm_mode=role_data.get('competitive_llm_mode'), volcano_model_id=role_data.get('volcano_model_id'), volcano_voice_type=role_data.get('volcano_voice_type'), tencent_voice_type=role_data.get('tencent_voice_type'), aliyun_voice_name=role_data.get('aliyun_voice_name'), minimax_voice_id=role_data.get('minimax_voice_id'), url=url, homophones=role_data.get('homophones'), enabled=True ) session.add(new_role) await session.commit() role_id = new_role.id session_logger.info("system", "import", f"已创建角色 {role_key} 的基本信息") for lang_code, lang_data in role_data.get('multilingual', {}).items(): new_lang = RoleLanguage( role_id=role_id, language_code=lang_code, name=lang_data.get('name'), content=lang_data.get('content'), # asr_provider=lang_data.get('asr_provider'), # llm_provider=lang_data.get('llm_provider'), # tts_provider=lang_data.get('tts_provider'), # volcano_voice_type=lang_data.get('volcano_voice_type'), # tencent_voice_type=lang_data.get('tencent_voice_type'), # aliyun_voice_name=lang_data.get('aliyun_voice_name'), minimax_voice_id=lang_data.get('minimax_voice_id'), url=lang_data.get('url') ) session.add(new_lang) await session.commit() session_logger.info("system", "import", f"已导入角色 {role_key} 的多语言配置") return True, f"成功导入角色 {role_key}" except Exception as e: await session.rollback() session_logger.error("system", "import", f"导入角色时发生错误: {str(e)}") return False, f"导入失败: {str(e)}" finally: await session.close() await db_manager.close() except Exception as e: session_logger.error("system", "import", f"处理YAML文件时发生错误: {str(e)}") return False, f"处理YAML文件失败: {str(e)}" async def main(): if len(sys.argv) != 2: print("用法: python import_role.py ") return yaml_file_path = sys.argv[1] success, message = await import_role_from_yaml(yaml_file_path) if success: print(f"成功: {message}") else: print(f"错误: {message}") sys.exit(1) if __name__ == "__main__": asyncio.run(main())