Merge branch 'main' into dev

This commit is contained in:
pengzhisheng
2025-03-08 02:54:17 +08:00
35 changed files with 2286 additions and 1227 deletions
Executable → Regular
View File
+1 -1
View File
@@ -23,7 +23,6 @@
<mybatisplus.version>3.5.5</mybatisplus.version> <mybatisplus.version>3.5.5</mybatisplus.version>
<hutool.version>5.8.24</hutool.version> <hutool.version>5.8.24</hutool.version>
<jsoup.version>1.19.1</jsoup.version> <jsoup.version>1.19.1</jsoup.version>
<jasypt.version>3.0.5</jasypt.version>
<knife4j.version>4.6.0</knife4j.version> <knife4j.version>4.6.0</knife4j.version>
<shiro.version>2.0.2</shiro.version> <shiro.version>2.0.2</shiro.version>
<captcha.version>1.6.2</captcha.version> <captcha.version>1.6.2</captcha.version>
@@ -67,6 +66,7 @@
</exclusion> </exclusion>
</exclusions> </exclusions>
</dependency> </dependency>
<!-- 验证码工具包 -->
<dependency> <dependency>
<groupId>com.github.whvcse</groupId> <groupId>com.github.whvcse</groupId>
<artifactId>easy-captcha</artifactId> <artifactId>easy-captcha</artifactId>
@@ -40,4 +40,6 @@ public interface ErrorCode {
int PASSWORD_LENGTH_ERROR = 10030; int PASSWORD_LENGTH_ERROR = 10030;
int PASSWORD_WEAK_ERROR = 10031; int PASSWORD_WEAK_ERROR = 10031;
int DEL_MYSELF_ERROR = 10032; int DEL_MYSELF_ERROR = 10032;
// 验证码错误
int VERIFICATION_CODE = 10033;
} }
@@ -0,0 +1,11 @@
package xiaozhi.modules.sys.service;
public interface TokenService {
/**
* 生成token
*
* @param userId
* @return
*/
String createToken(long userId);
}
@@ -0,0 +1,31 @@
package xiaozhi.modules.sys.service.impl;
import lombok.AllArgsConstructor;
import org.springframework.stereotype.Service;
import xiaozhi.common.redis.RedisUtils;
import xiaozhi.modules.security.oauth2.TokenGenerator;
import xiaozhi.modules.sys.service.TokenService;
import java.util.Date;
@AllArgsConstructor
@Service
public class TokenServiceImpl implements TokenService {
private final RedisUtils redisUtils;
/**
* 3小时无操作过期
*/
private final static int EXPIRE = 60 * 60 * 3;
@Override
public String createToken(long userId) {
//生成一个token
String token = TokenGenerator.generateValue();
//当前时间
Date now = new Date();
//过期时间
Date expireTime = new Date(now.getTime() + EXPIRE * 1000);
return token;
}
}
@@ -42,8 +42,3 @@ spring:
wall: wall:
config: config:
multi-statement-allow: true multi-statement-allow: true
logging:
level:
org.flowable.engine.impl.persistence.entity.*: debug
org.flowable.task.service.impl.persistence.entity.*: debug
@@ -1,25 +0,0 @@
package xiaozhi;
import org.jasypt.encryption.StringEncryptor;
import org.junit.Test;
import org.junit.runner.RunWith;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.test.context.SpringBootTest;
import org.springframework.test.context.junit4.SpringRunner;
/**
* 单元测试
*/
@RunWith(SpringRunner.class)
@SpringBootTest(webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT)
public class DbEncTest {
@Autowired
StringEncryptor stringEncryptor;
@Test
public void jiami() {
System.out.println("username:" + stringEncryptor.encrypt("07e43e8d669fb946e31ccd4ef5f32c9f2287619c79b766a5985d2c99ad7b7c7e"));
System.out.println("password:" + stringEncryptor.encrypt("042e94093fd2c2765ea45cf13ddbfd38e93026df4b6d5e4206ea5ac90956d63ab73e8b82c6daf7829f9aea7e27e1db5bb0a90944c4c4985af44db0ef49c46d6ad6"));
}
}
@@ -1,27 +0,0 @@
package xiaozhi;
import xiaozhi.common.redis.RedisUtils;
import xiaozhi.modules.sys.entity.SysUserEntity;
import jakarta.annotation.Resource;
import org.apache.commons.lang3.builder.ToStringBuilder;
import org.junit.Test;
import org.junit.runner.RunWith;
import org.springframework.boot.test.context.SpringBootTest;
import org.springframework.test.context.junit4.SpringRunner;
@RunWith(SpringRunner.class)
@SpringBootTest(webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT)
public class RedisTest {
@Resource
private RedisUtils redisUtils;
@Test
public void contextLoads() {
SysUserEntity user = new SysUserEntity();
user.setEmail("123456@qq.com");
redisUtils.set("user", user);
System.out.println(ToStringBuilder.reflectionToString(redisUtils.get("user")));
}
}
@@ -1,24 +0,0 @@
package xiaozhi.service;
import xiaozhi.modules.sys.dao.SysUserDao;
import xiaozhi.modules.sys.entity.SysUserEntity;
import jakarta.annotation.Resource;
import org.springframework.stereotype.Service;
/**
* 测试多数据源
*/
@Service
public class DynamicDataSourceTestService {
@Resource
private SysUserDao sysUserDao;
//@Transactional
public void updateUser(Long id) {
SysUserEntity user = new SysUserEntity();
user.setId(id);
user.setMobile("13500000000");
//sysUserDao.updateById(user);
System.out.println(sysUserDao.selectById(id));
}
}
+1074 -307
View File
File diff suppressed because it is too large Load Diff
+3 -1
View File
@@ -7,11 +7,13 @@
"build": "vue-cli-service build" "build": "vue-cli-service build"
}, },
"dependencies": { "dependencies": {
"axios": "^1.8.1",
"element-ui": "^2.15.14", "element-ui": "^2.15.14",
"flyio": "^0.6.14", "flyio": "^0.6.14",
"normalize.css": "^8.0.1", "normalize.css": "^8.0.1",
"vue": "^2.6.14", "vue": "^2.6.14",
"vue-router": "^3.5.1", "vue-axios": "^3.5.2",
"vue-router": "^3.6.5",
"vuex": "^3.6.2" "vuex": "^3.6.2"
}, },
"devDependencies": { "devDependencies": {
+6 -4
View File
@@ -1,5 +1,6 @@
// 引入各个模块的请求 // 引入各个模块的请求
import user from './module/user.js' import user from './module/user.js'
/** /**
* 接口地址 * 接口地址
* 在开发阶段,如果地址写的是相对路径,请与vue.config.js的devServer配置相结合,方便跨域请求 * 在开发阶段,如果地址写的是相对路径,请与vue.config.js的devServer配置相结合,方便跨域请求
@@ -11,12 +12,13 @@ const DEV_API_SERVICE = 'https://apifoxmock.com/m1/5931378-5618560-default'
* 根据开发环境返回接口url * 根据开发环境返回接口url
* @returns {string} * @returns {string}
*/ */
export function getServiceUrl () { export function getServiceUrl() {
return DEV_API_SERVICE return DEV_API_SERVICE
} }
/** request服务封装 */ /** request服务封装 */
export default { export default {
getServiceUrl, getServiceUrl,
user user,
} }
+27 -1
View File
@@ -16,5 +16,31 @@ export default {
this.login(loginForm, callback) this.login(loginForm, callback)
}) })
}).send() }).send()
} },
// 获取用户信息
getUserInfo(callback) {
RequestService.sendRequest().url(`${getServiceUrl()}/api/v1/user/info`).method('GET')
.success((res) => {
RequestService.clearRequestTime()
callback(res)
})
.fail(() => {
RequestService.reAjaxFun(() => {
this.getUserInfo()
})
}).send()
},
// 获取设备信息
getHomeList(callback) {
RequestService.sendRequest().url(`${getServiceUrl()}/api/v1/user/device/bind`).method('GET')
.success((res) => {
RequestService.clearRequestTime()
callback(res)
})
.fail(() => {
RequestService.reAjaxFun(() => {
this.getUserInfo()
})
}).send()
},
} }
+41 -13
View File
@@ -3,7 +3,7 @@
<el-container style="height: 100%;"> <el-container style="height: 100%;">
<el-header class="header"> <el-header class="header">
<div style="display: flex;justify-content: space-between;"> <div style="display: flex;justify-content: space-between;">
<div style="display: flex;align-items: center;gap: 8px;"> <div style="display: flex;align-items: center;gap: 8px;margin-top: 10px;">
<img src="@/assets/xiaozhi-logo.png" alt="" style="width: 45px;height: 45px;" /> <img src="@/assets/xiaozhi-logo.png" alt="" style="width: 45px;height: 45px;" />
<img src="@/assets/xiaozhi-ai.png" alt="" style="width: 70px;height: 13px;" /> <img src="@/assets/xiaozhi-ai.png" alt="" style="width: 70px;height: 13px;" />
<div class="equipment-management" @click="settingDevice=false"> <div class="equipment-management" @click="settingDevice=false">
@@ -19,7 +19,7 @@
<img src="@/assets/home/close.png" alt="" style="width: 6px;height: 6px;" /> <img src="@/assets/home/close.png" alt="" style="width: 6px;height: 6px;" />
</div> </div>
</div> </div>
<div style="display: flex;align-items: center;gap: 8px;"> <div style="display: flex;align-items: center;gap: 8px;margin-top: 10px">
<div class="serach-box"> <div class="serach-box">
<el-input placeholder="输入名称搜索.." v-model="serach" style="border: none; background: transparent;" /> <el-input placeholder="输入名称搜索.." v-model="serach" style="border: none; background: transparent;" />
<img src="@/assets/home/search.png" alt="" <img src="@/assets/home/search.png" alt=""
@@ -27,7 +27,8 @@
</div> </div>
<img src="@/assets/home/avatar.png" alt="" style="width: 21px;height: 21px;" /> <img src="@/assets/home/avatar.png" alt="" style="width: 21px;height: 21px;" />
<div class="user-info"> <div class="user-info">
158 3632 4642</div> {{ userInfo.mobile }}
</div>
</div> </div>
</div> </div>
</el-header> </el-header>
@@ -56,10 +57,11 @@
</div> </div>
<div <div
style="display: flex;flex-wrap: wrap;margin-top: 15px;gap: 15px;justify-content: space-between;box-sizing: border-box;"> style="display: flex;flex-wrap: wrap;margin-top: 15px;gap: 15px;justify-content: space-between;box-sizing: border-box;">
<div class="device-item" v-for="(item,index) in 10" :key="index"> <div class="device-item" v-for="(item,index) in deviceList" :key="index">
<div style="display: flex;justify-content: space-between;"> <div style="display: flex;justify-content: space-between;">
<div style="font-weight: 700;font-size: 18px;text-align: left;color: #3d4566;"> <div style="font-weight: 700;font-size: 18px;text-align: left;color: #3d4566;">
CC:ba:97:11:a6:ac <!-- CC:ba:97:11:a6:ac-->
{{item.list[0]?.mac_address}}
</div> </div>
<div> <div>
<img src="@/assets/home/delete.png" alt="" <img src="@/assets/home/delete.png" alt=""
@@ -68,7 +70,7 @@
</div> </div>
</div> </div>
<div class="device-name"> <div class="device-name">
设备型号:esp32-s3-touch-amoled-1.8 设备型号:{{item.list[0]?.device_type}}
</div> </div>
<div style="display: flex;gap: 8px;align-items: center;"> <div style="display: flex;gap: 8px;align-items: center;">
<div class="settings-btn" @click="clickSettingDevice"> <div class="settings-btn" @click="clickSettingDevice">
@@ -77,12 +79,12 @@
声纹识别</div> 声纹识别</div>
<div class="settings-btn"> <div class="settings-btn">
历史对话</div> 历史对话</div>
<el-switch v-model="switchValue" inactive-text="OTA升级:" :width="32" <el-switch :value="item.list[0]?.ota_upgrade && true || false" inactive-text="OTA升级:" :width="32"
style="margin-left: auto;" /> style="margin-left: auto;" />
</div> </div>
<div class="version-info"> <div class="version-info">
<div>最近对话:6天前</div> <div>最近对话:{{item.list[0]?.recent_chat_time}}</div>
<div>APP版本:1.1.0</div> <div>APP版本:{{item.list[0]?.app_version}}</div>
</div> </div>
</div> </div>
</div> </div>
@@ -185,7 +187,8 @@
</div> </div>
<div <div
style="font-size: 12px;font-weight: 400;margin-top: auto;padding-top: 30px;color: #979db1;"> style="font-size: 12px;font-weight: 400;margin-top: auto;padding-top: 30px;color: #979db1;">
©2025 xiaozhi-esp32-server</div> ©2025 xiaozhi-esp32-server
</div>
</el-main> </el-main>
</el-container> </el-container>
<el-dialog :visible.sync="addDeviceDialogVisible" width="400px" center> <el-dialog :visible.sync="addDeviceDialogVisible" width="400px" center>
@@ -220,6 +223,8 @@
<script> <script>
// @ is an alias to /src // @ is an alias to /src
import Api from '@/apis/api';
export default { export default {
name: 'home', name: 'home',
data() { data() {
@@ -242,8 +247,12 @@ export default {
}, { }, {
value: '选项2', value: '选项2',
label: '双皮奶' label: '双皮奶'
}] }],
} userInfo: {
mobile: '' // 初始化用户信息
},
deviceList:[]
};
}, },
methods: { methods: {
showAddDialog() { showAddDialog() {
@@ -251,10 +260,28 @@ export default {
}, },
clickSettingDevice() { clickSettingDevice() {
this.settingDevice = true this.settingDevice = true
},
// 获取用户信息
fetchUserInfo() {
Api.user.getUserInfo(({data}) => {
this.userInfo = data.data
});
},
// 获取已绑设备
getList(){
Api.user.getHomeList(({data})=>{
console.log(data.data)
this.deviceList = data.data
})
} }
},
mounted() {
this.fetchUserInfo(); // 组件加载时获取用户信息
this.getList()
} }
} }
</script> </script>
<style scoped lang="scss"> <style scoped lang="scss">
.welcome { .welcome {
min-width: 900px; min-width: 900px;
@@ -464,7 +491,7 @@ export default {
} }
} }
.device-item { .device-item {
width: 350px; width: 345px;
border-radius: 15px; border-radius: 15px;
background: #fafcfe; background: #fafcfe;
padding: 22px; padding: 22px;
@@ -507,6 +534,7 @@ audio::-webkit-media-controls-panel {
line-height: 34px; line-height: 34px;
box-sizing: border-box; box-sizing: border-box;
cursor: pointer; cursor: pointer;
font-size: 12px;
} }
.save-btn { .save-btn {
border-radius: 23px; border-radius: 23px;
+43 -1
View File
@@ -158,6 +158,15 @@ LLM:
base_url: http://homeassistant.local:8123 base_url: http://homeassistant.local:8123
agent_id: conversation.chatgpt agent_id: conversation.chatgpt
api_key: 你的home assistant api访问令牌 api_key: 你的home assistant api访问令牌
FastgptLLM:
# 定义LLM API类型
type: fastgpt
# 如果使用fastgpt,配置文件里prompt(提示词)是无效的,需要在fastgpt控制台设置提示词
base_url: https://host/api/v1
api_key: fastgpt-xxx
variables:
k: "v"
k2: "v2"
TTS: TTS:
# 当前支持的type为edge、doubao,可自行适配 # 当前支持的type为edge、doubao,可自行适配
EdgeTTS: EdgeTTS:
@@ -296,8 +305,11 @@ TTS:
type: aliyun type: aliyun
output_file: tmp/ output_file: tmp/
appkey: 你的阿里云智能语音交互服务项目Appkey appkey: 你的阿里云智能语音交互服务项目Appkey
token: 你的阿里云智能语音交互服务AccessToken token: 你的阿里云智能语音交互服务AccessToken,临时的24小时,要长期用下方的access_key_idaccess_key_secret
voice: xiaoyun voice: xiaoyun
access_key_id: 你的阿里云账号access_key_id
access_key_secret: 你的阿里云账号access_key_secret
# 以下可不用设置,使用默认设置 # 以下可不用设置,使用默认设置
# format: wav # format: wav
# sample_rate: 16000 # sample_rate: 16000
@@ -316,6 +328,36 @@ TTS:
voice: "zh_female_wanwanxiaohe_moon_bigtts" voice: "zh_female_wanwanxiaohe_moon_bigtts"
output_file: tmp/ output_file: tmp/
access_token: "你的302API密钥" access_token: "你的302API密钥"
ACGNTTS:
#在线网址:https://acgn.ttson.cn/
#token购买:www.ttson.cn
#开发相关疑问请提交至3497689533@qq.com
#角色id获取地址:ctrl+f快速检索角色——网站管理者不允许发布,可询问网站管理者:1069379506
#各参数意义见开发文档:https://www.yuque.com/alexuh/skmti9/wm6taqislegb02gd?singleDoc#
type: ttson
token: your_token
voice_id: 1695
speed_factor: 1
pitch_factor: 0
volume_change_dB: 0
to_lang: ZH
url: https://u95167-bd74-2aef8085.westx.seetacloud.com:8443/flashsummary/tts?token=
format: mp3
output_file: tmp/
emotion: 1
OpenAITTS:
# openai官方文本转语音服务,可支持全球大多数语种
type: openai
api_key: 你的openai api key
# 国内需要使用代理
api_url: https://api.openai.com/v1/audio/speech
# 可选tts-1或tts-1-hdtts-1速度更快tts-1-hd质量更好
model: tts-1
# 演讲者,可选alloy, echo, fable, onyx, nova, shimmer
voice: onyx
# 语速范围0.25-4.0
speed: 1
output_file: tmp/
# 模块测试配置 # 模块测试配置
module_test: module_test:
test_sentences: # 自定义测试语句 test_sentences: # 自定义测试语句
@@ -0,0 +1,65 @@
import json
from config.logger import setup_logging
import requests
from core.providers.llm.base import LLMProviderBase
TAG = __name__
logger = setup_logging()
class LLMProvider(LLMProviderBase):
def __init__(self, config):
self.api_key = config["api_key"]
self.base_url = config.get("base_url")
self.detail = config.get("detail", False)
self.variables = config.get("variables", {})
def response(self, session_id, dialogue):
try:
# 取最后一条用户消息
last_msg = next(m for m in reversed(dialogue) if m["role"] == "user")
# 发起流式请求
with requests.post(
f"{self.base_url}/chat/completions",
headers={"Authorization": f"Bearer {self.api_key}"},
json={
"stream": True,
"chatId": session_id,
"detail": self.detail,
"variables": self.variables,
"messages": [
{
"role": "user",
"content": last_msg["content"]
}
]
},
stream=True
) as r:
for line in r.iter_lines():
if line:
try:
if line.startswith(b'data: '):
if line[6:].decode('utf-8') == '[DONE]':
break
data = json.loads(line[6:])
if 'choices' in data and len(data['choices']) > 0:
delta = data['choices'][0].get('delta', {})
if delta and 'content' in delta and delta['content'] is not None:
content = delta['content']
if '<think>' in content:
continue
if '</think>' in content:
continue
yield content
except json.JSONDecodeError as e:
continue
except Exception as e:
continue
except Exception as e:
logger.bind(tag=TAG).error(f"Error in response generation: {e}")
yield "【服务响应异常】"
@@ -1,9 +1,7 @@
from config.logger import setup_logging
import google.generativeai as genai import google.generativeai as genai
from core.utils.util import check_model_key
from core.providers.llm.base import LLMProviderBase from core.providers.llm.base import LLMProviderBase
TAG = __name__
logger = setup_logging()
class LLMProvider(LLMProviderBase): class LLMProvider(LLMProviderBase):
def __init__(self, config): def __init__(self, config):
@@ -11,8 +9,9 @@ class LLMProvider(LLMProviderBase):
self.model_name = config.get("model_name", "gemini-1.5-pro") self.model_name = config.get("model_name", "gemini-1.5-pro")
self.api_key = config.get("api_key") self.api_key = config.get("api_key")
if not self.api_key or "" in self.api_key: have_key = check_model_key("LLM", self.api_key)
logger.bind(tag=TAG).error("你还没配置Gemini LLM的密钥,请在配置文件中配置密钥,否则无法正常工作")
if not have_key:
return return
try: try:
@@ -1,10 +1,7 @@
from config.logger import setup_logging
import openai import openai
from core.utils.util import check_model_key
from core.providers.llm.base import LLMProviderBase from core.providers.llm.base import LLMProviderBase
TAG = __name__
logger = setup_logging()
class LLMProvider(LLMProviderBase): class LLMProvider(LLMProviderBase):
def __init__(self, config): def __init__(self, config):
@@ -14,8 +11,7 @@ class LLMProvider(LLMProviderBase):
self.base_url = config.get("base_url") self.base_url = config.get("base_url")
else: else:
self.base_url = config.get("url") self.base_url = config.get("url")
if "" in self.api_key: check_model_key("LLM", self.api_key)
logger.bind(tag=TAG).error("你还没配置LLM的密钥,请在配置文件中配置密钥,否则无法正常工作")
self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url) self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url)
def response(self, session_id, dialogue): def response(self, session_id, dialogue):
@@ -1,5 +1,8 @@
import traceback
from ..base import MemoryProviderBase, logger from ..base import MemoryProviderBase, logger
from mem0 import MemoryClient from mem0 import MemoryClient
from core.utils.util import check_model_key
TAG = __name__ TAG = __name__
@@ -8,13 +11,19 @@ class MemoryProvider(MemoryProviderBase):
super().__init__(config) super().__init__(config)
self.api_key = config.get("api_key", "") self.api_key = config.get("api_key", "")
self.api_version = config.get("api_version", "v1.1") self.api_version = config.get("api_version", "v1.1")
if len(self.api_key) == 0 or "" in self.api_key: have_key = check_model_key("Mem0ai", self.api_key)
logger.bind(tag=TAG).error("你还没配置Mem0ai的密钥,请在配置文件中配置密钥,否则无法提供记忆服务") if not have_key :
self.use_mem0 = False self.use_mem0 = False
return return
else: else:
self.use_mem0 = True self.use_mem0 = True
self.client = MemoryClient(api_key=self.api_key) try:
self.client = MemoryClient(api_key=self.api_key)
logger.bind(tag=TAG).info("成功连接到 Mem0ai 服务")
except Exception as e:
logger.bind(tag=TAG).error(f"连接到 Mem0ai 服务时发生错误: {str(e)}")
logger.bind(tag=TAG).error(f"详细错误: {traceback.format_exc()}")
self.use_mem0 = False
async def save_memory(self, msgs): async def save_memory(self, msgs):
if not self.use_mem0: if not self.use_mem0:
@@ -1,19 +1,94 @@
import os import os
import uuid import uuid
import json import json
import hmac
import hashlib
import base64
import requests import requests
from datetime import datetime from datetime import datetime
from core.providers.tts.base import TTSProviderBase from core.providers.tts.base import TTSProviderBase
import http.client import http.client
import urllib.parse import urllib.parse
import time
import uuid
from urllib import parse
class AccessToken:
@staticmethod
def _encode_text(text):
encoded_text = parse.quote_plus(text)
return encoded_text.replace('+', '%20').replace('*', '%2A').replace('%7E', '~')
@staticmethod
def _encode_dict(dic):
keys = dic.keys()
dic_sorted = [(key, dic[key]) for key in sorted(keys)]
encoded_text = parse.urlencode(dic_sorted)
return encoded_text.replace('+', '%20').replace('*', '%2A').replace('%7E', '~')
@staticmethod
def create_token(access_key_id, access_key_secret):
parameters = {'AccessKeyId': access_key_id,
'Action': 'CreateToken',
'Format': 'JSON',
'RegionId': 'cn-shanghai',
'SignatureMethod': 'HMAC-SHA1',
'SignatureNonce': str(uuid.uuid1()),
'SignatureVersion': '1.0',
'Timestamp': time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
'Version': '2019-02-28'}
# 构造规范化的请求字符串
query_string = AccessToken._encode_dict(parameters)
print('规范化的请求字符串: %s' % query_string)
# 构造待签名字符串
string_to_sign = 'GET' + '&' + AccessToken._encode_text('/') + '&' + AccessToken._encode_text(query_string)
print('待签名的字符串: %s' % string_to_sign)
# 计算签名
secreted_string = hmac.new(bytes(access_key_secret + '&', encoding='utf-8'),
bytes(string_to_sign, encoding='utf-8'),
hashlib.sha1).digest()
signature = base64.b64encode(secreted_string)
print('签名: %s' % signature)
# 进行URL编码
signature = AccessToken._encode_text(signature)
print('URL编码后的签名: %s' % signature)
# 调用服务
full_url = 'http://nls-meta.cn-shanghai.aliyuncs.com/?Signature=%s&%s' % (signature, query_string)
print('url: %s' % full_url)
# 提交HTTP GET请求
response = requests.get(full_url)
if response.ok:
root_obj = response.json()
key = 'Token'
if key in root_obj:
token = root_obj[key]['Id']
expire_time = root_obj[key]['ExpireTime']
return token, expire_time
print(response.text)
return None, None
class TTSProvider(TTSProviderBase): class TTSProvider(TTSProviderBase):
def __init__(self, config, delete_audio_file): def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file) super().__init__(config, delete_audio_file)
# 新增空值判断逻辑
access_key_id = config.get("access_key_id")
access_key_secret = config.get("access_key_secret")
if access_key_id and access_key_secret:
# 使用密钥对生成临时token
token, expire_time = AccessToken.create_token(access_key_id, access_key_secret)
else:
# 直接使用预生成的长期token
token = config.get("token")
expire_time = None
print('token: %s, expire time(s): %s' % (token, expire_time))
self.appkey = config.get("appkey") self.appkey = config.get("appkey")
self.token = config.get("token") self.token = token
self.format = config.get("format", "wav") self.format = config.get("format", "wav")
self.sample_rate = config.get("sample_rate", 16000) self.sample_rate = config.get("sample_rate", 16000)
self.voice = config.get("voice", "xiaoyun") self.voice = config.get("voice", "xiaoyun")
@@ -4,6 +4,7 @@ import json
import base64 import base64
import requests import requests
from datetime import datetime from datetime import datetime
from core.utils.util import check_model_key
from core.providers.tts.base import TTSProviderBase from core.providers.tts.base import TTSProviderBase
@@ -17,6 +18,7 @@ class TTSProvider(TTSProviderBase):
self.api_url = config.get("api_url") self.api_url = config.get("api_url")
self.authorization = config.get("authorization") self.authorization = config.get("authorization")
self.header = {"Authorization": f"{self.authorization}{self.access_token}"} self.header = {"Authorization": f"{self.authorization}{self.access_token}"}
check_model_key("TTS", self.access_token)
def generate_filename(self, extension=".wav"): def generate_filename(self, extension=".wav"):
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}") return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
@@ -1,4 +1,3 @@
import base64 import base64
import os import os
import uuid import uuid
@@ -9,6 +8,7 @@ from pydantic import BaseModel, Field, conint, model_validator
from typing_extensions import Annotated from typing_extensions import Annotated
from datetime import datetime from datetime import datetime
from typing import Literal from typing import Literal
from core.utils.util import check_model_key
from core.providers.tts.base import TTSProviderBase from core.providers.tts.base import TTSProviderBase
from config.logger import setup_logging from config.logger import setup_logging
@@ -24,7 +24,7 @@ class ServeReferenceAudio(BaseModel):
def decode_audio(cls, values): def decode_audio(cls, values):
audio = values.get("audio") audio = values.get("audio")
if ( if (
isinstance(audio, str) and len(audio) > 255 isinstance(audio, str) and len(audio) > 255
): # Check if audio is a string (Base64) ): # Check if audio is a string (Base64)
try: try:
values["audio"] = base64.b64decode(audio) values["audio"] = base64.b64decode(audio)
@@ -36,6 +36,7 @@ class ServeReferenceAudio(BaseModel):
def __repr__(self) -> str: def __repr__(self) -> str:
return f"ServeReferenceAudio(text={self.text!r}, audio_size={len(self.audio)})" return f"ServeReferenceAudio(text={self.text!r}, audio_size={len(self.audio)})"
class ServeTTSRequest(BaseModel): class ServeTTSRequest(BaseModel):
text: str text: str
chunk_length: Annotated[int, conint(ge=100, le=300, strict=True)] = 200 chunk_length: Annotated[int, conint(ge=100, le=300, strict=True)] = 200
@@ -70,6 +71,7 @@ def audio_to_bytes(file_path):
wav = wav_file.read() wav = wav_file.read()
return wav return wav
def read_ref_text(ref_text): def read_ref_text(ref_text):
path = Path(ref_text) path = Path(ref_text)
if path.exists() and path.is_file(): if path.exists() and path.is_file():
@@ -77,31 +79,32 @@ def read_ref_text(ref_text):
return file.read() return file.read()
return ref_text return ref_text
class TTSProvider(TTSProviderBase): class TTSProvider(TTSProviderBase):
def __init__(self, config, delete_audio_file): def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file) super().__init__(config, delete_audio_file)
self.reference_id = config.get("reference_id") self.reference_id = config.get("reference_id")
self.reference_audio = config.get("reference_audio",[]) self.reference_audio = config.get("reference_audio", [])
self.reference_text = config.get("reference_text",[]) self.reference_text = config.get("reference_text", [])
self.format = config.get("format","wav") self.format = config.get("format", "wav")
self.channels = config.get("channels",1) self.channels = config.get("channels", 1)
self.rate = config.get("rate",44100) self.rate = config.get("rate", 44100)
self.api_key = config.get("api_key","YOUR_API_KEY") self.api_key = config.get("api_key", "YOUR_API_KEY")
if "" in self.api_key: have_key = check_model_key("FishSpeech TTS", self.api_key)
logger.bind(tag=TAG).error("你还没配置FishSpeech TTS的密钥,请在配置文件中配置密钥,否则无法正常工作") if not have_key:
return return
self.normalize = config.get("normalize",True) self.normalize = config.get("normalize", True)
self.max_new_tokens = config.get("max_new_tokens",1024) self.max_new_tokens = config.get("max_new_tokens", 1024)
self.chunk_length = config.get("chunk_length",200) self.chunk_length = config.get("chunk_length", 200)
self.top_p = config.get("top_p",0.7) self.top_p = config.get("top_p", 0.7)
self.repetition_penalty = config.get("repetition_penalty",1.2) self.repetition_penalty = config.get("repetition_penalty", 1.2)
self.temperature = config.get("temperature",0.7) self.temperature = config.get("temperature", 0.7)
self.streaming = config.get("streaming",False) self.streaming = config.get("streaming", False)
self.use_memory_cache = config.get("use_memory_cache","on") self.use_memory_cache = config.get("use_memory_cache", "on")
self.seed = config.get("seed") self.seed = config.get("seed")
self.api_url = config.get("api_url","http://127.0.0.1:8080/v1/tts") self.api_url = config.get("api_url", "http://127.0.0.1:8080/v1/tts")
def generate_filename(self, extension=".wav"): def generate_filename(self, extension=".wav"):
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}") return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
@@ -154,5 +157,3 @@ class TTSProvider(TTSProviderBase):
else: else:
print(f"Request failed with status code {response.status_code}") print(f"Request failed with status code {response.status_code}")
print(response.json()) print(response.json())
@@ -0,0 +1,40 @@
import os
import uuid
import requests
from datetime import datetime
from core.utils.util import check_model_key
from core.providers.tts.base import TTSProviderBase
class TTSProvider(TTSProviderBase):
def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file)
self.api_key = config.get("api_key")
self.api_url = config.get("api_url", "https://api.openai.com/v1/audio/speech")
self.model = config.get("model", "tts-1")
self.voice = config.get("voice", "alloy")
self.response_format = "wav"
self.speed = config.get("speed", 1.0)
self.output_file = config.get("output_file", "tmp/")
check_model_key("TTS", self.api_key)
def generate_filename(self, extension=".wav"):
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
async def text_to_speak(self, text, output_file):
headers = {
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json"
}
data = {
"model": self.model,
"input": text,
"voice": self.voice,
"response_format": "wav",
"speed": self.speed
}
response = requests.post(self.api_url, json=data, headers=headers)
if response.status_code == 200:
with open(output_file, "wb") as audio_file:
audio_file.write(response.content)
else:
raise Exception(f"OpenAI TTS请求失败: {response.status_code} - {response.text}")
@@ -0,0 +1,64 @@
import os
import uuid
import json
import requests
import shutil
from datetime import datetime
from core.providers.tts.base import TTSProviderBase
class TTSProvider(TTSProviderBase):
def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file)
self.url = config.get("url", "https://u95167-bd74-2aef8085.westx.seetacloud.com:8443/flashsummary/tts?token=")
self.voice_id = config.get("voice_id", 1695)
self.token = config.get("token")
self.to_lang = config.get("to_lang")
self.volume_change_dB = config.get("volume_change_dB", 0)
self.speed_factor = config.get("speed_factor", 1)
self.stream = config.get("stream", False)
self.output_file = config.get("output_file")
self.pitch_factor = config.get("pitch_factor", 0)
self.format = config.get("format", "mp3")
self.emotion = config.get("emotion", 1)
self.header = {
"Content-Type": "application/json"
}
def generate_filename(self, extension=".mp3"):
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
async def text_to_speak(self, text, output_file):
url = f'{self.url}{self.token}'
result = "firefly"
payload = json.dumps({
"to_lang": self.to_lang,
"text": text,
"emotion": self.emotion,
"format": self.format,
"volume_change_dB": self.volume_change_dB,
"voice_id": self.voice_id,
"pitch_factor": self.pitch_factor,
"speed_factor": self.speed_factor,
"token": self.token
})
resp = requests.request("POST", url, data=payload)
if resp.status_code != 200:
return None
resp_json = resp.json()
try:
result = resp_json['url'] + ':' + str(
resp_json[
'port']) + '/flashsummary/retrieveFileData?stream=True&token=' + self.token + '&voice_audio_path=' + \
resp_json['voice_path']
except Exception as e:
print("error:", e)
audio_content = requests.get(result)
with open(output_file, "wb") as f:
f.write(audio_content.content)
return True
voice_path = resp_json.get("voice_path")
des_path = output_file
shutil.move(voice_path, des_path)
+5 -27
View File
@@ -1,9 +1,9 @@
import os import os
import re
import json import json
import yaml import yaml
import socket import socket
import subprocess import subprocess
import logging
def get_project_dir(): def get_project_dir():
@@ -87,35 +87,13 @@ def remove_punctuation_and_length(text):
return 0, "" return 0, ""
return len(result), result return len(result), result
def check_model_key(modelType, modelKey):
def check_password(password): if "" in modelKey:
""" logging.error("你还没配置" + modelType + "的密钥,请在配置文件中配置密钥,否则无法正常工作")
检查密码是否满足以下条件:
1. 密码长度大于八位。
2. 密码包含英文和数字。
3. 密码不能包含“xiaozhi”字符。
:param password: 要检查的密码
:return: 如果密码满足条件,则返回True;否则返回False。
"""
# 检查密码长度
if len(password) < 8:
return False return False
# 检查是否包含英文字符和数字
if not re.search(r'[A-Za-z]', password) or not re.search(r'[0-9]', password):
return False
# 检查是否包含“xiaozhi”字符
if "xiaozhi" in password:
return False
if "1234" in password:
return False
# 如果满足所有条件,则返回True
return True return True
def check_ffmpeg_installed(): def check_ffmpeg_installed():
ffmpeg_installed = False ffmpeg_installed = False
try: try: