mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-27 17:43:55 +08:00
增加通过认证码绑定设备和管理账户的功能,需要开启私有设备配置
增加登录用户鉴权
This commit is contained in:
@@ -1 +1 @@
|
|||||||
VITE_API_BASE_URL=''
|
VITE_API_BASE_URL='http://127.0.0.1:8002'
|
||||||
|
|||||||
@@ -19,7 +19,7 @@
|
|||||||
<script setup>
|
<script setup>
|
||||||
import { ref } from 'vue';
|
import { ref } from 'vue';
|
||||||
import { useRouter } from 'vue-router';
|
import { useRouter } from 'vue-router';
|
||||||
import { API_BASE_URL } from '../config/api';
|
import apiClient from '../utils/api';
|
||||||
|
|
||||||
const router = useRouter();
|
const router = useRouter();
|
||||||
const username = ref('');
|
const username = ref('');
|
||||||
@@ -34,24 +34,16 @@ const handleLogin = async () => {
|
|||||||
|
|
||||||
isLoading.value = true;
|
isLoading.value = true;
|
||||||
try {
|
try {
|
||||||
const response = await fetch(`${API_BASE_URL}/api/login`, {
|
const response = await apiClient.post('/api/login', {
|
||||||
method: 'POST',
|
username: username.value,
|
||||||
headers: {
|
password: password.value
|
||||||
'Content-Type': 'application/json',
|
|
||||||
},
|
|
||||||
body: JSON.stringify({
|
|
||||||
username: username.value,
|
|
||||||
password: password.value
|
|
||||||
})
|
|
||||||
});
|
});
|
||||||
|
|
||||||
const data = await response.json();
|
const data = response.data;
|
||||||
|
|
||||||
if (data.success) {
|
if (data.success) {
|
||||||
// 存储token和登录状态
|
localStorage.setItem('session_id', data.session_id);
|
||||||
localStorage.setItem('token', data.token);
|
|
||||||
localStorage.setItem('isLoggedIn', 'true');
|
localStorage.setItem('isLoggedIn', 'true');
|
||||||
// 使用路由导航到panel页面
|
|
||||||
router.push('/panel');
|
router.push('/panel');
|
||||||
} else {
|
} else {
|
||||||
alert(data.message || '登录失败');
|
alert(data.message || '登录失败');
|
||||||
|
|||||||
@@ -2,11 +2,46 @@
|
|||||||
<div class="app">
|
<div class="app">
|
||||||
<NavBar current-tab="device" @tab-change="handleTabChange"/>
|
<NavBar current-tab="device" @tab-change="handleTabChange"/>
|
||||||
<main class="content">
|
<main class="content">
|
||||||
<div class="breadcrumb">
|
<div class="page-header">
|
||||||
<router-link to="/">首页</router-link> /
|
<div class="header-left">
|
||||||
<span>设备管理</span>
|
|
||||||
|
<div class="breadcrumb">
|
||||||
|
<router-link to="/">首页</router-link> /
|
||||||
|
<span>设备管理</span>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
</div>
|
</div>
|
||||||
|
<button class="add-btn" @click="showBindDialog = true">
|
||||||
|
<i class="icon-plus"></i>添加设备
|
||||||
|
</button>
|
||||||
|
<!-- 绑定设备弹窗 -->
|
||||||
|
<div v-if="showBindDialog" class="dialog-overlay">
|
||||||
|
<div class="dialog">
|
||||||
|
<h3>绑定新设备</h3>
|
||||||
|
<div class="form-group">
|
||||||
|
<label>请输入6位认证码:</label>
|
||||||
|
<input
|
||||||
|
type="text"
|
||||||
|
v-model="authCode"
|
||||||
|
maxlength="6"
|
||||||
|
pattern="\d*"
|
||||||
|
placeholder="请输入6位数字认证码"
|
||||||
|
@input="handleAuthCodeInput"
|
||||||
|
/>
|
||||||
|
</div>
|
||||||
|
<div class="dialog-buttons">
|
||||||
|
<button @click="showBindDialog = false">取消</button>
|
||||||
|
<button
|
||||||
|
class="primary"
|
||||||
|
@click="handleBindDevice"
|
||||||
|
:disabled="authCode.length !== 6 || isBinding"
|
||||||
|
>
|
||||||
|
{{ isBinding ? '绑定中...' : '确认绑定' }}
|
||||||
|
</button>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
|
||||||
<template v-if="devices.length > 0">
|
<template v-if="devices.length > 0">
|
||||||
<div class="device-list">
|
<div class="device-list">
|
||||||
<DeviceCard
|
<DeviceCard
|
||||||
@@ -25,14 +60,7 @@
|
|||||||
/>
|
/>
|
||||||
</div>
|
</div>
|
||||||
</template>
|
</template>
|
||||||
<template v-else>
|
|
||||||
<div class="empty-state">
|
|
||||||
<div class="empty-message">
|
|
||||||
<i class="icon-info"></i>
|
|
||||||
<p>目前没有设备,请确认是否启用私有配置,并且和设备进行一次对话</p>
|
|
||||||
</div>
|
|
||||||
</div>
|
|
||||||
</template>
|
|
||||||
</main>
|
</main>
|
||||||
</div>
|
</div>
|
||||||
</template>
|
</template>
|
||||||
@@ -42,12 +70,78 @@ import { ref, onMounted } from 'vue';
|
|||||||
import { useRouter } from 'vue-router';
|
import { useRouter } from 'vue-router';
|
||||||
import NavBar from './NavBar.vue';
|
import NavBar from './NavBar.vue';
|
||||||
import DeviceCard from './DeviceCard.vue';
|
import DeviceCard from './DeviceCard.vue';
|
||||||
import { API_BASE_URL } from '../config/api';
|
import apiClient from '../utils/api';
|
||||||
|
|
||||||
const router = useRouter();
|
const router = useRouter();
|
||||||
const baseUrl = API_BASE_URL;
|
|
||||||
const devices = ref([]);
|
const devices = ref([]);
|
||||||
|
|
||||||
|
// 绑定设备相关的状态
|
||||||
|
const showBindDialog = ref(false);
|
||||||
|
const authCode = ref('');
|
||||||
|
const isBinding = ref(false);
|
||||||
|
|
||||||
|
// 处理认证码输入,只允许数字
|
||||||
|
const handleAuthCodeInput = (event) => {
|
||||||
|
authCode.value = event.target.value.replace(/\D/g, '').slice(0, 6);
|
||||||
|
};
|
||||||
|
|
||||||
|
// 处理设备绑定
|
||||||
|
const handleBindDevice = async () => {
|
||||||
|
if (authCode.value.length !== 6) {
|
||||||
|
alert('请输入6位数字认证码');
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
isBinding.value = true;
|
||||||
|
try {
|
||||||
|
const response = await apiClient.post('/api/config/bind_device', {
|
||||||
|
auth_code: authCode.value
|
||||||
|
});
|
||||||
|
|
||||||
|
if (response.data.success) {
|
||||||
|
alert('设备绑定成功');
|
||||||
|
showBindDialog.value = false;
|
||||||
|
authCode.value = '';
|
||||||
|
// 刷新设备列表
|
||||||
|
loadDevices();
|
||||||
|
} else {
|
||||||
|
throw new Error(response.data.message);
|
||||||
|
}
|
||||||
|
} catch (error) {
|
||||||
|
alert(error.response?.data?.message || error.message || '绑定失败');
|
||||||
|
} finally {
|
||||||
|
isBinding.value = false;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
// 将现有的加载设备方法提取出来
|
||||||
|
const loadDevices = async () => {
|
||||||
|
try {
|
||||||
|
const response = await apiClient.get('/api/config/devices');
|
||||||
|
|
||||||
|
if (response.data.success) {
|
||||||
|
const deviceArray = Object.entries(response.data.data).map(([id, config]) => ({
|
||||||
|
id,
|
||||||
|
config,
|
||||||
|
type: '面包板(WiFi)',
|
||||||
|
version: '0.9.9',
|
||||||
|
lastActivity: '3 天前',
|
||||||
|
note: ''
|
||||||
|
}));
|
||||||
|
devices.value = deviceArray;
|
||||||
|
} else {
|
||||||
|
throw new Error(response.data.message || '加载设备失败');
|
||||||
|
}
|
||||||
|
} catch (error) {
|
||||||
|
console.error('Error loading devices:', error);
|
||||||
|
// Show error message to user
|
||||||
|
const errorMessage = error.message || '加载设备失败,请检查网络连接';
|
||||||
|
alert(errorMessage);
|
||||||
|
|
||||||
|
// If user is not logged in, redirect will be handled by api interceptor
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
const formatLastActivity = (timestamp) => {
|
const formatLastActivity = (timestamp) => {
|
||||||
if (!timestamp) return '从未对话';
|
if (!timestamp) return '从未对话';
|
||||||
|
|
||||||
@@ -82,20 +176,15 @@ const handleHistory = (device) => {
|
|||||||
|
|
||||||
const handleDelete = async (device) => {
|
const handleDelete = async (device) => {
|
||||||
try {
|
try {
|
||||||
const response = await fetch(`${baseUrl}/api/config/delete_device`, {
|
const response = await apiClient.post('/api/config/delete_device', {
|
||||||
method: 'POST',
|
device_id: device.id
|
||||||
headers: {
|
|
||||||
'Content-Type': 'application/json'
|
|
||||||
},
|
|
||||||
body: JSON.stringify({ device_id: device.id })
|
|
||||||
});
|
});
|
||||||
|
|
||||||
const data = await response.json();
|
if (response.data.success) {
|
||||||
if (data.success) {
|
|
||||||
devices.value = devices.value.filter(d => d.id !== device.id);
|
devices.value = devices.value.filter(d => d.id !== device.id);
|
||||||
alert('设备已删除');
|
alert('设备已删除');
|
||||||
} else {
|
} else {
|
||||||
throw new Error(data.message || '删除失败');
|
throw new Error(response.data.message || '删除失败');
|
||||||
}
|
}
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
console.error('Error deleting device:', error);
|
console.error('Error deleting device:', error);
|
||||||
@@ -110,23 +199,7 @@ const handleTabChange = (tab) => {
|
|||||||
};
|
};
|
||||||
|
|
||||||
// Load devices on mount
|
// Load devices on mount
|
||||||
onMounted(async () => {
|
onMounted(loadDevices);
|
||||||
try {
|
|
||||||
const response = await fetch(`${baseUrl}/api/config/devices`);
|
|
||||||
const data = await response.json();
|
|
||||||
if (data.success) {
|
|
||||||
devices.value = data.data.map(device => ({
|
|
||||||
...device,
|
|
||||||
type: '面包板(WiFi)',
|
|
||||||
version: '0.9.9',
|
|
||||||
lastActivity: '3 天前',
|
|
||||||
note: ''
|
|
||||||
}));
|
|
||||||
}
|
|
||||||
} catch (error) {
|
|
||||||
console.error('Error loading devices:', error);
|
|
||||||
}
|
|
||||||
});
|
|
||||||
</script>
|
</script>
|
||||||
|
|
||||||
<style scoped>
|
<style scoped>
|
||||||
@@ -305,4 +378,115 @@ onMounted(async () => {
|
|||||||
background-size: contain;
|
background-size: contain;
|
||||||
opacity: 0.6;
|
opacity: 0.6;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
.page-header {
|
||||||
|
display: flex;
|
||||||
|
justify-content: flex-start;
|
||||||
|
align-items: center;
|
||||||
|
margin-bottom: 20px;
|
||||||
|
}
|
||||||
|
|
||||||
|
.header-left {
|
||||||
|
display: flex;
|
||||||
|
align-items: center;
|
||||||
|
gap: 20px;
|
||||||
|
}
|
||||||
|
|
||||||
|
.add-btn {
|
||||||
|
display: flex;
|
||||||
|
align-items: center;
|
||||||
|
gap: 8px;
|
||||||
|
padding: 8px 16px;
|
||||||
|
background-color: #28a745;
|
||||||
|
color: white;
|
||||||
|
border: none;
|
||||||
|
border-radius: 4px;
|
||||||
|
cursor: pointer;
|
||||||
|
font-size: 14px;
|
||||||
|
transition: all 0.3s ease;
|
||||||
|
height: 36px;
|
||||||
|
}
|
||||||
|
|
||||||
|
.add-btn:hover {
|
||||||
|
background-color: #218838;
|
||||||
|
}
|
||||||
|
|
||||||
|
.icon-plus {
|
||||||
|
font-size: 16px;
|
||||||
|
}
|
||||||
|
|
||||||
|
.dialog-overlay {
|
||||||
|
position: fixed;
|
||||||
|
top: 0;
|
||||||
|
left: 0;
|
||||||
|
right: 0;
|
||||||
|
bottom: 0;
|
||||||
|
background-color: rgba(0, 0, 0, 0.5);
|
||||||
|
display: flex;
|
||||||
|
align-items: center;
|
||||||
|
justify-content: center;
|
||||||
|
z-index: 1000;
|
||||||
|
}
|
||||||
|
|
||||||
|
.dialog {
|
||||||
|
background: white;
|
||||||
|
padding: 24px;
|
||||||
|
border-radius: 8px;
|
||||||
|
width: 90%;
|
||||||
|
max-width: 400px;
|
||||||
|
}
|
||||||
|
|
||||||
|
.dialog h3 {
|
||||||
|
margin: 0 0 20px;
|
||||||
|
font-size: 18px;
|
||||||
|
color: #2c3e50;
|
||||||
|
}
|
||||||
|
|
||||||
|
.dialog .form-group {
|
||||||
|
margin-bottom: 20px;
|
||||||
|
}
|
||||||
|
|
||||||
|
.dialog label {
|
||||||
|
display: block;
|
||||||
|
margin-bottom: 8px;
|
||||||
|
color: #4a5568;
|
||||||
|
}
|
||||||
|
|
||||||
|
.dialog input {
|
||||||
|
width: 100%;
|
||||||
|
padding: 8px 12px;
|
||||||
|
border: 1px solid #e2e8f0;
|
||||||
|
border-radius: 4px;
|
||||||
|
font-size: 16px;
|
||||||
|
}
|
||||||
|
|
||||||
|
.dialog-buttons {
|
||||||
|
display: flex;
|
||||||
|
justify-content: flex-end;
|
||||||
|
gap: 12px;
|
||||||
|
}
|
||||||
|
|
||||||
|
.dialog-buttons button {
|
||||||
|
padding: 8px 16px;
|
||||||
|
border: 1px solid #e2e8f0;
|
||||||
|
border-radius: 4px;
|
||||||
|
cursor: pointer;
|
||||||
|
font-size: 14px;
|
||||||
|
}
|
||||||
|
|
||||||
|
.dialog-buttons button.primary {
|
||||||
|
background-color: #28a745;
|
||||||
|
color: white;
|
||||||
|
border-color: #28a745;
|
||||||
|
}
|
||||||
|
|
||||||
|
.dialog-buttons button.primary:disabled {
|
||||||
|
background-color: #90be9c;
|
||||||
|
border-color: #90be9c;
|
||||||
|
cursor: not-allowed;
|
||||||
|
}
|
||||||
|
|
||||||
|
.breadcrumb {
|
||||||
|
margin-bottom: 0;
|
||||||
|
}
|
||||||
</style>
|
</style>
|
||||||
@@ -0,0 +1,61 @@
|
|||||||
|
import axios from 'axios';
|
||||||
|
import { API_BASE_URL } from '../config/api';
|
||||||
|
|
||||||
|
// Add server status check utility
|
||||||
|
export const checkServerStatus = async () => {
|
||||||
|
try {
|
||||||
|
await axios.get(`${API_BASE_URL}/health`, { timeout: 5000 });
|
||||||
|
return true;
|
||||||
|
} catch (error) {
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
const apiClient = axios.create({
|
||||||
|
baseURL: API_BASE_URL,
|
||||||
|
headers: {
|
||||||
|
'Content-Type': 'application/json',
|
||||||
|
},
|
||||||
|
timeout: 10000 // Add timeout
|
||||||
|
});
|
||||||
|
|
||||||
|
// 添加请求拦截器,自动添加 session_id
|
||||||
|
apiClient.interceptors.request.use(config => {
|
||||||
|
const sessionId = localStorage.getItem('session_id');
|
||||||
|
if (sessionId) {
|
||||||
|
config.headers.Authorization = sessionId;
|
||||||
|
}
|
||||||
|
return config;
|
||||||
|
});
|
||||||
|
|
||||||
|
// 响应拦截器
|
||||||
|
apiClient.interceptors.response.use(
|
||||||
|
response => response,
|
||||||
|
error => {
|
||||||
|
// Network error or server not reachable
|
||||||
|
if (!error.response || error.code === 'ERR_NETWORK' || error.code === 'ECONNABORTED') {
|
||||||
|
localStorage.removeItem('session_id');
|
||||||
|
localStorage.removeItem('isLoggedIn');
|
||||||
|
|
||||||
|
const errorMessage = error.code === 'ECONNABORTED'
|
||||||
|
? '服务器响应超时'
|
||||||
|
: '无法连接到服务器,请检查服务器是否正常运行';
|
||||||
|
|
||||||
|
if (window.location.pathname !== '/login') {
|
||||||
|
window.location.href = '/login';
|
||||||
|
}
|
||||||
|
return Promise.reject(new Error(errorMessage));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unauthorized error
|
||||||
|
if (error.response && error.response.status === 401) {
|
||||||
|
localStorage.removeItem('session_id');
|
||||||
|
localStorage.removeItem('isLoggedIn');
|
||||||
|
window.location.href = '/login';
|
||||||
|
}
|
||||||
|
|
||||||
|
return Promise.reject(error);
|
||||||
|
}
|
||||||
|
);
|
||||||
|
|
||||||
|
export default apiClient;
|
||||||
+227
-100
@@ -2,60 +2,76 @@ import os
|
|||||||
import time
|
import time
|
||||||
import yaml
|
import yaml
|
||||||
import logging
|
import logging
|
||||||
from typing import Dict, Any
|
from typing import Dict, Any, Optional
|
||||||
from copy import deepcopy
|
from copy import deepcopy
|
||||||
from core.utils.util import get_project_dir
|
from core.utils.util import get_project_dir
|
||||||
from core.utils import asr, vad, llm, tts
|
from core.utils import asr, vad, llm, tts
|
||||||
|
from manager.api.user_manager import UserManager
|
||||||
|
from core.utils.lock_manager import FileLockManager
|
||||||
|
|
||||||
class PrivateConfig:
|
class PrivateConfig:
|
||||||
def __init__(self, device_id: str, default_config: Dict[str, Any]):
|
def __init__(self, device_id: str, default_config: Dict[str, Any], auth_code_gen=None):
|
||||||
self.device_id = device_id
|
self.device_id = device_id
|
||||||
self.default_config = default_config
|
self.default_config = default_config
|
||||||
self.config_path = get_project_dir() + 'data/.private_config.yaml'
|
self.config_path = get_project_dir() + 'data/.private_config.yaml'
|
||||||
self.logger = logging.getLogger(__name__)
|
self.logger = logging.getLogger(__name__)
|
||||||
self.private_config = {}
|
self.private_config = {}
|
||||||
|
self.auth_code_gen = auth_code_gen
|
||||||
|
self.user_manager = UserManager()
|
||||||
|
self.lock_manager = FileLockManager()
|
||||||
|
|
||||||
async def load_or_create(self):
|
async def load_or_create(self):
|
||||||
try:
|
try:
|
||||||
if os.path.exists(self.config_path):
|
await self.lock_manager.acquire_lock(self.config_path)
|
||||||
with open(self.config_path, 'r', encoding='utf-8') as f:
|
try:
|
||||||
all_configs = yaml.safe_load(f) or {}
|
if os.path.exists(self.config_path):
|
||||||
else:
|
with open(self.config_path, 'r', encoding='utf-8') as f:
|
||||||
all_configs = {}
|
all_configs = yaml.safe_load(f) or {}
|
||||||
|
else:
|
||||||
|
all_configs = {}
|
||||||
|
|
||||||
if self.device_id not in all_configs:
|
if self.device_id not in all_configs:
|
||||||
# Get selected module names
|
# Get selected module names
|
||||||
selected_modules = self.default_config['selected_module']
|
selected_modules = self.default_config['selected_module']
|
||||||
selected_tts = selected_modules['TTS']
|
selected_tts = selected_modules['TTS']
|
||||||
selected_llm = selected_modules['LLM']
|
selected_llm = selected_modules['LLM']
|
||||||
selected_asr = selected_modules['ASR']
|
selected_asr = selected_modules['ASR']
|
||||||
selected_vad = selected_modules['VAD']
|
selected_vad = selected_modules['VAD']
|
||||||
|
|
||||||
# Initialize device config with only necessary configurations
|
# 生成认证码
|
||||||
device_config = {
|
auth_code = None
|
||||||
'selected_module': deepcopy(selected_modules),
|
if self.auth_code_gen:
|
||||||
'prompt': self.default_config['prompt'],
|
auth_code = self.auth_code_gen.generate_code()
|
||||||
'LLM': {
|
|
||||||
selected_llm: deepcopy(self.default_config['LLM'][selected_llm])
|
# Initialize device config with only necessary configurations
|
||||||
},
|
device_config = {
|
||||||
'TTS': {
|
'selected_module': deepcopy(selected_modules),
|
||||||
selected_tts: deepcopy(self.default_config['TTS'][selected_tts])
|
'prompt': self.default_config['prompt'],
|
||||||
},
|
'LLM': {
|
||||||
'ASR': {
|
selected_llm: deepcopy(self.default_config['LLM'][selected_llm])
|
||||||
selected_asr: deepcopy(self.default_config['ASR'][selected_asr])
|
},
|
||||||
},
|
'TTS': {
|
||||||
'VAD': {
|
selected_tts: deepcopy(self.default_config['TTS'][selected_tts])
|
||||||
selected_vad: deepcopy(self.default_config['VAD'][selected_vad])
|
},
|
||||||
|
'ASR': {
|
||||||
|
selected_asr: deepcopy(self.default_config['ASR'][selected_asr])
|
||||||
|
},
|
||||||
|
'VAD': {
|
||||||
|
selected_vad: deepcopy(self.default_config['VAD'][selected_vad])
|
||||||
|
},
|
||||||
|
'auth_code': auth_code # 添加认证码字段
|
||||||
}
|
}
|
||||||
}
|
|
||||||
|
all_configs[self.device_id] = device_config
|
||||||
all_configs[self.device_id] = device_config
|
|
||||||
|
# Save updated configs
|
||||||
# Save updated configs
|
with open(self.config_path, 'w', encoding='utf-8') as f:
|
||||||
with open(self.config_path, 'w', encoding='utf-8') as f:
|
yaml.dump(all_configs, f, allow_unicode=True)
|
||||||
yaml.dump(all_configs, f, allow_unicode=True)
|
|
||||||
|
|
||||||
self.private_config = all_configs[self.device_id]
|
self.private_config = all_configs[self.device_id]
|
||||||
|
|
||||||
|
finally:
|
||||||
|
self.lock_manager.release_lock(self.config_path)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.logger.error(f"Error handling private config: {e}")
|
self.logger.error(f"Error handling private config: {e}")
|
||||||
@@ -70,41 +86,45 @@ class PrivateConfig:
|
|||||||
bool: 更新是否成功
|
bool: 更新是否成功
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
# Read main config to get full module configurations
|
await self.lock_manager.acquire_lock(self.config_path)
|
||||||
main_config = self.default_config
|
try:
|
||||||
|
# Read main config to get full module configurations
|
||||||
|
main_config = self.default_config
|
||||||
|
|
||||||
# Create new device config
|
# Create new device config
|
||||||
device_config = {
|
device_config = {
|
||||||
'selected_module': selected_modules,
|
'selected_module': selected_modules,
|
||||||
'prompt': prompt,
|
'prompt': prompt,
|
||||||
'nickname': nickname,
|
'nickname': nickname,
|
||||||
}
|
}
|
||||||
if self.private_config.get('last_chat_time'):
|
if self.private_config.get('last_chat_time'):
|
||||||
device_config['last_chat_time'] = self.private_config['last_chat_time']
|
device_config['last_chat_time'] = self.private_config['last_chat_time']
|
||||||
|
|
||||||
# Copy full module configurations from main config
|
# Copy full module configurations from main config
|
||||||
for module_type, selected_name in selected_modules.items():
|
for module_type, selected_name in selected_modules.items():
|
||||||
if selected_name and selected_name in main_config.get(module_type, {}):
|
if selected_name and selected_name in main_config.get(module_type, {}):
|
||||||
device_config[module_type] = {
|
device_config[module_type] = {
|
||||||
selected_name: main_config[module_type][selected_name]
|
selected_name: main_config[module_type][selected_name]
|
||||||
}
|
}
|
||||||
|
|
||||||
# Read all configs
|
# Read all configs
|
||||||
if os.path.exists(self.config_path):
|
if os.path.exists(self.config_path):
|
||||||
with open(self.config_path, 'r', encoding='utf-8') as f:
|
with open(self.config_path, 'r', encoding='utf-8') as f:
|
||||||
all_configs = yaml.safe_load(f) or {}
|
all_configs = yaml.safe_load(f) or {}
|
||||||
else:
|
else:
|
||||||
all_configs = {}
|
all_configs = {}
|
||||||
|
|
||||||
# Update device config
|
# Update device config
|
||||||
all_configs[self.device_id] = device_config
|
all_configs[self.device_id] = device_config
|
||||||
self.private_config = device_config
|
self.private_config = device_config
|
||||||
|
|
||||||
# Save back to file
|
# Save back to file
|
||||||
with open(self.config_path, 'w', encoding='utf-8') as f:
|
with open(self.config_path, 'w', encoding='utf-8') as f:
|
||||||
yaml.dump(all_configs, f, allow_unicode=True)
|
yaml.dump(all_configs, f, allow_unicode=True)
|
||||||
|
|
||||||
return True
|
return True
|
||||||
|
finally:
|
||||||
|
self.lock_manager.release_lock(self.config_path)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.logger.error(f"Error updating config: {e}")
|
self.logger.error(f"Error updating config: {e}")
|
||||||
@@ -116,25 +136,29 @@ class PrivateConfig:
|
|||||||
bool: 删除是否成功
|
bool: 删除是否成功
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
# 读取所有配置
|
await self.lock_manager.acquire_lock(self.config_path)
|
||||||
if os.path.exists(self.config_path):
|
try:
|
||||||
with open(self.config_path, 'r', encoding='utf-8') as f:
|
# 读取所有配置
|
||||||
all_configs = yaml.safe_load(f) or {}
|
if os.path.exists(self.config_path):
|
||||||
else:
|
with open(self.config_path, 'r', encoding='utf-8') as f:
|
||||||
return False
|
all_configs = yaml.safe_load(f) or {}
|
||||||
|
else:
|
||||||
|
return False
|
||||||
|
|
||||||
# 删除设备配置
|
# 删除设备配置
|
||||||
if self.device_id in all_configs:
|
if self.device_id in all_configs:
|
||||||
del all_configs[self.device_id]
|
del all_configs[self.device_id]
|
||||||
|
|
||||||
|
# 保存更新后的配置
|
||||||
|
with open(self.config_path, 'w', encoding='utf-8') as f:
|
||||||
|
yaml.dump(all_configs, f, allow_unicode=True)
|
||||||
|
|
||||||
|
self.private_config = {}
|
||||||
|
return True
|
||||||
|
|
||||||
# 保存更新后的配置
|
return False
|
||||||
with open(self.config_path, 'w', encoding='utf-8') as f:
|
finally:
|
||||||
yaml.dump(all_configs, f, allow_unicode=True)
|
self.lock_manager.release_lock(self.config_path)
|
||||||
|
|
||||||
self.private_config = {}
|
|
||||||
return True
|
|
||||||
|
|
||||||
return False
|
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.logger.error(f"Error deleting config: {e}")
|
self.logger.error(f"Error deleting config: {e}")
|
||||||
@@ -176,7 +200,7 @@ class PrivateConfig:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
def update_last_chat_time(self, timestamp=None):
|
async def update_last_chat_time(self, timestamp=None):
|
||||||
"""更新设备最近一次的聊天时间
|
"""更新设备最近一次的聊天时间
|
||||||
Args:
|
Args:
|
||||||
timestamp: 指定的时间戳,不传则使用当前时间
|
timestamp: 指定的时间戳,不传则使用当前时间
|
||||||
@@ -186,24 +210,127 @@ class PrivateConfig:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if timestamp is None:
|
await self.lock_manager.acquire_lock(self.config_path)
|
||||||
timestamp = int(time.time())
|
try:
|
||||||
|
if timestamp is None:
|
||||||
|
timestamp = int(time.time())
|
||||||
|
|
||||||
|
self.private_config['last_chat_time'] = timestamp
|
||||||
|
|
||||||
self.private_config['last_chat_time'] = timestamp
|
# 读取所有配置
|
||||||
|
with open(self.config_path, 'r', encoding='utf-8') as f:
|
||||||
# 读取所有配置
|
all_configs = yaml.safe_load(f) or {}
|
||||||
with open(self.config_path, 'r', encoding='utf-8') as f:
|
|
||||||
all_configs = yaml.safe_load(f) or {}
|
# 更新当前设备配置
|
||||||
|
all_configs[self.device_id] = self.private_config
|
||||||
|
|
||||||
# 更新当前设备配置
|
# 保存回文件
|
||||||
all_configs[self.device_id] = self.private_config
|
with open(self.config_path, 'w', encoding='utf-8') as f:
|
||||||
|
yaml.dump(all_configs, f, allow_unicode=True)
|
||||||
# 保存回文件
|
|
||||||
with open(self.config_path, 'w', encoding='utf-8') as f:
|
return True
|
||||||
yaml.dump(all_configs, f, allow_unicode=True)
|
finally:
|
||||||
|
self.lock_manager.release_lock(self.config_path)
|
||||||
|
|
||||||
return True
|
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.logger.error(f"Error updating last chat time: {e}")
|
self.logger.error(f"Error updating last chat time: {e}")
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
def get_auth_code(self) -> str:
|
||||||
|
"""获取设备的认证码
|
||||||
|
Returns:
|
||||||
|
str: 认证码,如果没有返回空字符串
|
||||||
|
"""
|
||||||
|
return self.private_config.get('auth_code', '')
|
||||||
|
|
||||||
|
async def bind_user(self, username: str) -> bool:
|
||||||
|
"""绑定用户到设备"""
|
||||||
|
try:
|
||||||
|
await self.lock_manager.acquire_lock(self.config_path)
|
||||||
|
try:
|
||||||
|
# 检查用户是否存在
|
||||||
|
if not self.user_manager.get_user(username):
|
||||||
|
self.logger.error(f"User {username} not found")
|
||||||
|
return False
|
||||||
|
|
||||||
|
# 读取所有配置
|
||||||
|
with open(self.config_path, 'r', encoding='utf-8') as f:
|
||||||
|
all_configs = yaml.safe_load(f) or {}
|
||||||
|
|
||||||
|
if self.device_id not in all_configs:
|
||||||
|
self.logger.error(f"Device {self.device_id} not found")
|
||||||
|
return False
|
||||||
|
|
||||||
|
# 删除认证码
|
||||||
|
auth_code = all_configs[self.device_id].get('auth_code')
|
||||||
|
self.logger.info(f"Binding user {username} to device {self.device_id}")
|
||||||
|
if auth_code:
|
||||||
|
del all_configs[self.device_id]['auth_code']
|
||||||
|
|
||||||
|
if self.auth_code_gen:
|
||||||
|
self.auth_code_gen.remove_code(auth_code)
|
||||||
|
|
||||||
|
# 更新设备所有者
|
||||||
|
all_configs[self.device_id]['owner'] = username
|
||||||
|
self.private_config = all_configs[self.device_id]
|
||||||
|
|
||||||
|
# 更新用户的设备列表
|
||||||
|
user_data = await self.user_manager.get_user(username)
|
||||||
|
if 'devices' not in user_data:
|
||||||
|
user_data['devices'] = []
|
||||||
|
if self.device_id not in user_data['devices']:
|
||||||
|
user_data['devices'].append(self.device_id)
|
||||||
|
await self.user_manager.update_user(username, user_data)
|
||||||
|
|
||||||
|
# 保存配置
|
||||||
|
with open(self.config_path, 'w', encoding='utf-8') as f:
|
||||||
|
yaml.dump(all_configs, f, allow_unicode=True)
|
||||||
|
|
||||||
|
return True
|
||||||
|
finally:
|
||||||
|
self.lock_manager.release_lock(self.config_path)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
self.logger.error(f"Error binding user: {e}")
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def unbind_user(self) -> bool:
|
||||||
|
"""解绑设备当前用户"""
|
||||||
|
try:
|
||||||
|
await self.lock_manager.acquire_lock(self.config_path)
|
||||||
|
try:
|
||||||
|
if not self.private_config.get('owner'):
|
||||||
|
return True
|
||||||
|
|
||||||
|
username = self.private_config['owner']
|
||||||
|
|
||||||
|
# 从用户数据中移除设备
|
||||||
|
user_data = self.user_manager.get_user(username)
|
||||||
|
if user_data and 'devices' in user_data:
|
||||||
|
if self.device_id in user_data['devices']:
|
||||||
|
user_data['devices'].remove(self.device_id)
|
||||||
|
self.user_manager.update_user(username, user_data)
|
||||||
|
|
||||||
|
# 从设备配置中移除所有者
|
||||||
|
with open(self.config_path, 'r', encoding='utf-8') as f:
|
||||||
|
all_configs = yaml.safe_load(f) or {}
|
||||||
|
|
||||||
|
if self.device_id in all_configs:
|
||||||
|
if 'owner' in all_configs[self.device_id]:
|
||||||
|
del all_configs[self.device_id]['owner']
|
||||||
|
self.private_config = all_configs[self.device_id]
|
||||||
|
|
||||||
|
with open(self.config_path, 'w', encoding='utf-8') as f:
|
||||||
|
yaml.dump(all_configs, f, allow_unicode=True)
|
||||||
|
|
||||||
|
return True
|
||||||
|
finally:
|
||||||
|
self.lock_manager.release_lock(self.config_path)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
self.logger.error(f"Error unbinding user: {e}")
|
||||||
|
return False
|
||||||
|
|
||||||
|
def get_owner(self) -> Optional[str]:
|
||||||
|
"""获取设备当前所有者"""
|
||||||
|
return self.private_config.get('owner')
|
||||||
+59
-14
@@ -17,6 +17,7 @@ from concurrent.futures import ThreadPoolExecutor, TimeoutError
|
|||||||
from core.handle.audioHandle import handleAudioMessage, sendAudioMessage
|
from core.handle.audioHandle import handleAudioMessage, sendAudioMessage
|
||||||
from config.private_config import PrivateConfig
|
from config.private_config import PrivateConfig
|
||||||
from core.auth import AuthMiddleware, AuthenticationError
|
from core.auth import AuthMiddleware, AuthenticationError
|
||||||
|
from core.utils.auth_code_gen import AuthCodeGenerator # 添加导入
|
||||||
|
|
||||||
class ConnectionHandler:
|
class ConnectionHandler:
|
||||||
def __init__(self, config: Dict[str, Any], _vad, _asr, _llm, _tts):
|
def __init__(self, config: Dict[str, Any], _vad, _asr, _llm, _tts):
|
||||||
@@ -76,6 +77,9 @@ class ConnectionHandler:
|
|||||||
self.max_cmd_length = len(cmd)
|
self.max_cmd_length = len(cmd)
|
||||||
|
|
||||||
self.private_config = None
|
self.private_config = None
|
||||||
|
self.auth_code_gen = AuthCodeGenerator.get_instance()
|
||||||
|
self.is_device_verified = False # 添加设备验证状态标志
|
||||||
|
|
||||||
|
|
||||||
async def handle_connection(self, ws):
|
async def handle_connection(self, ws):
|
||||||
try:
|
try:
|
||||||
@@ -92,21 +96,30 @@ class ConnectionHandler:
|
|||||||
bUsePrivateConfig = self.config.get("use_private_config", False)
|
bUsePrivateConfig = self.config.get("use_private_config", False)
|
||||||
logging.info(f"bUsePrivateConfig: {bUsePrivateConfig}, device_id: {device_id}")
|
logging.info(f"bUsePrivateConfig: {bUsePrivateConfig}, device_id: {device_id}")
|
||||||
if bUsePrivateConfig and device_id:
|
if bUsePrivateConfig and device_id:
|
||||||
self.private_config = PrivateConfig(device_id, self.config)
|
try:
|
||||||
await self.private_config.load_or_create()
|
self.private_config = PrivateConfig(device_id, self.config, self.auth_code_gen)
|
||||||
# Create private instances using private config
|
await self.private_config.load_or_create()
|
||||||
vad, asr, llm, tts = self.private_config.create_private_instances()
|
# 判断是否已经绑定
|
||||||
if vad is not None and asr is not None and llm is not None and tts is not None:
|
owner = self.private_config.get_owner()
|
||||||
self.vad = vad
|
self.is_device_verified = owner is not None
|
||||||
self.asr = asr
|
|
||||||
self.llm = llm
|
if self.is_device_verified:
|
||||||
self.tts = tts
|
await self.private_config.update_last_chat_time()
|
||||||
|
|
||||||
self.logger.info(f"Loaded private config and instances for device {device_id}")
|
vad, asr, llm, tts = self.private_config.create_private_instances()
|
||||||
self.private_config.update_last_chat_time()
|
if all([vad, asr, llm, tts]):
|
||||||
else:
|
self.vad = vad
|
||||||
self.logger.error(f"Failed to load private config for device {device_id}")
|
self.asr = asr
|
||||||
|
self.llm = llm
|
||||||
|
self.tts = tts
|
||||||
|
self.logger.info(f"Loaded private config and instances for device {device_id}")
|
||||||
|
else:
|
||||||
|
self.logger.error(f"Failed to create instances for device {device_id}")
|
||||||
|
self.private_config = None
|
||||||
|
except Exception as e:
|
||||||
|
self.logger.error(f"Error initializing private config: {e}")
|
||||||
self.private_config = None
|
self.private_config = None
|
||||||
|
raise
|
||||||
|
|
||||||
# 认证通过,继续处理
|
# 认证通过,继续处理
|
||||||
self.websocket = ws
|
self.websocket = ws
|
||||||
@@ -153,8 +166,40 @@ class ConnectionHandler:
|
|||||||
date_time = time.strftime("%Y-%m-%d %H:%M", time.localtime())
|
date_time = time.strftime("%Y-%m-%d %H:%M", time.localtime())
|
||||||
self.prompt = self.prompt.replace("{date_time}", date_time)
|
self.prompt = self.prompt.replace("{date_time}", date_time)
|
||||||
self.dialogue.put(Message(role="system", content=self.prompt))
|
self.dialogue.put(Message(role="system", content=self.prompt))
|
||||||
|
|
||||||
|
async def _check_and_broadcast_auth_code(self):
|
||||||
|
"""检查设备绑定状态并广播认证码"""
|
||||||
|
if not self.private_config.get_owner():
|
||||||
|
auth_code = self.private_config.get_auth_code()
|
||||||
|
if auth_code:
|
||||||
|
# 发送验证码语音提示
|
||||||
|
text = f"请在后台输入验证码:{' '.join(auth_code)}"
|
||||||
|
self.recode_first_last_text(text)
|
||||||
|
future = self.executor.submit(self.speak_and_play, text)
|
||||||
|
self.tts_queue.put(future)
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
def isNeedAuth(self):
|
||||||
|
bUsePrivateConfig = self.config.get("use_private_config", False)
|
||||||
|
if not bUsePrivateConfig:
|
||||||
|
# 如果不使用私有配置,就不需要验证
|
||||||
|
return False
|
||||||
|
return not self.is_device_verified
|
||||||
|
|
||||||
def chat(self, query):
|
def chat(self, query):
|
||||||
|
# 如果设备未验证,就发送验证码
|
||||||
|
if self.isNeedAuth():
|
||||||
|
self.llm_finish_task = True
|
||||||
|
# 创建一个新的事件循环来运行异步函数
|
||||||
|
loop = asyncio.new_event_loop()
|
||||||
|
asyncio.set_event_loop(loop)
|
||||||
|
try:
|
||||||
|
loop.run_until_complete(self._check_and_broadcast_auth_code())
|
||||||
|
finally:
|
||||||
|
loop.close()
|
||||||
|
return True
|
||||||
|
|
||||||
self.dialogue.put(Message(role="user", content=query))
|
self.dialogue.put(Message(role="user", content=query))
|
||||||
response_message = []
|
response_message = []
|
||||||
start = 0
|
start = 0
|
||||||
|
|||||||
@@ -0,0 +1,97 @@
|
|||||||
|
import random
|
||||||
|
import threading
|
||||||
|
import time
|
||||||
|
from typing import Set
|
||||||
|
|
||||||
|
class AuthCodeGenerator:
|
||||||
|
_instance = None
|
||||||
|
_instance_lock = threading.Lock()
|
||||||
|
|
||||||
|
def __new__(cls):
|
||||||
|
if not cls._instance:
|
||||||
|
with cls._instance_lock:
|
||||||
|
if not cls._instance:
|
||||||
|
cls._instance = super(AuthCodeGenerator, cls).__new__(cls)
|
||||||
|
# 初始化随机种子
|
||||||
|
random.seed(time.time())
|
||||||
|
return cls._instance
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
# 确保 __init__ 只被调用一次
|
||||||
|
if not hasattr(self, '_initialized'):
|
||||||
|
self._used_codes: Set[str] = set()
|
||||||
|
self._code_timestamps = {}
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
self._code_timeout = 3 * 24 * 60 * 60
|
||||||
|
self._initialized = True
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_instance(cls):
|
||||||
|
"""获取AuthCodeGenerator的单例实例"""
|
||||||
|
return cls()
|
||||||
|
|
||||||
|
def generate_code(self) -> str:
|
||||||
|
"""
|
||||||
|
生成6位数字认证码,确保不重复
|
||||||
|
返回: 6位数字字符串
|
||||||
|
"""
|
||||||
|
with self._lock:
|
||||||
|
self._clean_expired_codes() # 清理过期code
|
||||||
|
while True:
|
||||||
|
# 使用时间戳和已用码数量作为种子,确保每次生成不同的随机数
|
||||||
|
seed = int(time.time() * 1000) + len(self._used_codes)
|
||||||
|
random.seed(seed)
|
||||||
|
|
||||||
|
# 生成6位随机数字
|
||||||
|
code = ''.join(str(random.randint(0, 9)) for _ in range(6))
|
||||||
|
|
||||||
|
# 检查是否已存在
|
||||||
|
if code not in self._used_codes:
|
||||||
|
self._used_codes.add(code)
|
||||||
|
self._code_timestamps[code] = time.time()
|
||||||
|
return code
|
||||||
|
|
||||||
|
def remove_code(self, code: str) -> bool:
|
||||||
|
"""
|
||||||
|
删除已使用的认证码
|
||||||
|
参数:
|
||||||
|
code: 要删除的认证码
|
||||||
|
返回:
|
||||||
|
bool: 删除成功返回True,码不存在返回False
|
||||||
|
"""
|
||||||
|
print('remove_code', code)
|
||||||
|
with self._lock:
|
||||||
|
if code in self._used_codes:
|
||||||
|
self._used_codes.remove(code)
|
||||||
|
if code in self._code_timestamps:
|
||||||
|
del self._code_timestamps[code]
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
def is_code_used(self, code: str) -> bool:
|
||||||
|
"""
|
||||||
|
检查认证码是否已被使用
|
||||||
|
参数:
|
||||||
|
code: 要检查的认证码
|
||||||
|
返回:
|
||||||
|
bool: 如果码存在返回True,否则返回False
|
||||||
|
"""
|
||||||
|
with self._lock:
|
||||||
|
return code in self._used_codes
|
||||||
|
|
||||||
|
def clear_codes(self):
|
||||||
|
"""清空所有已使用的认证码"""
|
||||||
|
with self._lock:
|
||||||
|
self._used_codes.clear()
|
||||||
|
self._code_timestamps.clear()
|
||||||
|
|
||||||
|
def _clean_expired_codes(self):
|
||||||
|
"""清理过期的认证码"""
|
||||||
|
current_time = time.time()
|
||||||
|
expired_codes = [
|
||||||
|
code for code, timestamp in self._code_timestamps.items()
|
||||||
|
if (current_time - timestamp) > self._code_timeout
|
||||||
|
]
|
||||||
|
for code in expired_codes:
|
||||||
|
self._used_codes.remove(code)
|
||||||
|
del self._code_timestamps[code]
|
||||||
@@ -0,0 +1,38 @@
|
|||||||
|
import asyncio
|
||||||
|
from typing import Dict
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
class FileLockManager:
|
||||||
|
_instance = None
|
||||||
|
_locks: Dict[str, asyncio.Lock] = {}
|
||||||
|
|
||||||
|
def __new__(cls):
|
||||||
|
if cls._instance is None:
|
||||||
|
cls._instance = super(FileLockManager, cls).__new__(cls)
|
||||||
|
return cls._instance
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_lock(cls, file_path: str) -> asyncio.Lock:
|
||||||
|
"""获取指定文件的锁"""
|
||||||
|
if file_path not in cls._locks:
|
||||||
|
cls._locks[file_path] = asyncio.Lock()
|
||||||
|
return cls._locks[file_path]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def acquire_lock(cls, file_path: str):
|
||||||
|
"""获取锁"""
|
||||||
|
lock = cls.get_lock(file_path)
|
||||||
|
await lock.acquire()
|
||||||
|
logger.debug(f"Acquired lock for {file_path}")
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def release_lock(cls, file_path: str):
|
||||||
|
"""释放锁"""
|
||||||
|
if file_path in cls._locks:
|
||||||
|
try:
|
||||||
|
cls._locks[file_path].release()
|
||||||
|
logger.debug(f"Released lock for {file_path}")
|
||||||
|
except RuntimeError as e:
|
||||||
|
logger.warning(f"Failed to release lock for {file_path}: {e}")
|
||||||
+107
-32
@@ -4,18 +4,22 @@ import logging
|
|||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
from core.utils.util import get_project_dir
|
from core.utils.util import get_project_dir
|
||||||
from config.private_config import PrivateConfig
|
from config.private_config import PrivateConfig
|
||||||
|
from manager.api.user_manager import UserManager # 添加导入
|
||||||
|
from core.utils.auth_code_gen import AuthCodeGenerator # 添加导入
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
class ConfigHandler:
|
class ConfigHandler:
|
||||||
def __init__(self):
|
def __init__(self, session_manager):
|
||||||
|
self.session_manager = session_manager
|
||||||
|
self.user_manager = UserManager() # 添加 user_manager 实例
|
||||||
self.private_config_path = get_project_dir() + 'data/.private_config.yaml'
|
self.private_config_path = get_project_dir() + 'data/.private_config.yaml'
|
||||||
self.config_path = get_project_dir() + 'config.yaml'
|
self.config_path = get_project_dir() + 'config.yaml'
|
||||||
# 如果存在.config.yaml文件,则使用该文件
|
# 如果存在.config.yaml文件,则使用该文件
|
||||||
if os.path.exists(get_project_dir() + "data/.config.yaml"):
|
if os.path.exists(get_project_dir() + "data/.config.yaml"):
|
||||||
self.config_path = get_project_dir() + "data/.config.yaml"
|
self.config_path = get_project_dir() + "data/.config.yaml"
|
||||||
with open(self.config_path, 'r', encoding='utf-8') as f:
|
with open(self.config_path, 'r', encoding='utf-8') as f:
|
||||||
self.config = yaml.safe_load(f)
|
self.config = yaml.safe_load(f)
|
||||||
|
|
||||||
async def get_module_options(self, request):
|
async def get_module_options(self, request):
|
||||||
"""Get all available module options from config.yaml"""
|
"""Get all available module options from config.yaml"""
|
||||||
@@ -45,46 +49,41 @@ class ConfigHandler:
|
|||||||
})
|
})
|
||||||
|
|
||||||
async def get_private_configs(self, request):
|
async def get_private_configs(self, request):
|
||||||
"""获取所有私有配置设备列表及其配置"""
|
"""只返回用户绑定的设备配置"""
|
||||||
try:
|
try:
|
||||||
|
username = request['username']
|
||||||
|
logger.info(f"Getting devices for user: {username}")
|
||||||
|
|
||||||
|
# 从用户管理器获取用户的设备列表
|
||||||
|
user_devices = await self.user_manager.get_user_devices(username)
|
||||||
|
logger.info(f"User {username} has devices: {user_devices}")
|
||||||
|
|
||||||
|
# 读取所有配置
|
||||||
|
all_configs = {}
|
||||||
if os.path.exists(self.private_config_path):
|
if os.path.exists(self.private_config_path):
|
||||||
with open(self.private_config_path, 'r', encoding='utf-8') as f:
|
with open(self.private_config_path, 'r', encoding='utf-8') as f:
|
||||||
all_configs = yaml.safe_load(f) or {}
|
all_configs = yaml.safe_load(f) or {}
|
||||||
else:
|
|
||||||
all_configs = {}
|
# 只返回用户有权限的设备配置
|
||||||
|
user_configs = {
|
||||||
# 转换配置为前端友好的格式
|
device_id: config
|
||||||
devices = []
|
for device_id, config in all_configs.items()
|
||||||
for device_id, config in all_configs.items():
|
if device_id in user_devices
|
||||||
device_info = {
|
}
|
||||||
'id': device_id,
|
|
||||||
'config': {
|
logger.info(f"Returning {len(user_configs)} device configs for user {username}")
|
||||||
'selected_module': config.get('selected_module', {}),
|
|
||||||
'prompt': config.get('prompt', ''),
|
|
||||||
'last_chat_time': config.get('last_chat_time', ''),
|
|
||||||
'nickname': config.get('nickname', '小智'),
|
|
||||||
'modules': {
|
|
||||||
'LLM': config.get('LLM', {}),
|
|
||||||
'TTS': config.get('TTS', {}),
|
|
||||||
'ASR': config.get('ASR', {}),
|
|
||||||
'VAD': config.get('VAD', {})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
devices.append(device_info)
|
|
||||||
|
|
||||||
return web.json_response({
|
return web.json_response({
|
||||||
'success': True,
|
'success': True,
|
||||||
'data': devices,
|
'data': user_configs,
|
||||||
'message': '获取成功'
|
'message': '获取成功'
|
||||||
})
|
})
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error getting private configs: {str(e)}", exc_info=True)
|
logger.error(f"Error getting devices for user {request.get('username')}: {str(e)}", exc_info=True)
|
||||||
return web.json_response({
|
return web.json_response({
|
||||||
'success': False,
|
'success': False,
|
||||||
'message': '获取配置失败'
|
'message': f'获取设备列表失败: {str(e)}'
|
||||||
})
|
}, status=400)
|
||||||
|
|
||||||
async def save_device_config(self, request):
|
async def save_device_config(self, request):
|
||||||
"""保存单个设备的配置"""
|
"""保存单个设备的配置"""
|
||||||
@@ -92,6 +91,16 @@ class ConfigHandler:
|
|||||||
data = await request.json()
|
data = await request.json()
|
||||||
device_id = data.get('id')
|
device_id = data.get('id')
|
||||||
config = data.get('config')
|
config = data.get('config')
|
||||||
|
username = request['username'] # 从请求中获取用户名
|
||||||
|
|
||||||
|
# 检查设备所有权
|
||||||
|
user_devices = self.user_manager.get_user_devices(username)
|
||||||
|
if device_id not in user_devices:
|
||||||
|
return web.json_response({
|
||||||
|
'success': False,
|
||||||
|
'message': '无权操作此设备'
|
||||||
|
}, status=403)
|
||||||
|
|
||||||
logger.info(f"Device config updated: {device_id} :\n{config}")
|
logger.info(f"Device config updated: {device_id} :\n{config}")
|
||||||
if not device_id or not config:
|
if not device_id or not config:
|
||||||
return web.json_response({
|
return web.json_response({
|
||||||
@@ -130,10 +139,20 @@ class ConfigHandler:
|
|||||||
try:
|
try:
|
||||||
data = await request.json()
|
data = await request.json()
|
||||||
device_id = data.get('device_id')
|
device_id = data.get('device_id')
|
||||||
|
username = request['username']
|
||||||
|
|
||||||
|
# 检查设备所有权
|
||||||
|
user_devices = await self.user_manager.get_user_devices(username)
|
||||||
|
if device_id not in user_devices:
|
||||||
|
return web.json_response({
|
||||||
|
'success': False,
|
||||||
|
'message': '无权删除此设备'
|
||||||
|
}, status=403)
|
||||||
|
|
||||||
# 使用PrivateConfig处理配置删除
|
# 使用PrivateConfig处理配置删除
|
||||||
private_config = PrivateConfig(device_id, self.config)
|
private_config = PrivateConfig(device_id, self.config)
|
||||||
success = await private_config.delete_config()
|
success = await private_config.delete_config()
|
||||||
|
await self.user_manager.remove_device(username, device_id)
|
||||||
|
|
||||||
if not success:
|
if not success:
|
||||||
raise Exception("Failed to delete device config")
|
raise Exception("Failed to delete device config")
|
||||||
@@ -149,3 +168,59 @@ class ConfigHandler:
|
|||||||
'success': False,
|
'success': False,
|
||||||
'message': f'删除配置失败: {str(e)}'
|
'message': f'删除配置失败: {str(e)}'
|
||||||
})
|
})
|
||||||
|
|
||||||
|
async def bind_device(self, request):
|
||||||
|
"""绑定设备到用户"""
|
||||||
|
try:
|
||||||
|
data = await request.json()
|
||||||
|
auth_code = data.get('auth_code')
|
||||||
|
username = request['username']
|
||||||
|
|
||||||
|
if not auth_code or len(auth_code) != 6:
|
||||||
|
return web.json_response({
|
||||||
|
'success': False,
|
||||||
|
'message': '请输入6位认证码'
|
||||||
|
}, status=400)
|
||||||
|
|
||||||
|
# 读取所有设备配置
|
||||||
|
with open(self.private_config_path, 'r', encoding='utf-8') as f:
|
||||||
|
all_configs = yaml.safe_load(f) or {}
|
||||||
|
|
||||||
|
# 查找匹配认证码的设备
|
||||||
|
device_found = None
|
||||||
|
for device_id, config in all_configs.items():
|
||||||
|
if config.get('auth_code') == auth_code and not config.get('owner'):
|
||||||
|
device_found = device_id
|
||||||
|
break
|
||||||
|
|
||||||
|
if not device_found:
|
||||||
|
return web.json_response({
|
||||||
|
'success': False,
|
||||||
|
'message': '认证码无效或设备已被绑定'
|
||||||
|
}, status=400)
|
||||||
|
|
||||||
|
# 使用 PrivateConfig 进行绑定
|
||||||
|
private_config = PrivateConfig(device_found, self.config, AuthCodeGenerator())
|
||||||
|
await private_config.load_or_create()
|
||||||
|
|
||||||
|
# 绑定设备到用户 - 修改为异步调用
|
||||||
|
success = await private_config.bind_user(username)
|
||||||
|
if success:
|
||||||
|
# 同时更新用户的设备列表 - 修改为异步调用
|
||||||
|
await self.user_manager.add_device(username, device_found)
|
||||||
|
return web.json_response({
|
||||||
|
'success': True,
|
||||||
|
'message': '设备绑定成功'
|
||||||
|
})
|
||||||
|
else:
|
||||||
|
return web.json_response({
|
||||||
|
'success': False,
|
||||||
|
'message': '设备绑定失败'
|
||||||
|
}, status=500)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error binding device: {str(e)}", exc_info=True)
|
||||||
|
return web.json_response({
|
||||||
|
'success': False,
|
||||||
|
'message': f'绑定设备失败: {str(e)}'
|
||||||
|
}, status=500)
|
||||||
|
|||||||
+10
-7
@@ -5,8 +5,9 @@ import datetime
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
class LoginHandler:
|
class LoginHandler:
|
||||||
def __init__(self, config):
|
def __init__(self, user_manager, session_manager):
|
||||||
self.config = config
|
self.user_manager = user_manager
|
||||||
|
self.session_manager = session_manager
|
||||||
|
|
||||||
async def handle_login(self, request):
|
async def handle_login(self, request):
|
||||||
"""处理登录请求"""
|
"""处理登录请求"""
|
||||||
@@ -22,8 +23,8 @@ class LoginHandler:
|
|||||||
'message': '用户名和密码不能为空'
|
'message': '用户名和密码不能为空'
|
||||||
})
|
})
|
||||||
|
|
||||||
stored_user = self.config['get_user'](username)
|
stored_user = await self.user_manager.get_user(username)
|
||||||
if not stored_user or stored_user['password'] != self.config['hash_password'](password):
|
if not stored_user or stored_user['password'] != self.user_manager.hash_password(password):
|
||||||
logger.warning(f"Failed login attempt for user {username} from {request.remote}")
|
logger.warning(f"Failed login attempt for user {username} from {request.remote}")
|
||||||
return web.json_response({
|
return web.json_response({
|
||||||
'success': False,
|
'success': False,
|
||||||
@@ -31,14 +32,16 @@ class LoginHandler:
|
|||||||
})
|
})
|
||||||
|
|
||||||
# 更新最后登录时间
|
# 更新最后登录时间
|
||||||
self.config['update_user'](username, {
|
await self.user_manager.update_user(username, {
|
||||||
'last_login': datetime.datetime.now().isoformat()
|
'last_login': datetime.datetime.now().isoformat()
|
||||||
})
|
})
|
||||||
|
|
||||||
logger.info(f"Successful login for user {username} from {request.remote}")
|
# 创建会话并返回session_id
|
||||||
|
session_id = self.session_manager.create_session(username)
|
||||||
return web.json_response({
|
return web.json_response({
|
||||||
'success': True,
|
'success': True,
|
||||||
'message': '登录成功'
|
'message': '登录成功',
|
||||||
|
'session_id': session_id
|
||||||
})
|
})
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
+13
-25
@@ -1,14 +1,13 @@
|
|||||||
import logging
|
import logging
|
||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
import datetime
|
import datetime
|
||||||
from core.utils.util import check_password
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class RegisterHandler:
|
class RegisterHandler:
|
||||||
def __init__(self, config):
|
def __init__(self, user_manager):
|
||||||
self.config = config
|
self.user_manager = user_manager
|
||||||
|
|
||||||
async def handle_register(self, request):
|
async def handle_register(self, request):
|
||||||
"""处理注册请求"""
|
"""处理注册请求"""
|
||||||
@@ -17,12 +16,6 @@ class RegisterHandler:
|
|||||||
username = data.get('username')
|
username = data.get('username')
|
||||||
password = data.get('password')
|
password = data.get('password')
|
||||||
|
|
||||||
if not check_password(password):
|
|
||||||
return web.json_response({
|
|
||||||
'success': False,
|
|
||||||
'message': '密码必须包含大小写字母、数字且长度至少8位'
|
|
||||||
})
|
|
||||||
|
|
||||||
if not username or not password:
|
if not username or not password:
|
||||||
logger.warning(f"Registration attempt with empty credentials from {request.remote}")
|
logger.warning(f"Registration attempt with empty credentials from {request.remote}")
|
||||||
return web.json_response({
|
return web.json_response({
|
||||||
@@ -30,28 +23,23 @@ class RegisterHandler:
|
|||||||
'message': '用户名和密码不能为空'
|
'message': '用户名和密码不能为空'
|
||||||
})
|
})
|
||||||
|
|
||||||
users = self.config.get('users', {})
|
# 检查用户是否已存在
|
||||||
# 由于现在所有用户都能看到所有设备,从安全角度上考虑,只允许注册一个用户
|
if await self.user_manager.get_user(username):
|
||||||
# 未来绑定设备功能完成后,再放开任意注册
|
|
||||||
if len(users) >= 1:
|
|
||||||
return web.json_response({
|
|
||||||
'success': False,
|
|
||||||
'message': '系统已经初始化过了,如果忘记了密码,请直接删除“.secrets.yaml”文件,删除后重启本服务'
|
|
||||||
})
|
|
||||||
|
|
||||||
if username in users:
|
|
||||||
logger.warning(f"Registration attempt with existing username {username} from {request.remote}")
|
logger.warning(f"Registration attempt with existing username {username} from {request.remote}")
|
||||||
return web.json_response({
|
return web.json_response({
|
||||||
'success': False,
|
'success': False,
|
||||||
'message': '用户名已存在'
|
'message': '用户名已存在'
|
||||||
})
|
})
|
||||||
|
|
||||||
# 存储新用户
|
# 创建用户
|
||||||
self.config['users'][username] = {
|
user_data = {
|
||||||
'password': self.config['hash_password'](password),
|
'username': username,
|
||||||
'created_at': datetime.datetime.now().isoformat()
|
'password': self.user_manager.hash_password(password),
|
||||||
|
'devices': [],
|
||||||
|
'created_at': datetime.datetime.now().isoformat(),
|
||||||
|
'last_login': ''
|
||||||
}
|
}
|
||||||
self.config['save_user_data']()
|
await self.user_manager.add_user(username, user_data)
|
||||||
|
|
||||||
logger.info(f"Successfully registered new user {username} from {request.remote}")
|
logger.info(f"Successfully registered new user {username} from {request.remote}")
|
||||||
return web.json_response({
|
return web.json_response({
|
||||||
@@ -60,7 +48,7 @@ class RegisterHandler:
|
|||||||
})
|
})
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Registration error: {str(e)}", exc_info=True)
|
logger.error(f"Register error: {str(e)}", exc_info=True)
|
||||||
return web.json_response({
|
return web.json_response({
|
||||||
'success': False,
|
'success': False,
|
||||||
'message': '注册失败,请稍后重试'
|
'message': '注册失败,请稍后重试'
|
||||||
|
|||||||
+128
-23
@@ -3,16 +3,15 @@ import yaml
|
|||||||
import hashlib
|
import hashlib
|
||||||
import logging
|
import logging
|
||||||
from core.utils.util import get_project_dir
|
from core.utils.util import get_project_dir
|
||||||
|
from core.utils.lock_manager import FileLockManager
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class UserManager:
|
class UserManager:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.secrets_path = get_project_dir() + 'data/.secrets.yaml'
|
self.secrets_path = get_project_dir() + 'data/.secrets.yaml'
|
||||||
self.users = {}
|
self.lock_manager = FileLockManager()
|
||||||
self.ensure_secrets_file()
|
self.ensure_secrets_file()
|
||||||
self.load_user_data()
|
|
||||||
|
|
||||||
def ensure_secrets_file(self):
|
def ensure_secrets_file(self):
|
||||||
"""确保 .secrets.yaml 文件存在"""
|
"""确保 .secrets.yaml 文件存在"""
|
||||||
@@ -28,24 +27,51 @@ class UserManager:
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Failed to create .secrets.yaml: {e}")
|
logger.error(f"Failed to create .secrets.yaml: {e}")
|
||||||
raise
|
raise
|
||||||
|
|
||||||
def load_user_data(self):
|
async def _load_user_data_internal(self):
|
||||||
"""加载用户数据"""
|
"""内部加载用户数据方法 - 不获取锁"""
|
||||||
try:
|
try:
|
||||||
with open(self.secrets_path, 'r', encoding='utf-8') as f:
|
with open(self.secrets_path, 'r', encoding='utf-8') as f:
|
||||||
data = yaml.safe_load(f) or {'users': {}}
|
data = yaml.safe_load(f) or {'users': {}}
|
||||||
self.users = data['users']
|
users = data['users']
|
||||||
logger.info("Successfully loaded user data")
|
logger.debug("Successfully loaded user data")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Failed to load user data: {e}")
|
logger.error(f"Failed to load user data: {e}")
|
||||||
self.users = {}
|
users = {}
|
||||||
|
return users
|
||||||
|
|
||||||
|
async def load_user_data(self):
|
||||||
|
"""加载用户数据"""
|
||||||
|
try:
|
||||||
|
await self.lock_manager.acquire_lock(self.secrets_path)
|
||||||
|
try:
|
||||||
|
users = await self._load_user_data_internal()
|
||||||
|
finally:
|
||||||
|
self.lock_manager.release_lock(self.secrets_path)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Failed to load user data: {e}")
|
||||||
|
users = {}
|
||||||
|
return users
|
||||||
|
|
||||||
def save_user_data(self):
|
async def _save_user_data_internal(self, users):
|
||||||
"""保存用户数据"""
|
"""内部保存用户数据方法 - 不获取锁"""
|
||||||
try:
|
try:
|
||||||
with open(self.secrets_path, 'w', encoding='utf-8') as f:
|
with open(self.secrets_path, 'w', encoding='utf-8') as f:
|
||||||
yaml.dump({'users': self.users}, f)
|
yaml.dump({'users': users}, f)
|
||||||
logger.info("Successfully saved user data")
|
logger.debug("Successfully saved user data")
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Failed to save user data: {e}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
async def save_user_data(self, users):
|
||||||
|
"""外部保存用户数据方法 - 获取锁"""
|
||||||
|
try:
|
||||||
|
await self.lock_manager.acquire_lock(self.secrets_path)
|
||||||
|
try:
|
||||||
|
await self._save_user_data_internal(users)
|
||||||
|
finally:
|
||||||
|
self.lock_manager.release_lock(self.secrets_path)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Failed to save user data: {e}")
|
logger.error(f"Failed to save user data: {e}")
|
||||||
raise
|
raise
|
||||||
@@ -54,16 +80,95 @@ class UserManager:
|
|||||||
"""密码哈希"""
|
"""密码哈希"""
|
||||||
return hashlib.sha256(password.encode()).hexdigest()
|
return hashlib.sha256(password.encode()).hexdigest()
|
||||||
|
|
||||||
def get_users(self):
|
async def get_users(self):
|
||||||
"""获取所有用户"""
|
"""异步获取所有用户"""
|
||||||
return self.users
|
users = await self.load_user_data() # 确保获取最新数据
|
||||||
|
return users
|
||||||
|
|
||||||
def get_user(self, username):
|
async def get_user(self, username):
|
||||||
"""获取指定用户"""
|
"""异步获取指定用户"""
|
||||||
return self.users.get(username)
|
users = await self.load_user_data() # 确保获取最新数据
|
||||||
|
return users.get(username)
|
||||||
|
|
||||||
def update_user(self, username, data):
|
async def add_user(self, username: str, user_data: dict):
|
||||||
|
"""异步添加新用户"""
|
||||||
|
try:
|
||||||
|
await self.lock_manager.acquire_lock(self.secrets_path)
|
||||||
|
try:
|
||||||
|
users = await self._load_user_data_internal() # 确保获取最新数据
|
||||||
|
if username in users:
|
||||||
|
raise ValueError("User already exists")
|
||||||
|
users[username] = user_data
|
||||||
|
await self._save_user_data_internal(users)
|
||||||
|
finally:
|
||||||
|
self.lock_manager.release_lock(self.secrets_path)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error adding user: {e}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
async def update_user(self, username, data):
|
||||||
"""更新用户数据"""
|
"""更新用户数据"""
|
||||||
if username in self.users:
|
try:
|
||||||
self.users[username].update(data)
|
await self.lock_manager.acquire_lock(self.secrets_path)
|
||||||
self.save_user_data()
|
try:
|
||||||
|
users = await self._load_user_data_internal() # 确保获取最新数据
|
||||||
|
if username in users:
|
||||||
|
users[username].update(data)
|
||||||
|
await self._save_user_data_internal(users)
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
finally:
|
||||||
|
self.lock_manager.release_lock(self.secrets_path)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error updating user: {e}")
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def get_user_devices(self, username: str) -> list:
|
||||||
|
"""获取用户的设备列表"""
|
||||||
|
user = await self.get_user(username)
|
||||||
|
print(user)
|
||||||
|
if user and user.get('devices'):
|
||||||
|
return user['devices']
|
||||||
|
return []
|
||||||
|
|
||||||
|
async def add_device(self, username: str, device_id: str) -> bool:
|
||||||
|
"""添加设备到用户的设备列表"""
|
||||||
|
try:
|
||||||
|
await self.lock_manager.acquire_lock(self.secrets_path)
|
||||||
|
try:
|
||||||
|
users = await self._load_user_data_internal() # 确保获取最新数据
|
||||||
|
user = users.get(username) # 直接从内存获取,因为已经有锁
|
||||||
|
if not user:
|
||||||
|
return False
|
||||||
|
|
||||||
|
if 'devices' not in user:
|
||||||
|
user['devices'] = []
|
||||||
|
|
||||||
|
if device_id not in user['devices']:
|
||||||
|
user['devices'].append(device_id)
|
||||||
|
await self._save_user_data_internal(users)
|
||||||
|
return True
|
||||||
|
finally:
|
||||||
|
self.lock_manager.release_lock(self.secrets_path)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error adding device: {e}")
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def remove_device(self, username: str, device_id: str) -> bool:
|
||||||
|
"""从用户的设备列表中移除设备"""
|
||||||
|
try:
|
||||||
|
await self.lock_manager.acquire_lock(self.secrets_path)
|
||||||
|
try:
|
||||||
|
users = await self._load_user_data_internal() # 确保获取最新数据
|
||||||
|
user = users.get(username) # 直接从内存获取,因为已经有锁
|
||||||
|
if user and 'devices' in user:
|
||||||
|
if device_id in user['devices']:
|
||||||
|
user['devices'].remove(device_id)
|
||||||
|
await self._save_user_data_internal(users)
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
finally:
|
||||||
|
self.lock_manager.release_lock(self.secrets_path)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error removing device: {e}")
|
||||||
|
return False
|
||||||
|
|||||||
+48
-18
@@ -13,26 +13,33 @@ from manager.api.login import LoginHandler
|
|||||||
from manager.api.register import RegisterHandler
|
from manager.api.register import RegisterHandler
|
||||||
from manager.api.user_manager import UserManager
|
from manager.api.user_manager import UserManager
|
||||||
from manager.api.config import ConfigHandler
|
from manager.api.config import ConfigHandler
|
||||||
|
from manager.session import SessionManager
|
||||||
|
from functools import wraps
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
def auth_required(handler):
|
||||||
|
"""鉴权装饰器"""
|
||||||
|
@wraps(handler)
|
||||||
|
async def wrapper(self, request):
|
||||||
|
session_id = request.cookies.get('session_id')
|
||||||
|
username = self.session_manager.validate_session(session_id)
|
||||||
|
if not username:
|
||||||
|
return web.json_response({'error': 'Unauthorized'}, status=401)
|
||||||
|
# 将用户名添加到请求对象
|
||||||
|
request['username'] = username
|
||||||
|
return await handler(self, request)
|
||||||
|
return wrapper
|
||||||
|
|
||||||
class WebUI:
|
class WebUI:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.app = web.Application()
|
self.app = web.Application()
|
||||||
self.user_manager = UserManager()
|
self.user_manager = UserManager()
|
||||||
|
self.session_manager = SessionManager()
|
||||||
|
|
||||||
# 添加静态文件路径
|
# 添加静态文件路径
|
||||||
self.static_path = os.path.join(root_dir, 'manager', 'static', 'webui')
|
self.static_path = os.path.join(root_dir, 'manager', 'static', 'webui')
|
||||||
|
|
||||||
# 创建配置字典
|
|
||||||
self.config = {
|
|
||||||
'users': self.user_manager.get_users(),
|
|
||||||
'hash_password': self.user_manager.hash_password,
|
|
||||||
'save_user_data': self.user_manager.save_user_data,
|
|
||||||
'get_user': self.user_manager.get_user,
|
|
||||||
'update_user': self.user_manager.update_user
|
|
||||||
}
|
|
||||||
|
|
||||||
self.setup_routes()
|
self.setup_routes()
|
||||||
self.setup_cors()
|
self.setup_cors()
|
||||||
|
|
||||||
@@ -52,18 +59,21 @@ class WebUI:
|
|||||||
|
|
||||||
def setup_routes(self):
|
def setup_routes(self):
|
||||||
"""设置路由"""
|
"""设置路由"""
|
||||||
login_handler = LoginHandler(self.config)
|
login_handler = LoginHandler(self.user_manager, self.session_manager)
|
||||||
register_handler = RegisterHandler(self.config)
|
register_handler = RegisterHandler(self.user_manager)
|
||||||
config_handler = ConfigHandler()
|
config_handler = ConfigHandler(self.session_manager)
|
||||||
|
|
||||||
# API 路由
|
# Public APIs
|
||||||
self.app.router.add_post('/api/login', login_handler.handle_login)
|
self.app.router.add_post('/api/login', login_handler.handle_login)
|
||||||
self.app.router.add_post('/api/register', register_handler.handle_register)
|
self.app.router.add_post('/api/register', register_handler.handle_register)
|
||||||
self.app.router.add_get('/api/config/devices', config_handler.get_private_configs)
|
|
||||||
self.app.router.add_post('/api/config/device', config_handler.save_device_config)
|
# Protected APIs
|
||||||
self.app.router.add_get('/api/config/module-options', config_handler.get_module_options)
|
self.app.router.add_get('/api/config/devices', self.auth_wrapper(config_handler.get_private_configs))
|
||||||
self.app.router.add_post('/api/config/save_device_config', config_handler.save_device_config)
|
|
||||||
self.app.router.add_post('/api/config/delete_device', config_handler.delete_device_config)
|
self.app.router.add_get('/api/config/module-options', self.auth_wrapper(config_handler.get_module_options))
|
||||||
|
self.app.router.add_post('/api/config/save_device_config', self.auth_wrapper(config_handler.save_device_config))
|
||||||
|
self.app.router.add_post('/api/config/delete_device', self.auth_wrapper(config_handler.delete_device_config))
|
||||||
|
self.app.router.add_post('/api/config/bind_device', self.auth_wrapper(config_handler.bind_device))
|
||||||
|
|
||||||
# 添加静态文件服务
|
# 添加静态文件服务
|
||||||
self.app.router.add_static('/assets/', path=os.path.join(self.static_path, 'assets'))
|
self.app.router.add_static('/assets/', path=os.path.join(self.static_path, 'assets'))
|
||||||
@@ -77,6 +87,26 @@ class WebUI:
|
|||||||
return web.FileResponse(index_file)
|
return web.FileResponse(index_file)
|
||||||
return web.Response(status=404, text='Not found')
|
return web.Response(status=404, text='Not found')
|
||||||
|
|
||||||
|
def auth_wrapper(self, handler):
|
||||||
|
"""包装处理器添加鉴权"""
|
||||||
|
@wraps(handler)
|
||||||
|
async def wrapper(request):
|
||||||
|
# 从请求头获取session_id
|
||||||
|
session_id = request.headers.get('Authorization')
|
||||||
|
if not session_id:
|
||||||
|
logger.warning("No session_id in Authorization header")
|
||||||
|
return web.json_response({'error': 'Unauthorized'}, status=401)
|
||||||
|
|
||||||
|
username = self.session_manager.validate_session(session_id)
|
||||||
|
if not username:
|
||||||
|
logger.warning(f"Invalid session_id: {session_id}")
|
||||||
|
return web.json_response({'error': 'Unauthorized'}, status=401)
|
||||||
|
|
||||||
|
request['username'] = username
|
||||||
|
logger.debug(f"Auth success for user: {username}")
|
||||||
|
return await handler(request)
|
||||||
|
return wrapper
|
||||||
|
|
||||||
def run(self, host='0.0.0.0', port=8002):
|
def run(self, host='0.0.0.0', port=8002):
|
||||||
"""运行服务器"""
|
"""运行服务器"""
|
||||||
local_ip = get_local_ip()
|
local_ip = get_local_ip()
|
||||||
|
|||||||
@@ -0,0 +1,33 @@
|
|||||||
|
import time
|
||||||
|
from typing import Dict, Optional
|
||||||
|
|
||||||
|
class SessionManager:
|
||||||
|
def __init__(self):
|
||||||
|
self.sessions: Dict[str, Dict] = {}
|
||||||
|
self.session_timeout = 24 * 60 * 60 # 24小时过期
|
||||||
|
|
||||||
|
def create_session(self, username: str) -> str:
|
||||||
|
"""创建新会话"""
|
||||||
|
session_id = str(hash(f"{username}:{time.time()}"))
|
||||||
|
self.sessions[session_id] = {
|
||||||
|
'username': username,
|
||||||
|
'created_at': time.time()
|
||||||
|
}
|
||||||
|
return session_id
|
||||||
|
|
||||||
|
def validate_session(self, session_id: str) -> Optional[str]:
|
||||||
|
"""验证会话是否有效,返回用户名"""
|
||||||
|
if session_id not in self.sessions:
|
||||||
|
return None
|
||||||
|
|
||||||
|
session = self.sessions[session_id]
|
||||||
|
if time.time() - session['created_at'] > self.session_timeout:
|
||||||
|
del self.sessions[session_id]
|
||||||
|
return None
|
||||||
|
|
||||||
|
return session['username']
|
||||||
|
|
||||||
|
def remove_session(self, session_id: str):
|
||||||
|
"""删除会话"""
|
||||||
|
if session_id in self.sessions:
|
||||||
|
del self.sessions[session_id]
|
||||||
Reference in New Issue
Block a user