fix(agent): tolerate Jackson int-range decode in Redis draft/session reads

全局 Redisson codec 用 DefaultTyping.NON_FINAL,final 的 Long/record 不写类名:
数字草稿 id 读回 Integer、realtime 会话 record 读回 LinkedHashMap,confirm 与
工具调用路径双双 ClassCastException(生产 19:35 日志实证)。改为 Number/String
防御转换与 ObjectMapper.convertValue 显式还原,补 3 个回归测试,模块 693 全绿。
This commit is contained in:
2026-07-25 19:48:17 +08:00
parent bb8b2dde7e
commit c2358686ce
4 changed files with 96 additions and 3 deletions
@@ -37,7 +37,7 @@ public class AihrAgentActionService {
draftId = ticket(); draftId = ticket();
index.set(draftId, DRAFT_TTL); index.set(draftId, DRAFT_TTL);
} }
redissonClient.<Long>getBucket(PREFIX + "draft:" + draftId).set(candidate.id(), DRAFT_TTL); redissonClient.getBucket(PREFIX + "draft:" + draftId).set(candidate.id(), DRAFT_TTL);
return new ActionDraft(draftId, candidate.targetDomain(), candidate.version(), candidate); return new ActionDraft(draftId, candidate.targetDomain(), candidate.version(), candidate);
} }
@@ -53,13 +53,28 @@ public class AihrAgentActionService {
if (!valid(draftId)) { if (!valid(draftId)) {
throw new ServiceException("确认草稿不存在或已过期", 404); throw new ServiceException("确认草稿不存在或已过期", 404);
} }
Long candidateId = redissonClient.<Long>getBucket(PREFIX + "draft:" + draftId).get(); // Jackson 解码 int 范围内的数字为 Integer,不能直接强转 Long
Long candidateId = toCandidateId(redissonClient.getBucket(PREFIX + "draft:" + draftId).get());
if (candidateId == null || candidateId <= 0) { if (candidateId == null || candidateId <= 0) {
throw new ServiceException("确认草稿不存在或已过期", 404); throw new ServiceException("确认草稿不存在或已过期", 404);
} }
return candidateId; return candidateId;
} }
private static Long toCandidateId(Object raw) {
if (raw instanceof Number number) {
return number.longValue();
}
if (raw instanceof String text) {
try {
return Long.parseLong(text.trim());
} catch (NumberFormatException e) {
return null;
}
}
return null;
}
private static String ticket() { private static String ticket() {
byte[] bytes = new byte[32]; byte[] bytes = new byte[32];
RANDOM.nextBytes(bytes); RANDOM.nextBytes(bytes);
@@ -1,5 +1,6 @@
package org.dromara.aihr.service; package org.dromara.aihr.service;
import com.fasterxml.jackson.databind.ObjectMapper;
import org.dromara.aihr.domain.AihrPracticeDto.RealtimeSessionState; import org.dromara.aihr.domain.AihrPracticeDto.RealtimeSessionState;
import org.dromara.common.redis.utils.RedisUtils; import org.dromara.common.redis.utils.RedisUtils;
import org.redisson.api.RateType; import org.redisson.api.RateType;
@@ -17,13 +18,24 @@ public class AihrRealtimeSessionStore {
private static final Duration SESSION_TTL = Duration.ofMinutes(30); private static final Duration SESSION_TTL = Duration.ofMinutes(30);
private static final String SESSION_KEY_PREFIX = "aihr:practice:realtime:session:"; private static final String SESSION_KEY_PREFIX = "aihr:practice:realtime:session:";
private static final ObjectMapper MAPPER = new ObjectMapper();
public void save(RealtimeSessionState state) { public void save(RealtimeSessionState state) {
RedisUtils.setCacheObject(key(state.tenantId(), state.sessionId()), state, SESSION_TTL); RedisUtils.setCacheObject(key(state.tenantId(), state.sessionId()), state, SESSION_TTL);
} }
public RealtimeSessionState find(String tenantId, String sessionId) { public RealtimeSessionState find(String tenantId, String sessionId) {
return RedisUtils.getCacheObject(key(tenantId, sessionId)); return toState(RedisUtils.getCacheObject(key(tenantId, sessionId)));
}
/**
* record 是 final,全局 codec 的 NON_FINAL 默认类型不写入类名,读回为 LinkedHashMap,需显式转换。
*/
static RealtimeSessionState toState(Object raw) {
if (raw == null) {
return null;
}
return raw instanceof RealtimeSessionState state ? state : MAPPER.convertValue(raw, RealtimeSessionState.class);
} }
/** /**
@@ -44,6 +44,25 @@ class AihrAgentActionServiceTest {
verify(memory).confirm(7L, request); verify(memory).confirm(7L, request);
} }
@SuppressWarnings({"rawtypes", "unchecked"})
@Test
void confirmToleratesJacksonDecodingSmallIdsAsInteger() {
// TypedJsonJacksonCodec(Object.class) 把 int 范围内的 id 读回为 Integer(生产 ClassCastException 回归)
RedissonClient redis = mock(RedissonClient.class);
RBucket bucket = mock(RBucket.class);
AihrMemoryService memory = mock(AihrMemoryService.class);
when(redis.getBucket(anyString())).thenReturn(bucket);
when(bucket.get()).thenReturn(Integer.valueOf(7));
var service = new AihrAgentActionService(memory, redis);
var request = new ConfirmMemoryRequest(1L, "agent-confirm-1234", candidate().draft(), false, "PRIVATE");
when(memory.confirm(7L, request)).thenReturn(new ConfirmMemoryResponse("ASSISTANT_CAPTURE", 88L, 1));
var result = service.confirm("0123456789012345678901234567890123456789012", request);
assertEquals(88L, result.targetId());
verify(memory).confirm(7L, request);
}
@SuppressWarnings({"rawtypes", "unchecked"}) @SuppressWarnings({"rawtypes", "unchecked"})
@Test @Test
void dismissAlsoUsesTheExistingCandidateReference() { void dismissAlsoUsesTheExistingCandidateReference() {
@@ -0,0 +1,47 @@
package org.dromara.aihr.service;
import org.dromara.aihr.domain.AihrPracticeDto.RealtimeSessionState;
import org.junit.jupiter.api.Tag;
import org.junit.jupiter.api.Test;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Set;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertNull;
import static org.junit.jupiter.api.Assertions.assertSame;
@Tag("dev")
class AihrRealtimeSessionStoreTest {
@Test
void toStateConvertsLinkedHashMapReadBackFromUntypedCodec() {
// 全局 codec 的 NON_FINAL 默认类型不给 final record 写类名,生产读回为 LinkedHashMap(回归)
Map<String, Object> raw = new LinkedHashMap<>();
raw.put("sessionId", "s-1");
raw.put("personaId", "digital-mentor");
raw.put("username", "13800000000");
raw.put("userId", 7); // Jackson 同时会把 int 范围内的 Long 读回 Integer
raw.put("tenantId", "000000");
raw.put("allowedTools", new ArrayList<>(List.of("search_knowledge")));
raw.put("createdAt", 1721894400);
RealtimeSessionState state = AihrRealtimeSessionStore.toState(raw);
assertEquals("s-1", state.sessionId());
assertEquals(7L, state.userId());
assertEquals(Set.of("search_knowledge"), state.allowedTools());
assertEquals(1721894400L, state.createdAt());
}
@Test
void toStatePassesThroughTypedValueAndNull() {
RealtimeSessionState state = new RealtimeSessionState("s-1", "owner-calm", "u", 7L, "000000", Set.of(), 1L);
assertSame(state, AihrRealtimeSessionStore.toState(state));
assertNull(AihrRealtimeSessionStore.toState(null));
}
}