add banbanmini backend
This commit is contained in:
145
talkingq-url/test/import_role_simple.py
Normal file
145
talkingq-url/test/import_role_simple.py
Normal file
@@ -0,0 +1,145 @@
|
||||
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())
|
||||
Reference in New Issue
Block a user