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

197 lines
6.2 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 json
from typing import List, Dict, Any
import asyncio
from dashscope.audio.asr import VocabularyService
class AliyunHotwordManager:
"""阿里云热词管理类"""
def __init__(self, api_key: str = None):
"""初始化热词管理器
Args:
api_key: 阿里云API密钥。如不提供将尝试从环境变量或配置中获取
"""
self.api_key = "sk-7a50eca6856d4afb968ac3bf512f6d1b"
self._service = VocabularyService(api_key=self.api_key)
self.default_prefix = "talkingq"
self.default_model = "gummy-chat-v1" # 默认使用gummy-chat-v1模型
self._default_vocabulary_id = None
async def create_vocabulary(self,
hotwords: List[Dict[str, Any]],
prefix: str = None,
model: str = None) -> str:
"""创建热词表
Args:
hotwords: 热词列表每个热词是一个字典包含text、lang等字段
prefix: 热词表前缀默认使用self.default_prefix
model: 目标模型默认使用self.default_model
Returns:
热词表ID
"""
prefix = prefix or self.default_prefix
model = model or self.default_model
loop = asyncio.get_event_loop()
vocabulary_id = await loop.run_in_executor(
None,
lambda: self._service.create_vocabulary(
target_model=model,
prefix=prefix,
vocabulary=hotwords
)
)
return vocabulary_id
async def list_vocabularies(self, prefix: str = None,
page_index: int = 0,
page_size: int = 10) -> List[Dict]:
"""查询所有热词表
Args:
prefix: 热词表前缀,如果设置则只返回该前缀的热词表
page_index: 页码索引
page_size: 每页大小
Returns:
热词表列表
"""
loop = asyncio.get_event_loop()
result = await loop.run_in_executor(
None,
lambda: self._service.list_vocabularies(
prefix=prefix,
page_index=page_index,
page_size=page_size
)
)
return result
async def query_vocabulary(self, vocabulary_id: str) -> Dict[str, Any]:
"""查询指定热词表内容
Args:
vocabulary_id: 热词表ID
Returns:
热词表内容
"""
loop = asyncio.get_event_loop()
result = await loop.run_in_executor(
None,
lambda: self._service.query_vocabulary(vocabulary_id)
)
return result
async def update_vocabulary(self, vocabulary_id: str,
hotwords: List[Dict[str, Any]]) -> None:
"""更新热词表
Args:
vocabulary_id: 要更新的热词表ID
hotwords: 新的热词列表
"""
loop = asyncio.get_event_loop()
await loop.run_in_executor(
None,
lambda: self._service.update_vocabulary(
vocabulary_id=vocabulary_id,
vocabulary=hotwords
)
)
async def delete_vocabulary(self, vocabulary_id: str) -> None:
"""删除热词表
Args:
vocabulary_id: 要删除的热词表ID
"""
loop = asyncio.get_event_loop()
await loop.run_in_executor(
None,
lambda: self._service.delete_vocabulary(vocabulary_id)
)
async def get_or_create_default_vocabulary(self) -> str:
"""获取或创建默认热词表
如果已经有默认热词表ID直接返回否则创建一个新的热词表
Returns:
热词表ID
"""
if self._default_vocabulary_id:
return self._default_vocabulary_id
vocabularies = await self.list_vocabularies(prefix=self.default_prefix)
if vocabularies and len(vocabularies) > 0:
self._default_vocabulary_id = vocabularies[0].get('vocabulary_id')
return self._default_vocabulary_id
default_hotwords = [
{"text": "变成", "weight": 4, "lang": "zh"},
]
vocabulary_id = await self.create_vocabulary(default_hotwords)
self._default_vocabulary_id = vocabulary_id
return vocabulary_id
async def add_hotwords_to_vocabulary(self, vocabulary_id: str,
new_hotwords: List[Dict[str, Any]]) -> None:
"""向现有热词表添加新热词
Args:
vocabulary_id: 热词表ID
new_hotwords: 要添加的新热词列表
"""
current_vocab = await self.query_vocabulary(vocabulary_id)
current_hotwords = current_vocab.get('vocabulary', [])
updated_hotwords = current_hotwords + new_hotwords
await self.update_vocabulary(vocabulary_id, updated_hotwords)
print(f"成功添加 {len(new_hotwords)} 个热词到热词表 {vocabulary_id}")
hotword_manager = AliyunHotwordManager()
if __name__ == "__main__":
import asyncio
async def main():
vocabulary_id = await hotword_manager.get_or_create_default_vocabulary()
print(f"默认热词表ID: {vocabulary_id}")
vocabulary = await hotword_manager.query_vocabulary(vocabulary_id)
print(f"更新前热词表内容: {json.dumps(vocabulary, ensure_ascii=False, indent=2)}")
new_hotwords = [
{"text": "豚豚崽", "weight": 4, "lang": "zh"},
{"text": "Tuntunzai", "weight": 4, "lang": "en"}
]
await hotword_manager.add_hotwords_to_vocabulary(vocabulary_id, new_hotwords)
updated_vocabulary = await hotword_manager.query_vocabulary(vocabulary_id)
print(f"更新后热词表内容: {json.dumps(updated_vocabulary, ensure_ascii=False, indent=2)}")
asyncio.run(main())