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

225 lines
8.8 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import os
from typing import Dict, Optional, List
import asyncio
import time
from config import settings
from utils.logger import session_logger
from services.role_validator import role_validator
from sqlalchemy import select, and_
from services.database_service_base import DatabaseServiceBase
from database.models import Role, RoleLanguage
class RoleManager(DatabaseServiceBase):
def __init__(self):
super().__init__(service_name="role_manager")
self.roles: Dict[str, Dict[str, str]] = {}
self.invalid_roles: Dict[str, List[str]] = {}
self.lock = asyncio.Lock()
self._initialized = False
self.last_refresh_time = 0
self.refresh_interval = 300 # 缓存刷新间隔,单位秒
self.etags = {} # 角色配置的ETag缓存
async def initialize(self):
if self._initialized:
return
worker_id = os.environ.get("UVICORN_WID", "0")
is_main_process = worker_id == "0"
async with self.lock:
if self._initialized: # 双重检查锁定
return
try:
await self._init_database()
await self._load_roles_from_db()
self._initialized = True
self.last_refresh_time = time.time()
if is_main_process:
session_logger.info("system", "role_manager", "角色管理器初始化完成")
except Exception as e:
if is_main_process:
session_logger.error("system", "role_manager", f"角色管理器初始化失败: {str(e)}")
async def _maybe_refresh_roles(self):
"""如果缓存过期,刷新角色配置"""
current_time = time.time()
if current_time - self.last_refresh_time > self.refresh_interval:
await self.reload_roles()
async def _load_roles_from_db(self):
"""从数据库加载所有角色配置"""
self.roles = {}
db_session = await self.get_session()
try:
roles_query = select(Role).where(Role.enabled == True)
roles_result = await db_session.execute(roles_query)
db_roles = roles_result.scalars().all()
new_etags = {}
for db_role in db_roles:
role_dict = self._role_to_dict(db_role)
role_key = db_role.role_key.lower()
languages_query = select(RoleLanguage).where(RoleLanguage.role_id == db_role.id)
languages_result = await db_session.execute(languages_query)
lang_configs = languages_result.scalars().all()
if lang_configs:
multilingual = {}
for lang_config in lang_configs:
lang_dict = self._role_lang_to_dict(lang_config)
multilingual[lang_config.language_code] = lang_dict
role_dict["multilingual"] = multilingual
self.roles[role_key] = role_dict
import hashlib
import json
role_json = json.dumps(role_dict, sort_keys=True)
etag = hashlib.md5(role_json.encode()).hexdigest()
new_etags[role_key] = etag
self.etags = new_etags
except Exception as e:
session_logger.error("system", "role_manager", f"从数据库加载角色配置失败: {str(e)}")
raise
finally:
await db_session.close()
def _role_to_dict(self, db_role: Role) -> Dict:
"""将数据库Role对象转换为字典格式"""
role_dict = {
"role_key": db_role.role_key,
"name": db_role.name,
"content": db_role.content,
}
if db_role.description:
role_dict["description"] = db_role.description
if db_role.default_language:
role_dict["default_language"] = db_role.default_language
if db_role.volcano_model_id:
role_dict["volcano_model_id"] = db_role.volcano_model_id
if db_role.minimax_voice_id:
role_dict["minimax_voice_id"] = db_role.minimax_voice_id
if db_role.url:
role_dict["url"] = db_role.url
if db_role.homophones:
role_dict["homophones"] = db_role.homophones
return role_dict
def _role_lang_to_dict(self, lang_config: RoleLanguage) -> Dict:
"""将数据库RoleLanguage对象转换为字典格式"""
lang_dict = {}
if lang_config.name:
lang_dict["name"] = lang_config.name
if lang_config.content:
lang_dict["content"] = lang_config.content
if lang_config.minimax_voice_id:
lang_dict["minimax_voice_id"] = lang_config.minimax_voice_id
if lang_config.url:
lang_dict["url"] = lang_config.url
return lang_dict
async def get_role(self, role_key: str) -> Optional[Dict[str, str]]:
if not self._initialized:
await self.initialize()
await self._maybe_refresh_roles()
return self.roles.get(role_key.lower())
async def get_role_etag(self, role_key: str) -> Optional[str]:
"""获取角色配置的ETag"""
if not self._initialized:
await self.initialize()
return self.etags.get(role_key.lower())
async def get_role_config_for_language(
self, role_key: str, language: str = None
) -> Optional[Dict[str, str]]:
"""
根据指定的语言获取角色配置
Args:
role_key (str): 角色键名
language (str, optional): 语言代码。如果为None使用默认语言
Returns:
Dict: 角色配置
"""
if not self._initialized:
await self.initialize()
await self._maybe_refresh_roles()
role_config = self.roles.get(role_key.lower())
if not role_config:
return None
full_config = role_config.copy()
full_config["role_key"] = role_key
if "multilingual" not in role_config:
return full_config
multilingual = role_config.get("multilingual", {})
default_language = role_config.get("default_language", "zh")
selected_language = None
if language and language in multilingual:
selected_language = language
session_logger.info(
"system",
"role_manager",
f"角色 {role_key} 使用指定的语言: {selected_language}",
)
elif default_language in multilingual:
selected_language = default_language
session_logger.info(
"system",
"role_manager",
f"角色 {role_key} 使用默认语言: {selected_language}",
)
elif multilingual:
selected_language = next(iter(multilingual))
session_logger.info(
"system",
"role_manager",
f"角色 {role_key} 无法使用指定或默认语言,使用第一个可用语言: {selected_language}",
)
else:
return full_config # 如果没有多语言配置,返回原始配置
lang_specific_config = multilingual[selected_language]
for key, value in lang_specific_config.items():
full_config[key] = value
service_keys = ["tts_provider", "llm_provider", "asr_provider"]
for key in service_keys:
if key in lang_specific_config:
full_config[key] = lang_specific_config[key]
full_config["_selected_language"] = selected_language
session_logger.info(
"system",
"role_manager",
f"已加载角色 {role_key}{selected_language} 语言配置",
)
return full_config
async def get_role_errors(self, role_key: str) -> List[str]:
"""获取角色验证错误信息"""
if not self._initialized:
await self.initialize()
await self._maybe_refresh_roles()
return self.invalid_roles.get(role_key.lower(), [])
async def get_all_roles(self) -> Dict[str, Dict[str, str]]:
if not self._initialized:
await self.initialize()
await self._maybe_refresh_roles()
return self.roles
async def reload_roles(self):
"""重新加载所有角色配置"""
async with self.lock:
try:
await self._init_database()
await self._load_roles_from_db()
self.last_refresh_time = time.time()
session_logger.info("system", "role_manager", "角色配置已重新加载")
except Exception as e:
session_logger.error("system", "role_manager", f"重新加载角色配置失败: {str(e)}")
role_manager = RoleManager()