1+ import asyncio
2+ import json
3+ # 设置 websockets 的日志级别为 NONE
4+ import logging
5+ from collections import defaultdict
6+ from datetime import datetime , timedelta
7+
8+ from bisheng .api .v1 .schemas import resp_500
19from bisheng .database .base import session_getter
210from bisheng .database .models .message import ChatMessage
11+ from pydantic import BaseModel
12+ from websockets import connect
13+
14+ # 维护一个连接池
15+ connection_pool = defaultdict (asyncio .Queue )
16+ logging .getLogger ('websockets' ).setLevel (logging .ERROR )
17+
18+ expire = 600 # reids 60s 过期
19+
20+
21+ class TimedQueue :
22+
23+ def __init__ (self ):
24+ self .queue = asyncio .Queue ()
25+ self .last_active = datetime .now ()
26+
27+ async def put_nowait (self , item ):
28+ self .last_active = datetime .now ()
29+ await self .queue .put (item )
30+
31+ async def get_nowait (self ):
32+ self .last_active = datetime .now ()
33+ return await self .queue .get ()
34+
35+ def empty (self ):
36+ return self .queue .empty ()
37+
38+ def qsize (self ):
39+ return self .queue .qsize ()
40+
41+
42+ async def clean_inactive_queues (queue : defaultdict , timeout_threshold : timedelta ):
43+ while True :
44+ current_time = datetime .now ()
45+ for key , timed_queue in list (queue .items ()):
46+ # 如果队列超过设定的阈值时间没有活跃,则清理队列
47+ if current_time - timed_queue .last_active > timeout_threshold :
48+ while not timed_queue .empty ():
49+ timed_queue .get_nowait () # 从队列中移除任务
50+ del queue [key ] # 删除队列
51+ await asyncio .sleep (timeout_threshold .total_seconds ())
52+
53+
54+ # 维护一个连接池
55+ connection_pool = defaultdict (TimedQueue )
56+ clean_inactive_queues (connection_pool , timedelta (minutes = 5 ))
57+
58+
59+ async def get_connection (uri , identifier ):
60+ """
61+ 获取WebSocket连接。如果连接池中有可用的连接,则直接返回;
62+ 否则,创建新的连接并添加到连接池。
63+ """
64+ if connection_pool [identifier ].empty ():
65+ # 建立新的WebSocket连接
66+ websocket = await connect (uri )
67+
68+ await connection_pool [identifier ].put_nowait (websocket )
69+
70+ # 从连接池中获取连接
71+ websocket = await connection_pool [identifier ].get_nowait ()
72+ return websocket
73+
74+
75+ async def release_connection (identifier , websocket ):
76+ """
77+ 释放WebSocket连接,将其放回连接池。
78+ """
79+ await connection_pool [identifier ].put_nowait (websocket )
380
481
582def comment_answer (message_id : int , comment : str ):
@@ -9,3 +86,68 @@ def comment_answer(message_id: int, comment: str):
986 message .remark = comment [:4096 ]
1087 session .add (message )
1188 session .commit ()
89+
90+
91+ class ContentStreamResp (BaseModel ):
92+ role : str
93+ content : str
94+
95+
96+ class ChoiceStreamResp (BaseModel ):
97+ index : int = 0
98+ delta : ContentStreamResp = 0
99+ session_id : str
100+
101+ def __str__ (self ) -> str :
102+ jsonData = '{"index": "%s", "delta": %s, "session_id": "%s"}' % (
103+ self .index , json .dumps (self .delta .dict (), ensure_ascii = False ), self .session_id )
104+ return '{"choices":[%s]}\n \n ' % (jsonData )
105+
106+
107+ async def event_stream (
108+ webosocket : connect ,
109+ message : str ,
110+ session_id : str ,
111+ model : str ,
112+ streaming : bool ,
113+ ):
114+
115+ payload = {'inputs' : message , 'flow_id' : model , 'chat_id' : session_id }
116+ try :
117+ await webosocket .send (json .dumps (payload , ensure_ascii = False ))
118+ except Exception as e :
119+ yield json .dumps (resp_500 (message = str (e )).__dict__ )
120+ return
121+ sync = ''
122+ while True :
123+ try :
124+ msg = await webosocket .recv ()
125+ except Exception as e :
126+ yield json .dumps (resp_500 (message = str (e )).__dict__ )
127+ break
128+ if msg is None :
129+ continue
130+ # 判断msg 的类型
131+ res = json .loads (msg )
132+ if streaming :
133+ if res .get ('type' ) != 'end' and res .get ('message' ):
134+ delta = ContentStreamResp (role = 'assistant' , content = res .get ('message' ))
135+ yield str (ChoiceStreamResp (index = 0 , session_id = session_id , delta = delta ))
136+ else :
137+ # 通过此处控制下面的close是否发送消息
138+ if res .get ('type' ) == 'end' :
139+ sync = res .get ('message' )
140+
141+ if res .get ('type' ) == 'close' :
142+ if not streaming and sync :
143+ delta = ContentStreamResp (role = 'assistant' , content = sync )
144+ msg = ChoiceStreamResp (index = 0 ,
145+ session_id = session_id ,
146+ delta = delta ,
147+ finish_reason = 'stop' )
148+ yield '{"choices":[%s]}' % (json .dumps (msg .dict ()))
149+ # 释放连接
150+ elif streaming :
151+ yield 'data: [DONE]'
152+ await release_connection (session_id , webosocket )
153+ break
0 commit comments