RedisHotStore.java 21 KB

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