145 lines
6.8 KiB
Python
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()) |