/* * simcu::websocket.server —— WebSocket 服务端(RFC 6455),handler 接口式。 * * 用法: * let server = WebSocketServer(8080) * server.addHandler("/ws") { => MyHandler() } // 路径 -> 工厂,每次握手创建独立 handler 实例 * server.start() * ... * server.close() // 停止监听并断开所有连接 * * 多 handler 按注册路径路由:握手时按请求路径匹配 addHandler 的 path,匹配则调用工厂 * 创建一个新的 handler 实例(每连接一个),无匹配返回 404。 * 连接方法:send / ping / close / terminate(send 三重重载:WebSocketMessage / String / Array)。 * 组管理(链式):ws.group(name).join/leave/count/isEmpty/connIds,ws.conn(connId).groups(),ws.groups()。 * 发送门面(链式):group(name).send / conn(connId).send(均为三重重载)。 * 广播:broadcast(三重重载)/ connectionCount。 */ package simcu::websocket.server import std.collection.{ArrayList, HashMap} import std.net.IPSocketAddress import std.net.SocketException import std.net.TcpServerSocket import std.net.TcpSocket import std.sync.Mutex import simcu::websocket.common.FrameStream import simcu::websocket.common.WebSocketMessage import simcu::websocket.common.WebSocketException /// WebSocket 服务端(按路径工厂路由:每连接创建独立 handler 实例,支持组管理/单播/组播/广播)。 public class WebSocketServer { private let bindPort: UInt16 private let maxPayload: Int64 private var pathFactories = HashMap WebSocketHandler>() private var serverSock: ?TcpServerSocket = None private var running = false private var connections = HashMap() internal var connGroups = HashMap>() private var nextConnId: Int64 = 1 internal let connLock = Mutex() /// @param port 监听端口,0 表示随机空闲端口(start 后通过 localPort 读取)。 /// @param maxPayload 单条消息最大字节数,默认 64KB。 public init(port: Int64, maxPayload!: Int64 = 65536) { this.bindPort = UInt16(port) this.maxPayload = maxPayload } /// 注册一个 handler 工厂并绑定到端点路径(多 handler 按路径路由)。 /// 每次握手成功都会调用工厂创建一个新的 handler 实例(每连接一个), /// 因此不要在工厂里复用有状态对象(如 scoped 服务),应让各回调按需现场解析。 public func addHandler(path: String, factory: () -> WebSocketHandler): Unit { pathFactories[path] = factory } /// 启动监听:bind + accept 循环,每连接一个协程处理握手与读循环。 public func start(backlog!: Int64 = 128): Unit { if (running) { throw WebSocketException("服务端已在监听") } let s = TcpServerSocket(bindAt: bindPort) s.backlogSize = backlog s.bind() serverSock = Some(s) running = true var paths = "" for (p in pathFactories.keys()) { paths = if (paths.isEmpty()) { p } else { "${paths}, ${p}" } } println("[websocket] 已监听 端口=${localPort} 端点=[${paths}]") spawn { => acceptLoop() } } /// 实际监听端口(bindAt=0 时用于发现随机端口)。 public prop localPort: UInt16 { get() { let sa = serverSock.getOrThrow().localAddress if (let Some(ip) <- (sa as IPSocketAddress)) { ip.port } else { throw WebSocketException("无法解析本地监听端口") } } } /// 当前在线连接数。 public prop connectionCount: Int64 { get() { synchronized(connLock) { connections.size } } } /// 停止监听并强制断开所有连接(每连接会触发自身 close 回调)。 public func close(): Unit { running = false if (let Some(s) <- serverSock) { try { s.close() } catch (_) { () } serverSock = None } let snapshot = snapshotConnections() for (c in snapshot) { c.terminate() } } /// 向所有在线连接广播消息(单个连接发送失败不影响其余连接)。 public func broadcast(message: WebSocketMessage): Unit { let snapshot = snapshotConnections() for (c in snapshot) { try { c.send(message) } catch (_) { () } } } /// 广播文本消息(三重重载:WebSocketMessage / String / Array)。 public func broadcast(text: String): Unit { broadcast(WebSocketMessage.fromText(text)) } /// 广播二进制消息。 public func broadcast(data: Array): Unit { broadcast(WebSocketMessage.fromBinary(data)) } /// 所有组名(快照)。 public func groups(): ArrayList { synchronized(connLock) { let copy = ArrayList() for ((k, _) in connGroups) { copy.add(k) } copy } } // ===== 链式发送门面 ===== /// 组门面(链式): ws.group("authed").join(conn).send("hi").send(data);count/isEmpty/connIds 读组信息。 public func group(name: String): WebSocketGroupSender { WebSocketGroupSender(this, name) } /// 连接门面(链式): ws.conn(connId).send("hi").send(data);groups() 查连接所属组。 public func conn(connId: Int64): WebSocketConnSender { WebSocketConnSender(this, connId) } // ===== 内部连接管理 ===== internal func addConnection(conn: WebSocketConnection): Unit { synchronized(connLock) { connections[conn.connId] = conn } } internal func allocateConnId(): Int64 { synchronized(connLock) { let id = nextConnId nextConnId += 1 id } } internal func removeConnection(conn: WebSocketConnection): Unit { synchronized(connLock) { connections.remove(conn.connId) removeFromAllGroups(conn.connId) } } internal func getConn(connId: Int64): ?WebSocketConnection { synchronized(connLock) { connections.get(connId) } } private func removeFromAllGroups(connId: Int64): Unit { var toRemove = ArrayList() for ((g, list) in connGroups) { if (listHasConnId(list, connId)) { removeConnId(list, connId) if (list.size == 0) { toRemove.add(g) } } } for (g in toRemove) { connGroups.remove(g) } } private func snapshotConnections(): ArrayList { synchronized(connLock) { let copy = ArrayList() for (c in connections.values()) { copy.add(c) } copy } } // ===== 监听与握手 ===== /// 按请求路径匹配 handler 工厂。无匹配返回 None。 private func findFactory(path: String): ?() -> WebSocketHandler { pathFactories.get(path) } private func acceptLoop(): Unit { while (running) { try { let s = serverSock if (let Some(sock) <- s) { let client = sock.accept() spawn { => handleConnection(client) } } else { return } } catch (e: SocketException) { if (!running) { return } println("[websocket] accept 失败: ${e}") } catch (e: Exception) { if (!running) { return } println("[websocket] accept 失败: ${e}") } } } private func handleConnection(sock: TcpSocket): Unit { let fs = FrameStream(sock, maxPayload, false) try { let request = fs.readHttpHeader() let lines = request.split("\r\n") if (lines.size == 0 || !lines[0].startsWith("GET ")) { reject(sock, 400, "Bad Request") return } let parts = lines[0].split(" ") if (parts.size < 3) { reject(sock, 400, "Bad Request") return } // 拆出路径与 query(如 "/ws/server?serverId=srv-1" → path="/ws/server", query="serverId=srv-1") let target = parts[1] var requestPath = target var requestQuery = "" if (let Some(qi) <- target.indexOf("?")) { requestPath = target[0..qi] requestQuery = target[qi + 1..] } // 按 WsPath 路由到 handler 工厂,匹配则调用工厂创建新的 handler 实例(每连接一个) let factory = findFactory(requestPath) if (let None <- factory) { reject(sock, 404, "Not Found") return } var upgradeOk = false var connectionOk = false var key = "" for (i in 1..lines.size) { let line = lines[i] let ci = line.indexOf(":") if (let Some(c) <- ci) { let name = line[0..c].toAsciiLower() let value = line[c + 1..].trimAscii() if (name == "upgrade") { upgradeOk = value.toAsciiLower().contains("websocket") } if (name == "connection") { connectionOk = value.toAsciiLower().contains("upgrade") } if (name == "sec-websocket-key") { key = value } } } if (!upgradeOk || !connectionOk || key == "") { reject(sock, 400, "Bad Request") return } // 101 响应 let accept = FrameStream.computeAccept(key) let resp = StringBuilder() resp.append("HTTP/1.1 101 Switching Protocols\r\n") resp.append("Upgrade: websocket\r\n") resp.append("Connection: Upgrade\r\n") resp.append("Sec-WebSocket-Accept: ${accept}\r\n\r\n") sock.write(resp.toString().toArray()) sock.flush() let fac = factory.getOrThrow() let h = fac() let conn = WebSocketConnection(Some(this), h, fs, allocateConnId(), requestPath, requestQuery) addConnection(conn) h.onConnect(conn) conn.runReadLoop() } catch (e: Exception) { // 握手阶段失败:记录日志(连接尚未建立,无 handler 可回调) println("[websocket] 握手失败: ${e}") } } private func reject(sock: TcpSocket, status: Int64, message: String): Unit { try { let resp = "HTTP/1.1 ${status} ${message}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n" sock.write(resp.toArray()) sock.flush() } catch (_) { () } try { sock.close() } catch (_) { () } } }