RedisHotStore.java 18 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415
  1. package com.adx.tencent.storage;
  2. import com.adx.tencent.storage.model.BidRecord;
  3. import com.adx.tencent.storage.model.QueuedEvent;
  4. import com.adx.tencent.storage.model.TrackingRecord;
  5. import com.fasterxml.jackson.databind.ObjectMapper;
  6. import org.springframework.data.redis.connection.stream.*;
  7. import org.springframework.data.redis.core.StringRedisTemplate;
  8. import java.time.Duration;
  9. import java.time.Instant;
  10. import java.nio.charset.StandardCharsets;
  11. import java.util.*;
  12. import java.util.concurrent.ConcurrentHashMap;
  13. /**
  14. * 对应 Go internal/storage/redis.go 的 RedisHotStore。
  15. * 职责:
  16. * 1. bid 记录缓存(adx:bid:{qk} / adx:media:{media}:{traceId})
  17. * 2. Redis Stream 事件队列(adx:events)
  18. * 3. BidFinder 接口
  19. */
  20. public class RedisHotStore {
  21. public static final String BID_NOT_FOUND_MSG = "bid not found";
  22. private static final Duration DEFAULT_CONVERSION_SEEN_TTL = Duration.ofDays(45);
  23. private static final Duration DEFAULT_DEDUCTION_COUNTER_TTL = Duration.ofDays(7);
  24. private static final String DEFAULT_PREFIX = "adx:";
  25. private static final String DEFAULT_STREAM = "adx:events";
  26. private final StringRedisTemplate redis;
  27. private final ObjectMapper objectMapper;
  28. private final TiDBColdStore coldStore;
  29. private final String prefix;
  30. private final String stream;
  31. private final Duration bidTtl;
  32. private final Set<String> initializedGroups = ConcurrentHashMap.newKeySet();
  33. public RedisHotStore(StringRedisTemplate redis, ObjectMapper objectMapper,
  34. TiDBColdStore coldStore, String prefix, String stream, Duration bidTtl) {
  35. this.redis = redis;
  36. this.objectMapper = objectMapper;
  37. this.coldStore = coldStore;
  38. this.prefix = (prefix == null || prefix.isBlank()) ? DEFAULT_PREFIX : prefix;
  39. this.stream = (stream == null || stream.isBlank()) ? DEFAULT_STREAM : stream;
  40. this.bidTtl = (bidTtl == null || bidTtl.isZero()) ? Duration.ofHours(24) : bidTtl;
  41. }
  42. // ─── EventRecorder ───────────────────────────────────────────────────────
  43. /**
  44. * 对应 Go RecordBid:
  45. * Pipeline: SET bid:{qk} + SET media:{media}:{traceId} + XADD
  46. */
  47. public void recordBid(BidRecord record) {
  48. if (record.getCreatedAt() == null) {
  49. record.setCreatedAt(Instant.now());
  50. }
  51. try {
  52. // Redis KEY 只存关键字段(瘦身),降低内存占用
  53. String hotPayload = objectMapper.writeValueAsString(record.toHotRecord());
  54. // Stream 存全量数据,供 ColdWorker 持久化到 TiDB
  55. String fullPayload = objectMapper.writeValueAsString(record);
  56. redis.executePipelined((org.springframework.data.redis.connection.RedisConnection conn) -> {
  57. byte[] hotBytes = hotPayload.getBytes(StandardCharsets.UTF_8);
  58. long ttlMillis = bidTtl.toMillis();
  59. if (record.getQk() != null && !record.getQk().isEmpty()) {
  60. conn.stringCommands().set(
  61. bidKey(record.getQk()).getBytes(StandardCharsets.UTF_8),
  62. hotBytes,
  63. org.springframework.data.redis.core.types.Expiration.milliseconds(ttlMillis),
  64. org.springframework.data.redis.connection.RedisStringCommands.SetOption.UPSERT
  65. );
  66. }
  67. if (record.getMedia() != null && !record.getMedia().isEmpty()
  68. && record.getMediaTraceId() != null && !record.getMediaTraceId().isEmpty()) {
  69. conn.stringCommands().set(
  70. mediaTraceKey(record.getMedia(), record.getMediaTraceId()).getBytes(StandardCharsets.UTF_8),
  71. hotBytes,
  72. org.springframework.data.redis.core.types.Expiration.milliseconds(ttlMillis),
  73. org.springframework.data.redis.connection.RedisStringCommands.SetOption.UPSERT
  74. );
  75. }
  76. // XADD 全量数据到 Stream
  77. byte[] streamKey = stream.getBytes(StandardCharsets.UTF_8);
  78. Map<byte[], byte[]> body = new LinkedHashMap<>();
  79. body.put("type".getBytes(StandardCharsets.UTF_8), "bid".getBytes(StandardCharsets.UTF_8));
  80. body.put("payload".getBytes(StandardCharsets.UTF_8), fullPayload.getBytes(StandardCharsets.UTF_8));
  81. conn.streamCommands().xAdd(MapRecord.create(streamKey, body));
  82. return null;
  83. });
  84. } catch (Exception e) {
  85. throw new RuntimeException("recordBid failed", e);
  86. }
  87. }
  88. /**
  89. * 对应 Go RecordTracking:XADD 到 Stream。
  90. */
  91. public void recordTracking(TrackingRecord record) {
  92. if (record.getCreatedAt() == null) {
  93. record.setCreatedAt(Instant.now());
  94. }
  95. try {
  96. String payload = objectMapper.writeValueAsString(record);
  97. redis.opsForStream().add(
  98. MapRecord.create(stream, Map.of("type", "tracking", "payload", payload))
  99. .withStreamKey(stream)
  100. );
  101. } catch (Exception e) {
  102. throw new RuntimeException("recordTracking failed", e);
  103. }
  104. }
  105. // ─── BidFinder ───────────────────────────────────────────────────────────
  106. /**
  107. * 点击处理完成后,将 Redis 中的 bid 记录缩减为仅保留转化匹配所需字段。
  108. * 清除 clickUrls、landingPage、appStoreLink 以释放内存。
  109. */
  110. public void shrinkAfterClick(BidRecord record) {
  111. try {
  112. String slimPayload = objectMapper.writeValueAsString(record.toPostClickRecord());
  113. byte[] slimBytes = slimPayload.getBytes(StandardCharsets.UTF_8);
  114. // 获取剩余 TTL,保持原有过期时间
  115. Long remainingTtl = redis.getExpire(bidKey(record.getQk()), java.util.concurrent.TimeUnit.MILLISECONDS);
  116. if (remainingTtl == null || remainingTtl <= 0) remainingTtl = bidTtl.toMillis();
  117. long ttlMs = remainingTtl;
  118. redis.executePipelined((org.springframework.data.redis.connection.RedisConnection conn) -> {
  119. if (record.getQk() != null && !record.getQk().isEmpty()) {
  120. conn.stringCommands().set(
  121. bidKey(record.getQk()).getBytes(StandardCharsets.UTF_8),
  122. slimBytes,
  123. org.springframework.data.redis.core.types.Expiration.milliseconds(ttlMs),
  124. org.springframework.data.redis.connection.RedisStringCommands.SetOption.UPSERT
  125. );
  126. }
  127. if (record.getMedia() != null && !record.getMedia().isEmpty()
  128. && record.getMediaTraceId() != null && !record.getMediaTraceId().isEmpty()) {
  129. conn.stringCommands().set(
  130. mediaTraceKey(record.getMedia(), record.getMediaTraceId()).getBytes(StandardCharsets.UTF_8),
  131. slimBytes,
  132. org.springframework.data.redis.core.types.Expiration.milliseconds(ttlMs),
  133. org.springframework.data.redis.connection.RedisStringCommands.SetOption.UPSERT
  134. );
  135. }
  136. return null;
  137. });
  138. } catch (Exception e) {
  139. // 缩减失败不影响主流程,仅记录日志
  140. throw new RuntimeException("shrinkAfterClick failed", e);
  141. }
  142. }
  143. public BidRecord findBidByQk(String qk) {
  144. String payload = redis.opsForValue().get(bidKey(qk));
  145. if (payload != null) {
  146. return parseBidRecord(payload);
  147. }
  148. if (coldStore != null) {
  149. return coldStore.getBidByQk(qk);
  150. }
  151. return null;
  152. }
  153. public Map<String, BidRecord> findBidsByQks(List<String> qks) {
  154. if (qks == null || qks.isEmpty()) return Collections.emptyMap();
  155. Map<String, BidRecord> result = new HashMap<>();
  156. List<String> keys = new ArrayList<>(qks.size());
  157. Map<String, String> keyToQk = new HashMap<>();
  158. for (String qk : qks) {
  159. if (qk == null || qk.isBlank() || keyToQk.containsValue(qk)) continue;
  160. String key = bidKey(qk);
  161. keys.add(key);
  162. keyToQk.put(key, qk);
  163. }
  164. List<String> payloads = redis.opsForValue().multiGet(keys);
  165. Set<String> missing = new LinkedHashSet<>();
  166. for (int i = 0; i < keys.size(); i++) {
  167. String payload = payloads != null ? payloads.get(i) : null;
  168. String qk = keyToQk.get(keys.get(i));
  169. if (payload != null) {
  170. result.put(qk, parseBidRecord(payload));
  171. } else if (qk != null) {
  172. missing.add(qk);
  173. }
  174. }
  175. if (!missing.isEmpty() && coldStore != null) {
  176. for (BidRecord bid : coldStore.getBidsByQks(new ArrayList<>(missing))) {
  177. if (bid != null && bid.getQk() != null) {
  178. result.put(bid.getQk(), bid);
  179. }
  180. }
  181. }
  182. return result;
  183. }
  184. public BidRecord findBidByMediaTrace(String media, String traceId) {
  185. String payload = redis.opsForValue().get(mediaTraceKey(media, traceId));
  186. if (payload == null) return null;
  187. return parseBidRecord(payload);
  188. }
  189. /**
  190. * 首次见到某条转化时返回 true,并写入带 TTL 的判重标记。
  191. * 使用 SETNX 保证多实例并发下只会有一个消费者把它视为“新转化”。
  192. */
  193. public boolean markConversionSeen(String dedupeKey) {
  194. if (dedupeKey == null || dedupeKey.isBlank()) {
  195. return false;
  196. }
  197. Boolean created = redis.opsForValue().setIfAbsent(
  198. conversionSeenKey(dedupeKey),
  199. "1",
  200. DEFAULT_CONVERSION_SEEN_TTL
  201. );
  202. return Boolean.TRUE.equals(created);
  203. }
  204. /**
  205. * 按账户/事件/版本/日期维度维护一个递增计数器,供窗口扣量使用。
  206. */
  207. public long incrementDeductionCounter(String accountId, int act, String trackingVersion, String date) {
  208. String key = deductionCounterKey(accountId, act, trackingVersion, date);
  209. Long count = redis.opsForValue().increment(key);
  210. if (count == null) {
  211. throw new RuntimeException("increment deduction counter failed");
  212. }
  213. if (count == 1L) {
  214. redis.expire(key, DEFAULT_DEDUCTION_COUNTER_TTL);
  215. }
  216. return count;
  217. }
  218. // ─── EventQueue ──────────────────────────────────────────────────────────
  219. /**
  220. * 对应 Go EnsureGroup:XGROUP CREATE MKSTREAM(忽略 BUSYGROUP 错误)。
  221. */
  222. public void ensureGroup(String group) {
  223. if (group == null || group.isBlank()) return;
  224. if (initializedGroups.contains(group)) return;
  225. try {
  226. redis.opsForStream().createGroup(stream, ReadOffset.from("0"), group);
  227. initializedGroups.add(group);
  228. } catch (Exception e) {
  229. if (isBusyGroupError(e)) {
  230. initializedGroups.add(group);
  231. return; // 组已存在,忽略
  232. }
  233. // Stream 可能不存在,先创建 Stream 再重试
  234. try {
  235. Map<String, String> initMap = java.util.Collections.singletonMap("_init", "1");
  236. org.springframework.data.redis.connection.stream.RecordId id =
  237. redis.opsForStream().add(org.springframework.data.redis.connection.stream.StreamRecords
  238. .newRecord().ofMap(initMap).withStreamKey(stream));
  239. if (id != null) {
  240. redis.opsForStream().delete(stream, id);
  241. }
  242. redis.opsForStream().createGroup(stream, ReadOffset.from("0"), group);
  243. initializedGroups.add(group);
  244. } catch (Exception retryEx) {
  245. if (isBusyGroupError(retryEx)) {
  246. initializedGroups.add(group);
  247. return;
  248. }
  249. throw new RuntimeException("ensureGroup failed", retryEx);
  250. }
  251. }
  252. }
  253. private boolean isBusyGroupError(Throwable e) {
  254. while (e != null) {
  255. if (e.getMessage() != null && e.getMessage().contains("BUSYGROUP")) {
  256. return true;
  257. }
  258. e = e.getCause();
  259. }
  260. return false;
  261. }
  262. /**
  263. * 对应 Go Read:XREADGROUP GROUP {group} {consumer} COUNT {count} BLOCK 1000 STREAMS {stream} >
  264. */
  265. public List<QueuedEvent> read(String group, String consumer, int count) {
  266. if (count <= 0) count = 100;
  267. try {
  268. List<MapRecord<String, Object, Object>> records = redis.opsForStream().read(
  269. Consumer.from(group, consumer),
  270. StreamReadOptions.empty().count(count).block(Duration.ofSeconds(1)),
  271. StreamOffset.create(stream, ReadOffset.lastConsumed())
  272. );
  273. return toQueuedEvents(records);
  274. } catch (Exception e) {
  275. if (e.getMessage() != null && e.getMessage().contains("NOGROUP")) {
  276. return Collections.emptyList();
  277. }
  278. throw new RuntimeException("read stream failed", e);
  279. }
  280. }
  281. /**
  282. * 对应 Go ClaimStale:XAUTOCLAIM(转交超时未 ACK 的消息)。
  283. * 注意:StreamOperations.autoClaim 在 Spring Data Redis 3.2.x 中不可用,
  284. * 使用 XCLAIM + XPENDING 替代。
  285. */
  286. public List<QueuedEvent> claimStale(String group, String consumer, Duration minIdle, int count) {
  287. if (count <= 0) count = 100;
  288. if (minIdle == null || minIdle.isZero()) minIdle = Duration.ofMinutes(2);
  289. try {
  290. // 先获取 pending 消息
  291. var pendingMessages = redis.opsForStream().pending(
  292. stream, group, org.springframework.data.domain.Range.unbounded(), count);
  293. if (pendingMessages == null || !pendingMessages.iterator().hasNext()) {
  294. return Collections.emptyList();
  295. }
  296. // 筛选超时的消息 ID
  297. Duration threshold = minIdle;
  298. List<String> staleIds = new ArrayList<>();
  299. pendingMessages.forEach(msg -> {
  300. if (msg.getElapsedTimeSinceLastDelivery().compareTo(threshold) >= 0) {
  301. staleIds.add(msg.getIdAsString());
  302. }
  303. });
  304. if (staleIds.isEmpty()) return Collections.emptyList();
  305. // XCLAIM
  306. List<MapRecord<String, Object, Object>> claimed = redis.opsForStream().claim(
  307. stream, group, consumer, threshold,
  308. staleIds.stream().map(RecordId::of).toArray(RecordId[]::new));
  309. return toQueuedEvents(claimed);
  310. } catch (Exception e) {
  311. // 退化到空列表
  312. return Collections.emptyList();
  313. }
  314. }
  315. /**
  316. * 对应 Go Ack:XACK。
  317. */
  318. public void ack(String group, List<String> ids) {
  319. if (ids == null || ids.isEmpty()) return;
  320. String[] idArr = ids.toArray(new String[0]);
  321. redis.opsForStream().acknowledge(stream, group, idArr);
  322. }
  323. /**
  324. * 对应 Go Trim:XTRIMAPPROX。
  325. */
  326. public void trim(long maxLen) {
  327. if (maxLen <= 0) return;
  328. redis.opsForStream().trim(stream, maxLen, false);
  329. }
  330. // ─── 分布式锁支持(供 RedisLock 使用)────────────────────────────────────
  331. public StringRedisTemplate getRedisTemplate() {
  332. return redis;
  333. }
  334. // ─── 私有工具方法 ─────────────────────────────────────────────────────────
  335. private String bidKey(String qk) {
  336. return prefix + "bid:" + qk;
  337. }
  338. private String mediaTraceKey(String media, String traceId) {
  339. return prefix + "media:" + media + ":" + traceId;
  340. }
  341. private String conversionSeenKey(String dedupeKey) {
  342. return prefix + "conv:seen:" + dedupeKey;
  343. }
  344. private String deductionCounterKey(String accountId, int act, String trackingVersion, String date) {
  345. return String.format("%sdeduct:counter:tencent:acct:%s:act:%d:ver:%s:date:%s",
  346. prefix,
  347. sanitizeCounterPart(accountId, "unknown"),
  348. act,
  349. sanitizeCounterPart(trackingVersion, "v1"),
  350. sanitizeCounterPart(date, "unknown"));
  351. }
  352. private static String sanitizeCounterPart(String value, String fallback) {
  353. if (value == null || value.isBlank()) {
  354. return fallback;
  355. }
  356. return value.replace(':', '_').trim();
  357. }
  358. private BidRecord parseBidRecord(String payload) {
  359. try {
  360. return objectMapper.readValue(payload, BidRecord.class);
  361. } catch (Exception e) {
  362. throw new RuntimeException("parse BidRecord failed", e);
  363. }
  364. }
  365. @SuppressWarnings("unchecked")
  366. private List<QueuedEvent> toQueuedEvents(List<MapRecord<String, Object, Object>> records) {
  367. if (records == null) return Collections.emptyList();
  368. List<QueuedEvent> events = new ArrayList<>();
  369. for (MapRecord<String, Object, Object> rec : records) {
  370. Map<Object, Object> values = rec.getValue();
  371. QueuedEvent event = new QueuedEvent();
  372. event.setId(rec.getId().getValue());
  373. event.setType(String.valueOf(values.getOrDefault("type", "")));
  374. event.setPayload(String.valueOf(values.getOrDefault("payload", "")));
  375. events.add(event);
  376. }
  377. return events;
  378. }
  379. }