mirror of
https://github.com/open-webui/open-webui.git
synced 2026-10-02 02:12:25 +00:00
refac
This commit is contained in:
parent
7fa8673296
commit
0180efecf3
3 changed files with 86 additions and 22 deletions
|
|
@ -113,6 +113,7 @@ from open_webui.env import (
|
|||
SCIM_TOKEN,
|
||||
VERSION,
|
||||
WEBSOCKET_HEARTBEAT_INTERVAL,
|
||||
WEBSOCKET_MANAGER,
|
||||
# Admin Account Runtime Creation
|
||||
WEBUI_ADMIN_EMAIL,
|
||||
WEBUI_ADMIN_NAME,
|
||||
|
|
@ -189,6 +190,7 @@ from open_webui.socket.main import (
|
|||
get_user_id_from_session_pool,
|
||||
periodic_session_pool_cleanup,
|
||||
periodic_usage_pool_cleanup,
|
||||
redis_event_listener,
|
||||
)
|
||||
from open_webui.socket.main import (
|
||||
app as socket_app,
|
||||
|
|
@ -387,6 +389,9 @@ async def lifespan(app: FastAPI):
|
|||
if app.state.redis is not None:
|
||||
app.state.redis_task_command_listener = asyncio.create_task(redis_task_command_listener(app))
|
||||
|
||||
if WEBSOCKET_MANAGER == 'redis':
|
||||
app.state.redis_event_listener = asyncio.create_task(redis_event_listener())
|
||||
|
||||
app.state.periodic_usage_pool_cleanup = asyncio.create_task(periodic_usage_pool_cleanup())
|
||||
app.state.periodic_session_pool_cleanup = asyncio.create_task(periodic_session_pool_cleanup())
|
||||
|
||||
|
|
@ -473,6 +478,9 @@ async def lifespan(app: FastAPI):
|
|||
if hasattr(app.state, 'redis_task_command_listener'):
|
||||
app.state.redis_task_command_listener.cancel()
|
||||
|
||||
if hasattr(app.state, 'redis_event_listener'):
|
||||
app.state.redis_event_listener.cancel()
|
||||
|
||||
app.state.periodic_usage_pool_cleanup.cancel()
|
||||
app.state.periodic_session_pool_cleanup.cancel()
|
||||
app.state.scheduler_worker_loop.cancel()
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import logging
|
|||
import random
|
||||
import sys
|
||||
import time
|
||||
from contextlib import suppress
|
||||
from typing import Any
|
||||
|
||||
import pycrdt as Y
|
||||
|
|
@ -38,17 +39,23 @@ from open_webui.models.folders import Folders
|
|||
from open_webui.models.notes import Notes, NoteUpdateForm
|
||||
from open_webui.models.users import UserNameResponse, Users
|
||||
from open_webui.socket.utils import RedisDict, RedisLock, YdocManager
|
||||
from open_webui.tasks import create_task, stop_item_tasks
|
||||
from open_webui.tasks import (
|
||||
REDIS_PUBSUB_MAX_RECONNECT_INTERVAL,
|
||||
REDIS_PUBSUB_RECONNECT_INTERVAL,
|
||||
create_task,
|
||||
stop_item_tasks,
|
||||
)
|
||||
from open_webui.utils.access_control import has_permission
|
||||
from open_webui.utils.auth import get_verified_user_by_token
|
||||
from open_webui.utils.chat_id import is_saved_chat_id
|
||||
from open_webui.utils.json_codec import SOCKETIO_JSON
|
||||
from open_webui.utils.json_codec import SOCKETIO_JSON, JSONCodec, dumps_bytes
|
||||
from open_webui.utils.misc import get_output_text
|
||||
from open_webui.utils.redis import (
|
||||
build_sentinel_url,
|
||||
get_redis_connection,
|
||||
get_sentinels_from_env,
|
||||
)
|
||||
from redis.exceptions import RedisError
|
||||
from socketio.packet import Packet
|
||||
|
||||
logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL)
|
||||
|
|
@ -190,6 +197,11 @@ YDOC_MANAGER = YdocManager(
|
|||
redis_key_prefix=f'{REDIS_KEY_PREFIX}:ydoc:documents',
|
||||
)
|
||||
|
||||
REDIS_EVENT_CHANNEL = f'{REDIS_KEY_PREFIX}:direct_completion'
|
||||
|
||||
EVENT_QUEUES: dict[str, asyncio.Queue] = {}
|
||||
EVENT_PUBLISH_LOCK = asyncio.Lock()
|
||||
|
||||
|
||||
def get_session_pool_batches():
|
||||
"""All session pool entries, in bounded batches for the Redis backing."""
|
||||
|
|
@ -948,6 +960,62 @@ async def disconnect(sid, reason=None):
|
|||
# print(f"Unknown session ID {sid} disconnected")
|
||||
|
||||
|
||||
async def redis_event_listener() -> None:
|
||||
"""Route events received over Redis to their local queues."""
|
||||
reconnect_interval = REDIS_PUBSUB_RECONNECT_INTERVAL
|
||||
|
||||
while True:
|
||||
pubsub = None
|
||||
try:
|
||||
# RedisCluster can't route a pubsub subscribe until initialize() fills its slot cache.
|
||||
await REDIS.initialize()
|
||||
|
||||
pubsub = REDIS.pubsub()
|
||||
await pubsub.subscribe(REDIS_EVENT_CHANNEL)
|
||||
reconnect_interval = REDIS_PUBSUB_RECONNECT_INTERVAL
|
||||
|
||||
async for message in pubsub.listen():
|
||||
if message['type'] != 'message':
|
||||
continue
|
||||
event = JSONCodec.loads(message['data'])
|
||||
queue = EVENT_QUEUES.get(event['channel'])
|
||||
if queue is not None:
|
||||
await queue.put(event['data'])
|
||||
log.warning('Redis event listener stopped. Retrying.')
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception:
|
||||
log.exception('Redis event listener failed. Retrying.')
|
||||
finally:
|
||||
if pubsub:
|
||||
with suppress(Exception):
|
||||
await pubsub.aclose()
|
||||
|
||||
await asyncio.sleep(reconnect_interval)
|
||||
reconnect_interval = min(reconnect_interval * 2, REDIS_PUBSUB_MAX_RECONNECT_INTERVAL)
|
||||
|
||||
|
||||
@sio.on('*')
|
||||
async def socket_event_handler(event: Any, sid: str, *args: Any) -> None:
|
||||
"""Route user-owned stream events to a local queue or another worker."""
|
||||
if not isinstance(event, str) or event.count(':') != 2 or not args:
|
||||
return
|
||||
|
||||
user = await get_socket_session_user(sid)
|
||||
if not user or user.get('id') != event.split(':', 1)[0]:
|
||||
return
|
||||
|
||||
queue = EVENT_QUEUES.get(event)
|
||||
if queue is not None:
|
||||
await queue.put(args[0])
|
||||
elif WEBSOCKET_MANAGER == 'redis':
|
||||
try:
|
||||
async with EVENT_PUBLISH_LOCK:
|
||||
await REDIS.publish(REDIS_EVENT_CHANNEL, dumps_bytes({'channel': event, 'data': args[0]}))
|
||||
except RedisError as e:
|
||||
log.debug('Failed to relay socket event %s: %s', event, e)
|
||||
|
||||
|
||||
async def _make_channel_emitter(request_info):
|
||||
"""Event emitter that routes pipeline output to a channel message.
|
||||
|
||||
|
|
|
|||
|
|
@ -23,9 +23,9 @@ from open_webui.routers.pipelines import (
|
|||
process_pipeline_outlet_filter,
|
||||
)
|
||||
from open_webui.socket.main import (
|
||||
EVENT_QUEUES,
|
||||
get_event_call,
|
||||
get_event_emitter,
|
||||
sio,
|
||||
)
|
||||
from open_webui.utils.filter import (
|
||||
get_filter_functions,
|
||||
|
|
@ -71,19 +71,8 @@ async def generate_direct_chat_completion(
|
|||
logging.info('WebSocket channel: %s', channel)
|
||||
|
||||
if form_data.get('stream'):
|
||||
q = asyncio.Queue()
|
||||
|
||||
async def message_listener(sid, data):
|
||||
"""
|
||||
Handle received socket messages and push them into the queue.
|
||||
"""
|
||||
await q.put(data)
|
||||
|
||||
def remove_message_listener():
|
||||
sio.handlers['/'].pop(channel, None)
|
||||
|
||||
# Register the listener
|
||||
sio.on(channel, message_listener)
|
||||
queue = asyncio.Queue()
|
||||
EVENT_QUEUES[channel] = queue
|
||||
|
||||
# Start processing chat completion in background
|
||||
try:
|
||||
|
|
@ -103,16 +92,15 @@ async def generate_direct_chat_completion(
|
|||
|
||||
status = res.get('status', False)
|
||||
except BaseException:
|
||||
remove_message_listener()
|
||||
EVENT_QUEUES.pop(channel, None)
|
||||
raise
|
||||
|
||||
if status:
|
||||
# Define a generator to stream responses
|
||||
async def event_generator():
|
||||
nonlocal q
|
||||
try:
|
||||
while True:
|
||||
data = await q.get() # Wait for new messages
|
||||
data = await queue.get() # Wait for new messages
|
||||
if isinstance(data, dict):
|
||||
if 'done' in data and data['done']:
|
||||
break # Stop streaming when 'done' is received
|
||||
|
|
@ -127,16 +115,16 @@ async def generate_direct_chat_completion(
|
|||
log.debug('Error in event generator: %s', e)
|
||||
pass
|
||||
finally:
|
||||
remove_message_listener()
|
||||
EVENT_QUEUES.pop(channel, None)
|
||||
|
||||
# Define a background task to run the event generator
|
||||
async def background():
|
||||
remove_message_listener()
|
||||
EVENT_QUEUES.pop(channel, None)
|
||||
|
||||
# Return the streaming response
|
||||
return StreamingResponse(event_generator(), media_type='text/event-stream', background=background)
|
||||
else:
|
||||
remove_message_listener()
|
||||
EVENT_QUEUES.pop(channel, None)
|
||||
raise Exception(str(res))
|
||||
else:
|
||||
res = await event_caller(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue