Files
banban/talkingq-url/test/import_role_simple.py
2026-03-24 15:04:36 +08:00

145 lines
6.8 KiB
Python

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