## BaraDB WebSocket Server — real-time subscriptions import std/asyncdispatch import std/asyncnet import std/strutils import std/tables import std/base64 import std/sets import config import jwt as jwtlib type WsFrame = object fin: bool opcode: uint8 masked: bool payloadLen: uint64 maskKey: array[4, byte] payload: string WsClient* = ref object socket: AsyncSocket id: int subscriptions: HashSet[string] WsServer* = ref object clients*: Table[int, WsClient] nextId: int running: bool config*: BaraConfig secretKey*: string onInsert*: proc (table, key, value: string) {.closure.} onDelete*: proc (table, key: string) {.closure.} proc newWsServer*(cfg: BaraConfig = defaultConfig(), secret: string = ""): WsServer = WsServer(clients: initTable[int, WsClient](), nextId: 1, running: false, config: cfg, secretKey: secret) # ---------------------------------------------------------------------- # WebSocket frame encoding/decoding (RFC 6455) # ---------------------------------------------------------------------- proc encodeFrame(opcode: uint8, payload: string): string = result = "" let isMasked = false var b0 = 0x80'u8 or opcode result.add(char(b0)) var b1 = 0'u8 if not isMasked: if payload.len < 126: b1 = uint8(payload.len) elif payload.len <= 65535: b1 = 126 else: b1 = 127 result.add(char(b1)) if payload.len >= 126 and payload.len <= 65535: var len16 = uint16(payload.len) result.add(char((len16 shr 8) and 0xFF)) result.add(char(len16 and 0xFF)) elif payload.len > 65535: var len64 = uint64(payload.len) for i in countdown(7, 0): result.add(char((len64 shr (i * 8)) and 0xFF)) result.add(payload) proc decodeFrame(data: string): (WsFrame, int) = if data.len < 2: return (WsFrame(), 0) var frame = WsFrame() let b0 = uint8(data[0]) let b1 = uint8(data[1]) frame.fin = (b0 and 0x80) != 0 frame.opcode = b0 and 0x0F frame.masked = (b1 and 0x80) != 0 var len = uint64(b1 and 0x7F) var offset = 2 if len == 126: if data.len < 4: return (WsFrame(), 0) len = (uint64(uint8(data[2])) shl 8) or uint64(uint8(data[3])) offset = 4 elif len == 127: if data.len < 10: return (WsFrame(), 0) len = 0 for i in 0..7: len = (len shl 8) or uint64(uint8(data[2 + i])) offset = 10 if frame.masked: if data.len < offset + 4: return (WsFrame(), 0) for i in 0..3: frame.maskKey[i] = byte(data[offset + i]) offset += 4 if uint64(data.len) < uint64(offset) + len: return (Wsframe(), 0) let plen = int(len) if frame.masked: for i in 0..= 2: let (frame, consumed) = decodeFrame(buf) if consumed == 0: break case frame.opcode of 0x8: # close client.close() server.clients.del(id) return of 0x9: # ping let pong = encodeFrame(0xA, frame.payload) await client.send(pong) of 0x1: # text let msg = frame.payload if msg.startsWith("SUBSCRIBE "): let table = msg[10..^1].strip() wsClient.subscribe(table) let ack = encodeFrame(0x1, "OK subscribed to " & table) await client.send(ack) elif msg.startsWith("UNSUBSCRIBE "): let table = msg[12..^1].strip() wsClient.unsubscribe(table) let ack = encodeFrame(0x1, "OK unsubscribed from " & table) await client.send(ack) else: let echo = encodeFrame(0x1, "ECHO: " & msg) await client.send(echo) else: discard buf = buf[consumed..^1] except: discard finally: echo "WebSocket client ", id, " disconnected" server.clients.del(id) client.close() # ---------------------------------------------------------------------- # HTTP upgrade + WebSocket handoff # ---------------------------------------------------------------------- proc handleConnection(server: WsServer, client: AsyncSocket) {.async.} = let firstLine = await client.recvLine() if firstLine.len == 0: client.close() return var headers = initTable[string, string]() var wsKey = "" while true: let line = await client.recvLine() if line == "\r" or line == "": break let parts = line.split(":", maxSplit = 1) if parts.len >= 2: let key = parts[0].strip().toLower() let val = parts[1].strip() headers[key] = val if key == "sec-websocket-key": wsKey = val if wsKey.len == 0: await client.send("HTTP/1.1 400 Bad Request\r\n\r\n") client.close() return # Auth check if server.config.authEnabled: let authHeader = headers.getOrDefault("authorization", "") if authHeader.len == 0 or not authHeader.startsWith("Bearer "): await client.send("HTTP/1.1 401 Unauthorized\r\n\r\n") client.close() return let tokenStr = authHeader[7..^1] try: let token = tokenStr.toJWT() if not token.verify(server.secretKey, HS256): await client.send("HTTP/1.1 401 Unauthorized\r\n\r\n") client.close() return except: await client.send("HTTP/1.1 401 Unauthorized\r\n\r\n") client.close() return let acceptKey = computeAcceptKey(wsKey) var response = "HTTP/1.1 101 Switching Protocols\r\n" response &= "Upgrade: websocket\r\n" response &= "Connection: Upgrade\r\n" response &= "Sec-WebSocket-Accept: " & acceptKey & "\r\n" response &= "Access-Control-Allow-Origin: *\r\n" response &= "\r\n" await client.send(response) inc server.nextId asyncCheck server.handleWsClient(client, server.nextId) proc run*(server: WsServer, port: int = 9471) {.async.} = server.running = true let sock = newAsyncSocket() sock.setSockOpt(OptReuseAddr, true) sock.bindAddr(Port(port)) sock.listen() echo "BaraDB WebSocket listening on port ", port while server.running: let client = await sock.accept() asyncCheck server.handleConnection(client) proc stop*(server: WsServer) = server.running = false