Files
websocket-cj/src/server/websocket_server.cj
T

333 lines
11 KiB
Plaintext
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
/*
* 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 / terminatesend 三重重载:WebSocketMessage / String / Array<Byte>)。
* 组管理(链式):ws.group(name).join/leave/count/isEmpty/connIdsws.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<String, () -> WebSocketHandler>()
private var serverSock: ?TcpServerSocket = None
private var running = false
private var connections = HashMap<Int64, WebSocketConnection>()
internal var connGroups = HashMap<String, ArrayList<Int64>>()
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<Byte>)。
public func broadcast(text: String): Unit {
broadcast(WebSocketMessage.fromText(text))
}
/// 广播二进制消息。
public func broadcast(data: Array<Byte>): Unit {
broadcast(WebSocketMessage.fromBinary(data))
}
/// 所有组名(快照)。
public func groups(): ArrayList<String> {
synchronized(connLock) {
let copy = ArrayList<String>()
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<String>()
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<WebSocketConnection> {
synchronized(connLock) {
let copy = ArrayList<WebSocketConnection>()
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 (_) {
()
}
}
}