diff --git a/app.py b/app.py index 70499281..d4833195 100644 --- a/app.py +++ b/app.py @@ -2,15 +2,32 @@ import asyncio from config.logger import setup_logging from config.settings import load_config from core.server import WebSocketServer - +from core.http_server import ConfigServer async def main(): setup_logging() # 最先初始化日志 config = load_config() - - server = WebSocketServer(config) - await server.start() - + + # 启动 WebSocket 服务器 + ws_server = WebSocketServer(config) + ws_task = asyncio.create_task(ws_server.start()) + + # 启动 HTTP 配置服务器 + http_runner = None + if config['server'].get('http', {}).get('enabled', False): + config_server = ConfigServer(config) + try: + http_runner = await config_server.start() + except Exception as e: + print(f"Failed to start HTTP server: {e}") + + try: + # 等待 WebSocket 服务器运行 + await ws_task + finally: + # 清理 HTTP 服务器 + if http_runner: + await http_runner.cleanup() if __name__ == "__main__": asyncio.run(main()) diff --git a/config.yaml b/config.yaml index 7a9aebce..a4bd0e02 100644 --- a/config.yaml +++ b/config.yaml @@ -17,6 +17,11 @@ server: # 可选:设备白名单,如果设置了白名单,那么白名单的机器无论是什么token都可以连接。 #allowed_devices: # - "24:0A:C4:1D:3B:F0" # MAC地址列表 + http: + enabled: false + ip: 0.0.0.0 + port: 8001 + token: password xiaozhi: type: hello diff --git a/core/http_server.py b/core/http_server.py new file mode 100644 index 00000000..1cee536b --- /dev/null +++ b/core/http_server.py @@ -0,0 +1,145 @@ +import logging +import yaml +from aiohttp import web +from core.utils.util import get_project_dir + +logger = logging.getLogger(__name__) + +class ConfigServer: + def __init__(self, config: dict): + self.config = config + self.app = web.Application() + self.setup_routes() + + def setup_routes(self): + self.app.router.add_get('/', self.index) + self.app.router.add_get('/api/prompt', self.get_prompt) + self.app.router.add_post('/api/prompt', self.update_prompt) + self.app.router.add_options('/api/prompt', self.handle_options) + + async def index(self, request): + html = ''' + + + + 小智提示词配置 + + + + +

小智提示词配置

+ +
+ + +
+ + + + ''' + return web.Response(text=html, content_type='text/html') + + async def handle_options(self, request): + headers = { + 'Access-Control-Allow-Origin': '*', + 'Access-Control-Allow-Methods': 'GET, POST, OPTIONS', + 'Access-Control-Allow-Headers': 'Content-Type, Authorization' + } + return web.Response(headers=headers) + + async def verify_token(self, request): + if 'token' not in self.config['server']['http']: + return True + + expected_token = self.config['server']['http']['token'] + token = request.headers.get('Authorization', '').replace('Bearer ', '') + + if not token or token != expected_token: + return False + return True + + async def get_prompt(self, request): + if not await self.verify_token(request): + return web.json_response({'error': 'Unauthorized'}, status=401) + + headers = {'Access-Control-Allow-Origin': '*'} + return web.json_response({ + 'prompt': self.config['prompt'] + }, headers=headers) + + async def update_prompt(self, request): + if not await self.verify_token(request): + return web.json_response({'error': 'Unauthorized'}, status=401) + + try: + data = await request.json() + if 'prompt' not in data: + return web.json_response({'error': 'Missing prompt field'}, status=400) + + # 更新内存中的配置 + self.config['prompt'] = data['prompt'] + + # 更新配置文件 + config_path = get_project_dir() + 'config.yaml' + with open(config_path, 'w', encoding='utf-8') as f: + yaml.dump(self.config, f, allow_unicode=True) + + headers = {'Access-Control-Allow-Origin': '*'} + return web.json_response({'success': True}, headers=headers) + + except Exception as e: + logger.error(f"Failed to update prompt: {e}") + return web.json_response({'error': str(e)}, status=500) + + async def start(self): + try: + http_config = self.config['server']['http'] + if not http_config.get('enabled', False): + logger.info("HTTP server is disabled") + return + + runner = web.AppRunner(self.app) + await runner.setup() + site = web.TCPSite(runner, http_config['ip'], http_config['port']) + await site.start() + logger.info(f"Config HTTP server is running at http://{http_config['ip']}:{http_config['port']}") + return runner # 返回runner以便后续清理 + except Exception as e: + logger.error(f"Failed to start HTTP server: {e}") + raise diff --git a/requirements.txt b/requirements.txt index 9fa53a9c..77523f3e 100755 --- a/requirements.txt +++ b/requirements.txt @@ -11,4 +11,4 @@ openai==1.61.0 google-generativeai==0.8.4 edge_tts==7.0.0 httpx==0.27.2 -google-generativeai==0.8.4 \ No newline at end of file +aiohttp==3.9.3 \ No newline at end of file