RedisHotStore.java 21 KB

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