- refactored
This commit is contained in:
@@ -0,0 +1,39 @@
|
||||
import asyncio
|
||||
import websockets
|
||||
import abc
|
||||
|
||||
|
||||
class WsServer:
|
||||
def __init__(self, loop=None):
|
||||
if loop is None:
|
||||
self.loop = asyncio.get_event_loop()
|
||||
else:
|
||||
self.loop = loop
|
||||
|
||||
async def listen(self, host, port):
|
||||
return await websockets.serve(self.handler, host, port, loop=self.loop)
|
||||
|
||||
def disconnect(self):
|
||||
pass
|
||||
|
||||
async def handler(self, websocket, path):
|
||||
print("handler: got connection from {}".format(websocket.remote_address))
|
||||
await self.register(websocket)
|
||||
try:
|
||||
consumer_task = asyncio.ensure_future(self.handler_recv(websocket, path), loop=self.loop)
|
||||
producer_task = asyncio.ensure_future(self.handler_send(websocket, path), loop=self.loop)
|
||||
done, pending = await asyncio.wait([consumer_task, producer_task], return_when=asyncio.FIRST_COMPLETED,)
|
||||
for task in pending:
|
||||
task.cancel()
|
||||
finally:
|
||||
print("handler: lost connection from {}".format(websocket.remote_address))
|
||||
await self.unregister(websocket)
|
||||
|
||||
@abc.abstractmethod
|
||||
async def handler_recv(self, websocket, path):
|
||||
pass
|
||||
|
||||
@abc.abstractmethod
|
||||
async def handler_send(self, websocket, path):
|
||||
pass
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
import asyncio
|
||||
from ws.server.ws_server import WsServer
|
||||
from ws.user import User, UserSet, update
|
||||
import abc
|
||||
from ws.connection import IConnection
|
||||
|
||||
|
||||
class IWsServer:
|
||||
@abc.abstractmethod
|
||||
def on_recv(self, data, loop):
|
||||
pass
|
||||
|
||||
@abc.abstractmethod
|
||||
def on_send(self, loop):
|
||||
pass
|
||||
|
||||
|
||||
class WsServerMultiUser(WsServer):
|
||||
def __init__(self, loop=None, listener: IConnection = None):
|
||||
WsServer.__init__(self, loop)
|
||||
self.listener = listener
|
||||
self.USERS = UserSet()
|
||||
self.global_state = {}
|
||||
|
||||
async def notify_state(self, data):
|
||||
if self.USERS: # asyncio.wait doesn't accept an empty list
|
||||
await asyncio.wait([user.send(data) for user in self.USERS])
|
||||
|
||||
async def notify_users(self):
|
||||
if self.USERS: # asyncio.wait doesn't accept an empty list
|
||||
message = {"info": {"type": "users", "count": len(self.USERS)}}
|
||||
await asyncio.wait([user.send(message) for user in self.USERS])
|
||||
|
||||
async def register(self, websocket):
|
||||
usr = User(websocket)
|
||||
usr.path_add("info")
|
||||
self.USERS.add(websocket, usr)
|
||||
await self.notify_users()
|
||||
|
||||
async def unregister(self, websocket):
|
||||
self.USERS.remove(websocket)
|
||||
await self.notify_users()
|
||||
|
||||
async def handler_recv(self, websocket, path):
|
||||
while True:
|
||||
try:
|
||||
data = await websocket.recv()
|
||||
usr = self.USERS.get(websocket)
|
||||
processed = await usr.process(data)
|
||||
if processed:
|
||||
await usr.send(self.global_state)
|
||||
else:
|
||||
await self.listener.on_recv(data)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
break
|
||||
|
||||
async def handler_send(self, websocket, path):
|
||||
while True:
|
||||
data = await self.listener.on_send()
|
||||
try:
|
||||
await self.notify_state(data)
|
||||
# Update global state
|
||||
self.global_state = update(self.global_state, data)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
break
|
||||
Reference in New Issue
Block a user