重构 WebSocket 推送服务:统一使用 /ws 路径和 channel 订阅模式

- 后端重构:
  - 统一 WebSocket 路径为 /ws,通过 channel 区分不同推送服务
  - 实现 UnifiedWebSocketHandler 统一处理所有推送频道
  - 实现 WebSocketSubscriptionService 管理订阅和推送
  - 消息类型改为 int 类型(1:SUB, 2:UNSUB, 3:DATA, 4:SUB_ACK, 5:PING, 6:PONG)
  - status 字段改为 int 类型(0: success, 非0: error)
  - 移除旧的 /ws/positions 路由和 PositionWebSocketHandler
  - 修复首推数据:订阅 position 频道后立即发送全量数据

- 前端重构:
  - 实现全局 WebSocket 管理器(单例模式)
  - 应用启动时立即建立全局 WebSocket 连接
  - 实现 useWebSocketSubscription hook 用于订阅频道
  - PositionList 完全依赖 WebSocket 推送,移除 HTTP 轮询
  - 添加连接状态显示和自动重连机制
  - 实现心跳保活机制(PING/PONG)

- 配置更新:
  - 添加 WebSocket 相关配置项
  - 更新 vite.config.ts 添加 /ws 代理配置
This commit is contained in:
WrBug
2025-11-27 02:56:24 +08:00
parent f53ba1ca49
commit 47357995af
16 changed files with 2529 additions and 101 deletions
@@ -1,6 +1,7 @@
package com.wrbug.polymarketbot.config
import com.wrbug.polymarketbot.websocket.PolymarketWebSocketHandler
import com.wrbug.polymarketbot.websocket.UnifiedWebSocketHandler
import org.springframework.context.annotation.Configuration
import org.springframework.web.socket.config.annotation.EnableWebSocket
import org.springframework.web.socket.config.annotation.WebSocketConfigurer
@@ -13,12 +14,19 @@ import org.springframework.web.socket.config.annotation.WebSocketHandlerRegistry
@Configuration
@EnableWebSocket
class WebSocketConfig(
private val polymarketWebSocketHandler: PolymarketWebSocketHandler
private val polymarketWebSocketHandler: PolymarketWebSocketHandler,
private val unifiedWebSocketHandler: UnifiedWebSocketHandler
) : WebSocketConfigurer {
override fun registerWebSocketHandlers(registry: WebSocketHandlerRegistry) {
// Polymarket RTDS 转发端点(转发外部 Polymarket 实时数据流)
registry.addHandler(polymarketWebSocketHandler, "/ws/polymarket")
.setAllowedOrigins("*") // 生产环境应该配置具体的域名
// 统一 WebSocket 端点(所有推送服务统一使用此路径,通过 channel 区分)
// 支持的频道:position(仓位推送)、order(订单推送,待实现)等
registry.addHandler(unifiedWebSocketHandler, "/ws")
.setAllowedOrigins("*") // 生产环境应该配置具体的域名
}
}
@@ -0,0 +1,29 @@
package com.wrbug.polymarketbot.dto
/**
* 仓位推送消息类型
*/
enum class PositionPushMessageType {
FULL, // 全量推送
INCREMENTAL // 增量推送
}
/**
* 仓位推送消息
*/
data class PositionPushMessage(
val type: PositionPushMessageType, // 消息类型:FULL(全量)或 INCREMENTAL(增量)
val timestamp: Long, // 消息时间戳
val currentPositions: List<AccountPositionDto> = emptyList(), // 当前仓位列表(全量或增量)
val historyPositions: List<AccountPositionDto> = emptyList(), // 历史仓位列表(全量或增量)
val removedPositionKeys: List<String> = emptyList() // 已删除的仓位键(仅增量推送时使用)
)
/**
* 仓位键(用于唯一标识一个仓位)
* 格式:accountId-marketId-side
*/
fun AccountPositionDto.getPositionKey(): String {
return "${accountId}-${marketId}-${side}"
}
@@ -0,0 +1,35 @@
package com.wrbug.polymarketbot.dto
/**
* WebSocket 消息类型
*/
enum class WebSocketMessageType(val value: Int) {
SUB(1), // 订阅
UNSUB(2), // 取消订阅
DATA(3), // 数据推送
SUB_ACK(4), // 订阅确认
PING(5), // 心跳
PONG(6); // 心跳响应
companion object {
/**
* 根据 int 值获取枚举
*/
fun fromValue(value: Int): WebSocketMessageType? {
return values().find { it.value == value }
}
}
}
/**
* WebSocket 消息
*/
data class WebSocketMessage(
val type: Int, // WebSocketMessageType 的 int 值(1:SUB, 2:UNSUB, 3:DATA, 4:SUB_ACK, 5:PING, 6:PONG
val channel: String? = null,
val payload: Any? = null, // 可以是 PositionPushMessage 或其他类型
val timestamp: Long? = null,
val status: Int? = null, // 0: success, 非0: error
val message: String? = null // 错误信息
)
@@ -24,9 +24,9 @@ class AccountService(
private val blockchainService: BlockchainService,
private val apiKeyService: PolymarketApiKeyService
) {
private val logger = LoggerFactory.getLogger(AccountService::class.java)
/**
* 通过私钥导入账户
*/
@@ -37,19 +37,19 @@ class AccountService(
if (!isValidWalletAddress(request.walletAddress)) {
return Result.failure(IllegalArgumentException("无效的钱包地址格式"))
}
// 2. 检查地址是否已存在
if (accountRepository.existsByWalletAddress(request.walletAddress)) {
return Result.failure(IllegalArgumentException("该钱包地址已存在"))
}
// 3. 验证私钥和地址的对应关系
// 注意:前端已经验证了私钥和地址的对应关系,这里只做格式验证
// 如果需要更严格的验证,可以使用以太坊库(如 web3j)进行验证
if (!isValidPrivateKey(request.privateKey)) {
return Result.failure(IllegalArgumentException("无效的私钥格式"))
}
// 4. 自动获取或创建 API Key(必须成功,否则导入失败)
logger.info("开始自动获取或创建 API Key: ${request.walletAddress}")
val apiKeyCreds = runBlocking {
@@ -58,7 +58,7 @@ class AccountService(
walletAddress = request.walletAddress,
chainId = 137L // Polygon 主网
)
if (result.isSuccess) {
val creds = result.getOrNull()
if (creds != null) {
@@ -74,7 +74,7 @@ class AccountService(
throw IllegalStateException("自动获取 API Key 失败: ${error?.message}。请确保私钥有效且账户已激活")
}
}
// 5. 如果设置为默认账户,取消其他账户的默认状态
if (request.isDefault) {
accountRepository.findByIsDefaultTrue()?.let { defaultAccount ->
@@ -82,7 +82,7 @@ class AccountService(
accountRepository.save(updated)
}
}
// 6. 获取代理地址(必须成功,否则导入失败)
val proxyAddress = runBlocking {
val proxyResult = blockchainService.getProxyAddress(request.walletAddress)
@@ -101,7 +101,7 @@ class AccountService(
throw IllegalStateException("获取代理地址失败: ${error?.message}。请确保已配置 Ethereum RPC URL 且 RPC 节点可用")
}
}
// 7. 创建账户
val account = Account(
privateKey = request.privateKey,
@@ -115,17 +115,17 @@ class AccountService(
createdAt = System.currentTimeMillis(),
updatedAt = System.currentTimeMillis()
)
val saved = accountRepository.save(account)
logger.info("成功导入账户: ${saved.id}, ${saved.walletAddress}, 代理地址: ${saved.proxyAddress}")
Result.success(toDto(saved))
} catch (e: Exception) {
logger.error("导入账户失败", e)
Result.failure(e)
}
}
/**
* 更新账户信息
*/
@@ -134,10 +134,10 @@ class AccountService(
return try {
val account = accountRepository.findById(request.accountId)
.orElse(null) ?: return Result.failure(IllegalArgumentException("账户不存在"))
// 更新账户名称
val updatedAccountName = request.accountName ?: account.accountName
// 如果设置为默认账户,取消其他账户的默认状态
val updatedIsDefault = request.isDefault ?: account.isDefault
if (updatedIsDefault && !account.isDefault) {
@@ -146,23 +146,23 @@ class AccountService(
accountRepository.save(updated)
}
}
val updated = account.copy(
accountName = updatedAccountName,
isDefault = updatedIsDefault,
updatedAt = System.currentTimeMillis()
)
val saved = accountRepository.save(updated)
logger.info("成功更新账户: ${saved.id}")
Result.success(toDto(saved))
} catch (e: Exception) {
logger.error("更新账户失败", e)
Result.failure(e)
}
}
/**
* 删除账户
*/
@@ -171,15 +171,15 @@ class AccountService(
return try {
val account = accountRepository.findById(accountId)
.orElse(null) ?: return Result.failure(IllegalArgumentException("账户不存在"))
// 注意:不再检查活跃订单,允许用户删除有活跃订单的账户
// 前端会显示确认提示框,由用户决定是否删除
// 如果删除的是默认账户,需要先设置其他账户为默认
if (account.isDefault) {
val otherAccounts = accountRepository.findAllByOrderByCreatedAtAsc()
.filter { it.id != accountId }
if (otherAccounts.isNotEmpty()) {
val newDefault = otherAccounts.first().copy(
isDefault = true,
@@ -190,17 +190,17 @@ class AccountService(
return Result.failure(IllegalStateException("不能删除最后一个账户"))
}
}
accountRepository.delete(account)
logger.info("成功删除账户: $accountId")
Result.success(Unit)
} catch (e: Exception) {
logger.error("删除账户失败", e)
Result.failure(e)
}
}
/**
* 查询账户列表
*/
@@ -208,17 +208,19 @@ class AccountService(
return try {
val accounts = accountRepository.findAllByOrderByCreatedAtAsc()
val accountDtos = accounts.map { toDto(it) }
Result.success(AccountListResponse(
list = accountDtos,
total = accountDtos.size.toLong()
))
Result.success(
AccountListResponse(
list = accountDtos,
total = accountDtos.size.toLong()
)
)
} catch (e: Exception) {
logger.error("查询账户列表失败", e)
Result.failure(e)
}
}
/**
* 查询账户详情
*/
@@ -229,16 +231,16 @@ class AccountService(
} else {
accountRepository.findByIsDefaultTrue()
}
account ?: return Result.failure(IllegalArgumentException("账户不存在"))
Result.success(toDto(account))
} catch (e: Exception) {
logger.error("查询账户详情失败", e)
Result.failure(e)
}
}
/**
* 查询账户余额
* 通过链上 RPC 查询 USDC 余额,并通过 Subgraph API 查询持仓信息
@@ -250,15 +252,15 @@ class AccountService(
} else {
accountRepository.findByIsDefaultTrue()
}
account ?: return Result.failure(IllegalArgumentException("账户不存在"))
// 检查代理地址是否存在
if (account.proxyAddress.isBlank()) {
logger.error("账户 ${account.id} 的代理地址为空,无法查询余额")
return Result.failure(IllegalStateException("账户代理地址不存在,无法查询余额。请重新导入账户以获取代理地址"))
}
// 查询 USDC 余额和持仓信息
val balanceResult = runBlocking {
try {
@@ -280,7 +282,7 @@ class AccountService(
logger.warn("持仓信息查询失败: ${positionsResult.exceptionOrNull()?.message}")
emptyList()
}
// 使用 /value 接口获取仓位总价值(而不是累加)
val positionBalanceResult = blockchainService.getTotalValue(account.proxyAddress)
val positionBalance = if (positionBalanceResult.isSuccess) {
@@ -289,7 +291,7 @@ class AccountService(
logger.warn("仓位总价值查询失败: ${positionBalanceResult.exceptionOrNull()?.message}")
"0"
}
// 查询可用余额(通过 RPC 查询 USDC 余额)
// 必须使用代理地址查询
val availableBalanceResult = blockchainService.getUsdcBalance(
@@ -304,10 +306,10 @@ class AccountService(
logger.error("USDC 可用余额 RPC 查询失败: ${error?.message}")
throw Exception("USDC 可用余额查询失败: ${error?.message}。请确保已配置 Ethereum RPC URL")
}
// 计算总余额 = 可用余额 + 仓位余额
val totalBalance = availableBalance.toSafeBigDecimal().add(positionBalance.toSafeBigDecimal())
AccountBalanceResponse(
availableBalance = availableBalance,
positionBalance = positionBalance,
@@ -319,14 +321,14 @@ class AccountService(
throw e
}
}
Result.success(balanceResult)
} catch (e: Exception) {
logger.error("查询账户余额失败", e)
Result.failure(e)
}
}
/**
* 设置默认账户
*/
@@ -335,7 +337,7 @@ class AccountService(
return try {
val account = accountRepository.findById(accountId)
.orElse(null) ?: return Result.failure(IllegalArgumentException("账户不存在"))
// 取消其他账户的默认状态
accountRepository.findByIsDefaultTrue()?.let { defaultAccount ->
if (defaultAccount.id != account.id) {
@@ -343,11 +345,11 @@ class AccountService(
accountRepository.save(updated)
}
}
// 设置当前账户为默认
val updated = account.copy(isDefault = true, updatedAt = System.currentTimeMillis())
accountRepository.save(updated)
logger.info("成功设置默认账户: $accountId")
Result.success(Unit)
} catch (e: Exception) {
@@ -355,7 +357,7 @@ class AccountService(
Result.failure(e)
}
}
/**
* 转换为 DTO
* 包含交易统计数据(总订单数、总盈亏、活跃订单数、已完成订单数、持仓数量)
@@ -379,7 +381,7 @@ class AccountService(
)
}
}
/**
* 获取账户交易统计数据
*/
@@ -388,27 +390,27 @@ class AccountService(
// 如果账户没有配置 API 凭证,无法查询统计数据
if (account.apiKey == null || account.apiSecret == null || account.apiPassphrase == null) {
return AccountStatistics(
totalOrders = null,
totalOrders = null,
totalPnl = null,
activeOrders = null,
completedOrders = null,
positionCount = null
)
}
// 使用 API 凭证(直接使用,无需解密)
val apiKey = account.apiKey
val apiSecret = account.apiSecret
val apiPassphrase = account.apiPassphrase
// 创建带认证的 API 客户端(需要钱包地址用于 POLY_ADDRESS 请求头)
val clobApi = retrofitFactory.createClobApi(apiKey, apiSecret, apiPassphrase, account.walletAddress)
// 1. 查询活跃订单数量(open/active 状态)
val activeOrdersResult = try {
var totalActiveOrders = 0L
var nextCursor: String? = null
// 分页查询所有活跃订单
do {
val response = clobApi.getActiveOrders(
@@ -425,13 +427,13 @@ class AccountService(
break
}
} while (nextCursor != null && nextCursor.isNotEmpty())
Result.success(totalActiveOrders)
} catch (e: Exception) {
logger.warn("查询活跃订单失败: ${e.message}", e)
Result.failure(e)
}
// 2. 查询已完成订单数
// 注意:交易记录数不等于已完成订单数,因为一个订单可能产生多笔交易
// 已完成订单应该是指已完全成交或已关闭的订单
@@ -442,7 +444,7 @@ class AccountService(
// 使用代理地址查询交易记录(作为 maker 的交易)
var allTrades = mutableListOf<TradeResponse>()
var nextCursor: String? = null
// 分页查询所有交易(作为 maker)
do {
val response = clobApi.getTrades(
@@ -457,22 +459,22 @@ class AccountService(
break
}
} while (nextCursor != null && nextCursor.isNotEmpty())
// 注意:Polymarket API 的 getTrades 接口只支持查询 maker_address
// 如果需要查询作为 taker 的交易,可能需要使用其他接口或查询方式
// 目前只统计作为 maker 的交易记录
// 由于 TradeResponse 没有 orderId 字段,我们无法直接去重订单
// 这里使用交易记录数作为已完成订单数的近似值
// 更准确的方式需要查询所有订单并统计状态为 "filled" 的订单
val completedOrdersCount = allTrades.size.toLong()
Result.success(completedOrdersCount)
} catch (e: Exception) {
logger.warn("查询交易记录失败: ${e.message}", e)
Result.failure(e)
}
// 3. 查询仓位信息计算总盈亏(已实现盈亏)和持仓数量
val positionsResult = try {
val positions = blockchainService.getPositions(account.proxyAddress)
@@ -496,13 +498,13 @@ class AccountService(
logger.warn("查询仓位信息失败: ${e.message}", e)
Result.failure(e)
}
val activeOrders = activeOrdersResult.getOrNull() ?: 0L
val completedOrders = completedOrdersResult.getOrNull() ?: 0L
// 总订单数 = 活跃订单数 + 已完成订单数
val totalOrders = activeOrders + completedOrders
val (totalPnl, positionCount) = positionsResult.getOrNull() ?: Pair(null, null)
AccountStatistics(
totalOrders = totalOrders,
totalPnl = totalPnl,
@@ -513,7 +515,7 @@ class AccountService(
} catch (e: Exception) {
logger.warn("获取账户统计数据失败: ${e.message}", e)
AccountStatistics(
totalOrders = null,
totalOrders = null,
totalPnl = null,
activeOrders = null,
completedOrders = null,
@@ -521,7 +523,7 @@ class AccountService(
)
}
}
/**
* 账户统计数据
*/
@@ -532,7 +534,7 @@ class AccountService(
val completedOrders: Long?,
val positionCount: Long?
)
/**
* 验证钱包地址格式
*/
@@ -540,7 +542,7 @@ class AccountService(
// 以太坊地址格式:0x 开头,42 位字符
return address.startsWith("0x") && address.length == 42 && address.matches(Regex("^0x[0-9a-fA-F]{40}$"))
}
/**
* 验证私钥格式
*/
@@ -549,7 +551,7 @@ class AccountService(
val cleanKey = if (privateKey.startsWith("0x")) privateKey.substring(2) else privateKey
return cleanKey.length == 64 && cleanKey.matches(Regex("^[0-9a-fA-F]{64}$"))
}
/**
* 查询所有账户的仓位列表
* 返回所有账户的仓位信息,包括账户信息
@@ -559,24 +561,24 @@ class AccountService(
val accounts = accountRepository.findAll()
val currentPositions = mutableListOf<AccountPositionDto>()
val historyPositions = mutableListOf<AccountPositionDto>()
// 遍历所有账户,查询每个账户的仓位
accounts.forEach { account ->
if (account.proxyAddress.isNotBlank()) {
try {
// 查询所有仓位(不限制 sortBy,获取当前和历史仓位)
val positionsResult = blockchainService.getPositions(account.proxyAddress, sortBy = null)
val positionsResult = blockchainService.getPositions(account.proxyAddress)
if (positionsResult.isSuccess) {
val positions = positionsResult.getOrNull() ?: emptyList()
// 遍历所有仓位,区分当前仓位和历史仓位
positions.forEach { pos ->
val currentValue = pos.currentValue?.toSafeBigDecimal() ?: BigDecimal.ZERO
val curPrice = pos.curPrice?.toSafeBigDecimal() ?: BigDecimal.ZERO
// 判断是否为当前仓位:currentValue != 0 且 curPrice != 0
// 使用 eq 方法判断值是否等于 0
val isCurrent = !currentValue.eq(BigDecimal.ZERO) && !curPrice.eq(BigDecimal.ZERO)
val positionDto = AccountPositionDto(
accountId = account.id!!,
accountName = account.accountName,
@@ -601,7 +603,7 @@ class AccountService(
endDate = pos.endDate,
isCurrent = isCurrent // 标识是当前仓位还是历史仓位
)
// 根据 isCurrent 分别添加到对应的列表
if (isCurrent) {
currentPositions.add(positionDto)
@@ -615,19 +617,21 @@ class AccountService(
}
}
}
// 按照接口返回的顺序返回,不进行排序
// 前端负责本地排序
Result.success(PositionListResponse(
currentPositions = currentPositions,
historyPositions = historyPositions
))
Result.success(
PositionListResponse(
currentPositions = currentPositions,
historyPositions = historyPositions
)
)
} catch (e: Exception) {
logger.error("查询所有仓位失败: ${e.message}", e)
Result.failure(e)
}
}
/**
* 检查账户是否有活跃订单
* 使用账户的 API Key 查询该账户的活跃订单
@@ -639,15 +643,15 @@ class AccountService(
logger.debug("账户 ${account.id} 未配置 API 凭证,无法查询活跃订单,允许删除")
return false
}
// 使用 API 凭证(直接使用,无需解密)
val apiKey = account.apiKey
val apiSecret = account.apiSecret
val apiPassphrase = account.apiPassphrase
// 创建带认证的 API 客户端(需要钱包地址用于 POLY_ADDRESS 请求头)
val clobApi = retrofitFactory.createClobApi(apiKey, apiSecret, apiPassphrase, account.walletAddress)
// 查询活跃订单(只查询第一条,用于判断是否有订单)
// 使用 next_cursor 参数进行分页,这里只查询第一页
val response = clobApi.getActiveOrders(
@@ -656,7 +660,7 @@ class AccountService(
asset_id = null,
next_cursor = null // null 表示从第一页开始
)
if (response.isSuccessful && response.body() != null) {
val ordersResponse = response.body()!!
val hasOrders = ordersResponse.data.isNotEmpty()
@@ -0,0 +1,337 @@
package com.wrbug.polymarketbot.service
import com.wrbug.polymarketbot.dto.AccountPositionDto
import com.wrbug.polymarketbot.dto.PositionPushMessage
import com.wrbug.polymarketbot.dto.PositionPushMessageType
import com.wrbug.polymarketbot.dto.getPositionKey
import jakarta.annotation.PostConstruct
import jakarta.annotation.PreDestroy
import kotlinx.coroutines.*
import org.slf4j.LoggerFactory
import org.springframework.beans.factory.annotation.Value
import org.springframework.stereotype.Service
import java.util.concurrent.ConcurrentHashMap
/**
* 仓位推送服务
* 轮询仓位接口,比较差异并推送增量更新
*/
@Service
class PositionPushService(
private val accountService: AccountService
) {
private val logger = LoggerFactory.getLogger(PositionPushService::class.java)
@Value("\${position.push.polling-interval:3000}")
private var pollingInterval: Long = 3000 // 轮询间隔(毫秒),默认3秒
// 存储客户端会话和对应的推送回调
private val clientCallbacks = ConcurrentHashMap<String, (PositionPushMessage) -> Unit>()
// 存储上一次的仓位数据快照(用于比较差异)
private var lastCurrentPositions: Map<String, AccountPositionDto> = emptyMap()
private var lastHistoryPositions: Map<String, AccountPositionDto> = emptyMap()
// 协程作用域和任务
private val scope = CoroutineScope(Dispatchers.Default + SupervisorJob())
private var pollingJob: Job? = null
// 同步锁,确保轮询任务的启动和停止是线程安全的
private val lock = Any()
/**
* 初始化服务(不自动启动轮询,等待有客户端连接时再启动)
*/
@PostConstruct
fun init() {
logger.info("仓位推送服务已初始化,轮询间隔: ${pollingInterval}ms,等待客户端连接...")
}
/**
* 清理资源
*/
@PreDestroy
fun destroy() {
logger.info("停止仓位推送服务")
synchronized(lock) {
pollingJob?.cancel()
pollingJob = null
}
scope.cancel()
}
/**
* 订阅仓位推送(新接口)
*/
fun subscribe(sessionId: String, callback: (PositionPushMessage) -> Unit) {
logger.info("订阅仓位推送: $sessionId")
registerSession(sessionId, callback)
}
/**
* 取消订阅仓位推送(新接口)
*/
fun unsubscribe(sessionId: String) {
logger.info("取消订阅仓位推送: $sessionId")
unregisterSession(sessionId)
}
/**
* 注册客户端会话(兼容旧接口)
* 如果有第一个客户端连接,启动轮询任务
*/
fun registerSession(sessionId: String, callback: (PositionPushMessage) -> Unit) {
logger.info("注册仓位推送客户端会话: $sessionId")
synchronized(lock) {
val wasEmpty = clientCallbacks.isEmpty()
clientCallbacks[sessionId] = callback
// 如果是第一个客户端连接,启动轮询任务
if (wasEmpty && clientCallbacks.isNotEmpty()) {
logger.info("检测到第一个客户端连接,启动轮询任务")
startPolling()
}
}
}
/**
* 注销客户端会话(兼容旧接口)
* 如果没有客户端连接了,停止轮询任务
*/
fun unregisterSession(sessionId: String) {
logger.info("注销仓位推送客户端会话: $sessionId")
synchronized(lock) {
clientCallbacks.remove(sessionId)
// 如果没有客户端连接了,停止轮询任务
if (clientCallbacks.isEmpty()) {
logger.info("没有客户端连接了,停止轮询任务")
stopPolling()
}
}
}
/**
* 发送全量数据给指定客户端
*/
suspend fun sendFullData(sessionId: String) {
try {
val result = accountService.getAllPositions()
if (result.isSuccess) {
val positions = result.getOrNull()
if (positions != null) {
val message = PositionPushMessage(
type = PositionPushMessageType.FULL,
timestamp = System.currentTimeMillis(),
currentPositions = positions.currentPositions,
historyPositions = positions.historyPositions
)
// 更新快照
lastCurrentPositions = positions.currentPositions.associateBy { it.getPositionKey() }
lastHistoryPositions = positions.historyPositions.associateBy { it.getPositionKey() }
// 发送给指定客户端
clientCallbacks[sessionId]?.invoke(message)
logger.debug("已发送全量仓位数据给客户端: $sessionId")
}
} else {
logger.warn("获取仓位数据失败,无法发送全量数据: ${result.exceptionOrNull()?.message}")
}
} catch (e: Exception) {
logger.error("发送全量仓位数据失败: $sessionId, ${e.message}", e)
}
}
/**
* 启动轮询任务
*/
private fun startPolling() {
synchronized(lock) {
// 如果已经有轮询任务在运行,先取消
pollingJob?.cancel()
// 启动新的轮询任务
pollingJob = scope.launch {
logger.info("轮询任务已启动,间隔: ${pollingInterval}ms")
while (isActive) {
try {
pollAndPush()
} catch (e: Exception) {
logger.error("轮询仓位数据失败: ${e.message}", e)
}
delay(pollingInterval)
}
}
}
}
/**
* 停止轮询任务
*/
private fun stopPolling() {
synchronized(lock) {
pollingJob?.cancel()
pollingJob = null
logger.info("轮询任务已停止")
}
}
/**
* 轮询仓位数据并推送增量更新
*/
private suspend fun pollAndPush() {
// 双重检查:如果没有客户端连接,跳过轮询(虽然理论上不应该发生,但作为安全措施)
if (clientCallbacks.isEmpty()) {
logger.debug("没有客户端连接,跳过本次轮询")
return
}
try {
val result = accountService.getAllPositions()
if (result.isSuccess) {
val positions = result.getOrNull()
if (positions != null) {
// 比较差异
val incremental = calculateIncremental(
newCurrentPositions = positions.currentPositions,
newHistoryPositions = positions.historyPositions
)
// 如果有变化,推送增量更新
if (incremental != null) {
val message = PositionPushMessage(
type = PositionPushMessageType.INCREMENTAL,
timestamp = System.currentTimeMillis(),
currentPositions = incremental.currentPositions,
historyPositions = incremental.historyPositions,
removedPositionKeys = incremental.removedKeys
)
// 推送给所有连接的客户端
clientCallbacks.values.forEach { callback ->
try {
callback(message)
} catch (e: Exception) {
logger.error("推送增量更新失败: ${e.message}", e)
}
}
logger.debug("已推送仓位增量更新,当前仓位变化: ${incremental.currentPositions.size}, 历史仓位变化: ${incremental.historyPositions.size}, 删除: ${incremental.removedKeys.size}")
}
// 更新快照
lastCurrentPositions = positions.currentPositions.associateBy { it.getPositionKey() }
lastHistoryPositions = positions.historyPositions.associateBy { it.getPositionKey() }
}
} else {
logger.warn("获取仓位数据失败: ${result.exceptionOrNull()?.message}")
}
} catch (e: Exception) {
logger.error("轮询仓位数据异常: ${e.message}", e)
}
}
/**
* 计算增量更新
* 返回 null 表示没有变化
*/
private fun calculateIncremental(
newCurrentPositions: List<AccountPositionDto>,
newHistoryPositions: List<AccountPositionDto>
): IncrementalUpdate? {
val newCurrentMap = newCurrentPositions.associateBy { it.getPositionKey() }
val newHistoryMap = newHistoryPositions.associateBy { it.getPositionKey() }
// 找出新增或更新的当前仓位
val updatedCurrentPositions = mutableListOf<AccountPositionDto>()
newCurrentMap.forEach { (key, newPos) ->
val oldPos = lastCurrentPositions[key]
if (oldPos == null || hasChanged(oldPos, newPos)) {
updatedCurrentPositions.add(newPos)
}
}
// 找出新增或更新的历史仓位
val updatedHistoryPositions = mutableListOf<AccountPositionDto>()
newHistoryMap.forEach { (key, newPos) ->
val oldPos = lastHistoryPositions[key]
if (oldPos == null || hasChanged(oldPos, newPos)) {
updatedHistoryPositions.add(newPos)
}
}
// 找出已删除的仓位(从当前仓位变为历史仓位,或完全删除)
val removedKeys = mutableListOf<String>()
// 检查上次的当前仓位是否还在当前仓位列表中
lastCurrentPositions.forEach { (key, _) ->
if (!newCurrentMap.containsKey(key)) {
// 如果不在当前仓位中,检查是否移到了历史仓位
if (!newHistoryMap.containsKey(key)) {
// 完全删除
removedKeys.add(key)
} else {
// 从当前仓位移到历史仓位,需要更新历史仓位
newHistoryMap[key]?.let { updatedHistoryPositions.add(it) }
}
}
}
// 检查上次的历史仓位是否还在历史仓位列表中
lastHistoryPositions.forEach { (key, _) ->
if (!newHistoryMap.containsKey(key)) {
// 如果不在历史仓位中,检查是否移到了当前仓位
if (!newCurrentMap.containsKey(key)) {
// 完全删除
removedKeys.add(key)
} else {
// 从历史仓位移到当前仓位,需要更新当前仓位
newCurrentMap[key]?.let { updatedCurrentPositions.add(it) }
}
}
}
// 如果没有变化,返回 null
if (updatedCurrentPositions.isEmpty() && updatedHistoryPositions.isEmpty() && removedKeys.isEmpty()) {
return null
}
return IncrementalUpdate(
currentPositions = updatedCurrentPositions,
historyPositions = updatedHistoryPositions,
removedKeys = removedKeys
)
}
/**
* 检查仓位是否有变化
* 比较关键字段:数量、价格、价值、盈亏等
*/
private fun hasChanged(old: AccountPositionDto, new: AccountPositionDto): Boolean {
return old.quantity != new.quantity ||
old.avgPrice != new.avgPrice ||
old.currentPrice != new.currentPrice ||
old.currentValue != new.currentValue ||
old.pnl != new.pnl ||
old.percentPnl != new.percentPnl ||
old.realizedPnl != new.realizedPnl ||
old.percentRealizedPnl != new.percentRealizedPnl ||
old.redeemable != new.redeemable ||
old.mergeable != new.mergeable ||
old.isCurrent != new.isCurrent
}
/**
* 增量更新数据
*/
private data class IncrementalUpdate(
val currentPositions: List<AccountPositionDto>,
val historyPositions: List<AccountPositionDto>,
val removedKeys: List<String>
)
}
@@ -0,0 +1,152 @@
package com.wrbug.polymarketbot.service
import com.wrbug.polymarketbot.dto.PositionPushMessage
import com.wrbug.polymarketbot.dto.WebSocketMessage as WsMessage
import com.wrbug.polymarketbot.dto.WebSocketMessageType
import kotlinx.coroutines.*
import org.slf4j.LoggerFactory
import org.springframework.stereotype.Service
import java.util.concurrent.ConcurrentHashMap
/**
* WebSocket 订阅管理服务
* 管理所有频道的订阅和数据推送
*/
@Service
class WebSocketSubscriptionService(
private val positionPushService: PositionPushService
) {
private val logger = LoggerFactory.getLogger(WebSocketSubscriptionService::class.java)
// 协程作用域,用于异步发送首推数据
private val scope = CoroutineScope(Dispatchers.Default + SupervisorJob())
// 存储会话和对应的推送回调
private val sessionCallbacks = ConcurrentHashMap<String, (WsMessage) -> Unit>()
// 存储每个会话的订阅频道:sessionId -> Set<channel>
private val sessionSubscriptions = ConcurrentHashMap<String, MutableSet<String>>()
// 存储每个频道的订阅会话数:channel -> Set<sessionId>
private val channelSubscriptions = ConcurrentHashMap<String, MutableSet<String>>()
/**
* 注册会话
*/
fun registerSession(sessionId: String, callback: (WsMessage) -> Unit) {
logger.info("注册 WebSocket 会话: $sessionId")
sessionCallbacks[sessionId] = callback
sessionSubscriptions[sessionId] = mutableSetOf()
}
/**
* 注销会话
*/
fun unregisterSession(sessionId: String) {
logger.info("注销 WebSocket 会话: $sessionId")
// 取消所有订阅
val channels = sessionSubscriptions.remove(sessionId) ?: emptySet()
channels.forEach { channel ->
unsubscribe(sessionId, channel)
}
sessionCallbacks.remove(sessionId)
}
/**
* 订阅频道
*/
fun subscribe(sessionId: String, channel: String, payload: Map<*, *>?) {
logger.info("订阅频道: $sessionId -> $channel")
// 检查是否已经订阅
val sessionChannels = sessionSubscriptions.getOrPut(sessionId) { mutableSetOf() }
if (sessionChannels.contains(channel)) {
logger.debug("会话 $sessionId 已经订阅了频道 $channel,跳过重复订阅")
sendSubscribeAck(sessionId, channel, true)
return
}
// 记录订阅关系
sessionChannels.add(channel)
channelSubscriptions.getOrPut(channel) { mutableSetOf() }.add(sessionId)
// 发送订阅确认
sendSubscribeAck(sessionId, channel, true)
// 根据频道类型启动推送服务
when (channel) {
"position" -> {
positionPushService.subscribe(sessionId) { message ->
pushData(sessionId, channel, message)
}
// 立即发送首推数据(全量数据)
scope.launch {
try {
positionPushService.sendFullData(sessionId)
logger.info("已发送仓位首推数据给会话: $sessionId")
} catch (e: Exception) {
logger.error("发送仓位首推数据失败: $sessionId, ${e.message}", e)
}
}
}
else -> {
logger.warn("未知的频道: $channel")
sendSubscribeAck(sessionId, channel, false, "未知的频道")
}
}
}
/**
* 取消订阅
*/
fun unsubscribe(sessionId: String, channel: String) {
logger.info("取消订阅频道: $sessionId -> $channel")
// 移除订阅关系
sessionSubscriptions[sessionId]?.remove(channel)
channelSubscriptions[channel]?.remove(sessionId)
// 取消推送服务的订阅(推送服务内部会处理是否停止轮询)
when (channel) {
"position" -> positionPushService.unsubscribe(sessionId)
}
}
/**
* 推送数据到指定会话
*/
private fun pushData(sessionId: String, channel: String, payload: Any) {
val callback = sessionCallbacks[sessionId]
if (callback != null) {
val message = WsMessage(
type = WebSocketMessageType.DATA.value,
channel = channel,
payload = payload,
timestamp = System.currentTimeMillis()
)
callback(message)
} else {
logger.warn("会话 $sessionId 的回调不存在,无法推送数据")
}
}
/**
* 发送订阅确认
*/
private fun sendSubscribeAck(sessionId: String, channel: String, success: Boolean, errorMessage: String? = null) {
val callback = sessionCallbacks[sessionId]
if (callback != null) {
val message = WsMessage(
type = WebSocketMessageType.SUB_ACK.value,
channel = channel,
status = if (success) 0 else 1, // 0: success, 非0: error
message = errorMessage
)
callback(message)
}
}
}
@@ -0,0 +1,219 @@
package com.wrbug.polymarketbot.websocket
import com.fasterxml.jackson.databind.ObjectMapper
import com.wrbug.polymarketbot.dto.WebSocketMessage as WsMessage
import com.wrbug.polymarketbot.dto.WebSocketMessageType
import com.wrbug.polymarketbot.service.WebSocketSubscriptionService
import jakarta.annotation.PostConstruct
import jakarta.annotation.PreDestroy
import kotlinx.coroutines.*
import org.slf4j.LoggerFactory
import org.springframework.beans.factory.annotation.Value
import org.springframework.stereotype.Component
import org.springframework.web.socket.*
import java.util.concurrent.ConcurrentHashMap
/**
* 统一 WebSocket 处理器
* 处理所有推送频道的订阅和数据推送
*/
@Component
class UnifiedWebSocketHandler(
private val objectMapper: ObjectMapper,
private val subscriptionService: WebSocketSubscriptionService
) : WebSocketHandler {
private val logger = LoggerFactory.getLogger(UnifiedWebSocketHandler::class.java)
@Value("\${websocket.heartbeat-timeout:60000}")
private var heartbeatTimeout: Long = 60000
// 存储客户端会话
private val clientSessions = ConcurrentHashMap<String, WebSocketSession>()
// 存储每个连接的最后活动时间
private val lastActivityTime = ConcurrentHashMap<String, Long>()
// 协程作用域
private val scope = CoroutineScope(Dispatchers.Default + SupervisorJob())
private var cleanupJob: Job? = null
@PostConstruct
fun init() {
logger.info("统一 WebSocket 处理器已初始化,心跳超时: ${heartbeatTimeout}ms")
startCleanupTask()
}
@PreDestroy
fun destroy() {
logger.info("停止统一 WebSocket 处理器")
cleanupJob?.cancel()
scope.cancel()
}
override fun afterConnectionEstablished(session: WebSocketSession) {
logger.info("WebSocket 客户端连接建立: ${session.id}")
clientSessions[session.id] = session
lastActivityTime[session.id] = System.currentTimeMillis()
// 注册会话到订阅服务
subscriptionService.registerSession(session.id) { wsMessage ->
sendMessageToClient(session.id, wsMessage)
}
}
override fun handleMessage(session: WebSocketSession, message: WebSocketMessage<*>) {
val payload = message.payload.toString()
// 处理心跳
if (payload == "PING" || payload == "ping") {
lastActivityTime[session.id] = System.currentTimeMillis()
try {
session.sendMessage(TextMessage("PONG"))
logger.debug("收到心跳并响应: ${session.id}")
} catch (e: Exception) {
logger.error("发送心跳响应失败: ${session.id}, ${e.message}", e)
}
return
}
// 更新活动时间
lastActivityTime[session.id] = System.currentTimeMillis()
// 解析消息
try {
val wsMessage: WsMessage = objectMapper.readValue(payload, WsMessage::class.java)
handleWebSocketMessage(session.id, wsMessage)
} catch (e: Exception) {
logger.error("解析 WebSocket 消息失败: ${session.id}, ${e.message}", e)
}
}
/**
* 处理 WebSocket 消息
*/
private fun handleWebSocketMessage(sessionId: String, message: WsMessage) {
val messageType = WebSocketMessageType.fromValue(message.type)
when (messageType) {
WebSocketMessageType.SUB -> {
val channel = message.channel
if (channel != null) {
val payload = message.payload as? Map<*, *>
subscriptionService.subscribe(sessionId, channel, payload)
} else {
logger.warn("订阅消息缺少 channel 字段: $sessionId")
}
}
WebSocketMessageType.UNSUB -> {
val channel = message.channel
if (channel != null) {
subscriptionService.unsubscribe(sessionId, channel)
} else {
logger.warn("取消订阅消息缺少 channel 字段: $sessionId")
}
}
null -> {
logger.warn("未知的消息类型: ${message.type}")
}
else -> {
logger.warn("不支持的消息类型: ${messageType}")
}
}
}
override fun handleTransportError(session: WebSocketSession, exception: Throwable) {
logger.error("WebSocket 传输错误: ${session.id}, ${exception.message}", exception)
cleanup(session.id)
}
override fun afterConnectionClosed(session: WebSocketSession, closeStatus: CloseStatus) {
logger.info("WebSocket 客户端连接关闭: ${session.id}, 状态: $closeStatus")
cleanup(session.id)
}
override fun supportsPartialMessages(): Boolean = false
/**
* 发送消息给客户端
*/
private fun sendMessageToClient(sessionId: String, message: WsMessage) {
val session = clientSessions[sessionId]
if (session != null && session.isOpen) {
try {
val json = objectMapper.writeValueAsString(message)
session.sendMessage(TextMessage(json))
lastActivityTime[sessionId] = System.currentTimeMillis()
} catch (e: Exception) {
logger.error("发送消息失败: $sessionId, ${e.message}", e)
cleanup(sessionId)
}
} else {
logger.warn("客户端会话不存在或已关闭: $sessionId")
cleanup(sessionId)
}
}
/**
* 清理资源
*/
private fun cleanup(sessionId: String) {
try {
val session = clientSessions.remove(sessionId)
lastActivityTime.remove(sessionId)
subscriptionService.unregisterSession(sessionId)
if (session != null && session.isOpen) {
try {
session.close(CloseStatus.NORMAL)
} catch (e: Exception) {
logger.debug("关闭会话失败: $sessionId, ${e.message}")
}
}
logger.info("已清理 WebSocket 资源: $sessionId")
} catch (e: Exception) {
logger.error("清理 WebSocket 资源时发生错误: $sessionId, ${e.message}", e)
}
}
/**
* 启动清理任务
*/
private fun startCleanupTask() {
cleanupJob = scope.launch {
while (isActive) {
try {
cleanupInactiveConnections()
} catch (e: Exception) {
logger.error("清理不活跃连接失败: ${e.message}", e)
}
delay(30000)
}
}
}
/**
* 清理不活跃的连接
*/
private fun cleanupInactiveConnections() {
val now = System.currentTimeMillis()
val inactiveSessions = mutableListOf<String>()
lastActivityTime.forEach { (sessionId, lastActivity) ->
val inactiveTime = now - lastActivity
if (inactiveTime > heartbeatTimeout) {
inactiveSessions.add(sessionId)
}
}
inactiveSessions.forEach { sessionId ->
logger.warn("检测到不活跃连接,准备清理: $sessionId, 不活跃时间: ${now - (lastActivityTime[sessionId] ?: 0)}ms")
cleanup(sessionId)
}
if (inactiveSessions.isNotEmpty()) {
logger.info("已清理 ${inactiveSessions.size} 个不活跃连接")
}
}
}
@@ -40,3 +40,12 @@ polymarket.data-api.base-url=https://data-api.polymarket.com
# 示例:https://polygon-rpc.com 或 https://polygon-mainnet.infura.io/v3/YOUR_PROJECT_ID
ethereum.rpc.url=${ETHEREUM_RPC_URL:https://polygon-rpc.com}
# 仓位推送配置
# 轮询间隔(毫秒),默认3秒
position.push.polling-interval=${POSITION_PUSH_POLLING_INTERVAL:3000}
# 心跳超时时间(毫秒),默认60秒,超过此时间未收到心跳则清理连接
position.push.heartbeat-timeout=${POSITION_PUSH_HEARTBEAT_TIMEOUT:60000}
# WebSocket 配置
websocket.heartbeat-timeout=${WEBSOCKET_HEARTBEAT_TIMEOUT:60000}