RelationAnalysisServiceTest.java 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354
  1. package com.etotem.num.service;
  2. import com.etotem.num.common.BizException;
  3. import com.etotem.num.entity.ChatMessage;
  4. import com.etotem.num.entity.Relation;
  5. import com.etotem.num.entity.RelationAnalysisRecord;
  6. import com.etotem.num.entity.User;
  7. import com.etotem.num.repository.ChatMessageRepository;
  8. import com.etotem.num.repository.RelationAnalysisRecordRepository;
  9. import com.etotem.num.repository.RelationRepository;
  10. import com.google.gson.Gson;
  11. import org.junit.jupiter.api.BeforeEach;
  12. import org.junit.jupiter.api.Test;
  13. import org.mockito.ArgumentCaptor;
  14. import org.springframework.beans.factory.annotation.Autowired;
  15. import org.springframework.boot.test.context.SpringBootTest;
  16. import org.springframework.boot.test.mock.mockito.MockBean;
  17. import java.time.LocalDate;
  18. import java.time.LocalDateTime;
  19. import java.util.*;
  20. import static org.junit.jupiter.api.Assertions.*;
  21. import static org.mockito.Mockito.*;
  22. @SpringBootTest
  23. class RelationAnalysisServiceTest {
  24. @Autowired
  25. private RelationAnalysisService relationAnalysisService;
  26. @MockBean
  27. private RelationAnalysisRecordRepository recordRepository;
  28. @MockBean
  29. private RelationRepository relationRepository;
  30. @MockBean
  31. private ChatMessageRepository chatMessageRepository;
  32. @MockBean
  33. private UserService userService;
  34. @MockBean
  35. private DifyService difyService;
  36. @MockBean
  37. private CalculatorService calculatorService;
  38. // Gson is final and cannot be mocked by Mockito; use the real bean
  39. @Autowired
  40. private Gson gson;
  41. private final Long userId = 1L;
  42. private final Long recordId = 100L;
  43. private User mockUser;
  44. private Relation mockRelation;
  45. private RelationAnalysisRecord mockRecord;
  46. @BeforeEach
  47. void setUp() {
  48. mockUser = new User();
  49. mockUser.setId(userId);
  50. mockUser.setNickname("测试用户");
  51. mockUser.setBirthYear(1990);
  52. mockUser.setBirthMonth(6);
  53. mockUser.setBirthDay(15);
  54. mockUser.setGender(1);
  55. mockRelation = new Relation();
  56. mockRelation.setId(10L);
  57. mockRelation.setUserId(userId);
  58. mockRelation.setName("测试关系人");
  59. mockRelation.setRelationType("spouse");
  60. mockRelation.setRelationTypeLabel("配偶");
  61. mockRelation.setBirthDate(LocalDate.of(1992, 3, 20));
  62. mockRelation.setGender(0);
  63. mockRecord = new RelationAnalysisRecord();
  64. mockRecord.setId(recordId);
  65. mockRecord.setUserId(userId);
  66. mockRecord.setTitle("测试用户 & 测试关系人(配偶)");
  67. mockRecord.setMemberSnapshot("[{\"name\":\"测试用户\"}]");
  68. mockRecord.setStatus("active");
  69. mockRecord.setChatSessionId("chat_ses_abc");
  70. mockRecord.setCreateTime(LocalDateTime.now());
  71. mockRecord.setUpdateTime(LocalDateTime.now());
  72. java.util.Map<String, Object> positions = new java.util.HashMap<>();
  73. positions.put("O", 6);
  74. java.util.Map<String, Object> triangleResult = new java.util.HashMap<>();
  75. triangleResult.put("mainCharacter", 6);
  76. triangleResult.put("positions", positions);
  77. when(calculatorService.calculateFullTriangle(anyInt(), anyInt(), anyInt()))
  78. .thenReturn(triangleResult);
  79. }
  80. // ─── listRecords ────────────────────────────────────────────────
  81. @Test
  82. void testListRecords() {
  83. when(recordRepository.findByUserIdOrderByCreateTimeDesc(userId))
  84. .thenReturn(Collections.singletonList(mockRecord));
  85. List<RelationAnalysisRecord> records = relationAnalysisService.listRecords(userId);
  86. assertEquals(1, records.size());
  87. assertEquals(recordId, records.get(0).getId());
  88. verify(recordRepository).findByUserIdOrderByCreateTimeDesc(userId);
  89. }
  90. @Test
  91. void testListRecords_empty() {
  92. when(recordRepository.findByUserIdOrderByCreateTimeDesc(userId))
  93. .thenReturn(Collections.emptyList());
  94. List<RelationAnalysisRecord> records = relationAnalysisService.listRecords(userId);
  95. assertTrue(records.isEmpty());
  96. }
  97. // ─── getRecordDetail ────────────────────────────────────────────
  98. @Test
  99. void testGetRecordDetail_success() {
  100. when(recordRepository.findByIdAndUserId(recordId, userId))
  101. .thenReturn(Optional.of(mockRecord));
  102. RelationAnalysisRecord result = relationAnalysisService.getRecordDetail(userId, recordId);
  103. assertNotNull(result);
  104. assertEquals(recordId, result.getId());
  105. }
  106. @Test
  107. void testGetRecordDetail_notFound() {
  108. when(recordRepository.findByIdAndUserId(recordId, userId))
  109. .thenReturn(Optional.empty());
  110. assertThrows(BizException.class,
  111. () -> relationAnalysisService.getRecordDetail(userId, recordId));
  112. }
  113. // ─── deleteRecord ───────────────────────────────────────────────
  114. @Test
  115. void testDeleteRecord_success() {
  116. when(recordRepository.findByIdAndUserId(recordId, userId))
  117. .thenReturn(Optional.of(mockRecord));
  118. relationAnalysisService.deleteRecord(userId, recordId);
  119. verify(recordRepository).delete(mockRecord);
  120. }
  121. @Test
  122. void testDeleteRecord_notFound() {
  123. when(recordRepository.findByIdAndUserId(recordId, userId))
  124. .thenReturn(Optional.empty());
  125. assertThrows(BizException.class,
  126. () -> relationAnalysisService.deleteRecord(userId, recordId));
  127. verify(recordRepository, never()).delete(any());
  128. }
  129. // ─── getChatHistory ─────────────────────────────────────────────
  130. @Test
  131. void testGetChatHistory_success() {
  132. List<ChatMessage> messages = Arrays.asList(
  133. createChatMessage(1L, "user", "你好"),
  134. createChatMessage(2L, "ai", "你好!有什么可以帮助你的?")
  135. );
  136. when(recordRepository.findByIdAndUserId(recordId, userId))
  137. .thenReturn(Optional.of(mockRecord));
  138. when(chatMessageRepository.findByRelationRecordIdOrderByCreatedAtAsc(recordId))
  139. .thenReturn(messages);
  140. List<ChatMessage> result = relationAnalysisService.getChatHistory(userId, recordId);
  141. assertEquals(2, result.size());
  142. assertEquals("你好", result.get(0).getContent());
  143. }
  144. @Test
  145. void testGetChatHistory_recordNotFound() {
  146. when(recordRepository.findByIdAndUserId(recordId, userId))
  147. .thenReturn(Optional.empty());
  148. assertThrows(BizException.class,
  149. () -> relationAnalysisService.getChatHistory(userId, recordId));
  150. }
  151. // ─── sendChatMessage ────────────────────────────────────────────
  152. @Test
  153. void testSendChatMessage_success() {
  154. String query = "我们的关系如何?";
  155. String aiResponse = "你们的关系非常和谐。";
  156. when(recordRepository.findByIdAndUserId(recordId, userId))
  157. .thenReturn(Optional.of(mockRecord));
  158. when(difyService.invokeRelationChatflow(mockRecord.getMemberSnapshot(), query, userId.toString()))
  159. .thenReturn(aiResponse);
  160. String result = relationAnalysisService.sendChatMessage(userId, recordId, query);
  161. assertEquals(aiResponse, result);
  162. // Verify user message saved
  163. ArgumentCaptor<ChatMessage> userMsgCaptor = ArgumentCaptor.forClass(ChatMessage.class);
  164. verify(chatMessageRepository, times(2)).save(userMsgCaptor.capture());
  165. List<ChatMessage> savedMessages = userMsgCaptor.getAllValues();
  166. assertEquals("user", savedMessages.get(0).getRole());
  167. assertEquals(query, savedMessages.get(0).getContent());
  168. assertEquals("ai", savedMessages.get(1).getRole());
  169. assertEquals(aiResponse, savedMessages.get(1).getContent());
  170. // Verify record timestamp updated
  171. verify(recordRepository).save(mockRecord);
  172. }
  173. @Test
  174. void testSendChatMessage_difyFailure() {
  175. when(recordRepository.findByIdAndUserId(recordId, userId))
  176. .thenReturn(Optional.of(mockRecord));
  177. when(difyService.invokeRelationChatflow(anyString(), anyString(), anyString()))
  178. .thenThrow(new RuntimeException("Dify unavailable"));
  179. String result = relationAnalysisService.sendChatMessage(userId, recordId, "提问");
  180. assertEquals("AI 解读服务暂时不可用,请稍后再试。", result);
  181. }
  182. @Test
  183. void testSendChatMessage_recordNotFound() {
  184. when(recordRepository.findByIdAndUserId(recordId, userId))
  185. .thenReturn(Optional.empty());
  186. assertThrows(BizException.class,
  187. () -> relationAnalysisService.sendChatMessage(userId, recordId, "提问"));
  188. }
  189. // ─── startAnalysis ──────────────────────────────────────────────
  190. @Test
  191. void testStartAnalysis_success() {
  192. String question = "我们之间的能量关系如何?";
  193. List<Long> relationIds = Collections.singletonList(10L);
  194. when(userService.getById(userId)).thenReturn(mockUser);
  195. when(relationRepository.findAllById(relationIds)).thenReturn(Collections.singletonList(mockRelation));
  196. when(recordRepository.save(any(RelationAnalysisRecord.class))).thenAnswer(invocation -> {
  197. RelationAnalysisRecord saved = invocation.getArgument(0);
  198. saved.setId(recordId);
  199. return saved;
  200. });
  201. when(difyService.invokeRelationChatflow(anyString(), eq(question), eq(userId.toString())))
  202. .thenReturn("分析结果文本");
  203. RelationAnalysisRecord result = relationAnalysisService.startAnalysis(userId, relationIds, question);
  204. assertNotNull(result);
  205. assertEquals(userId, result.getUserId());
  206. assertEquals("active", result.getStatus());
  207. assertTrue(result.getTitle().contains("测试用户"));
  208. assertTrue(result.getTitle().contains("测试关系人"));
  209. assertTrue(result.getMemberSnapshot().contains("测试用户"));
  210. // Verify chat messages saved
  211. verify(chatMessageRepository, times(2)).save(any(ChatMessage.class));
  212. }
  213. @Test
  214. void testStartAnalysis_noRelations() {
  215. String question = "我的个人能量如何?";
  216. List<Long> relationIds = Collections.emptyList();
  217. when(userService.getById(userId)).thenReturn(mockUser);
  218. when(relationRepository.findAllById(relationIds)).thenReturn(Collections.emptyList());
  219. when(recordRepository.save(any(RelationAnalysisRecord.class))).thenAnswer(invocation -> {
  220. RelationAnalysisRecord saved = invocation.getArgument(0);
  221. saved.setId(recordId);
  222. return saved;
  223. });
  224. when(difyService.invokeRelationChatflow(anyString(), eq(question), eq(userId.toString())))
  225. .thenReturn("个人能量分析结果");
  226. RelationAnalysisRecord result = relationAnalysisService.startAnalysis(userId, relationIds, question);
  227. assertNotNull(result);
  228. assertEquals("个人能量分析", result.getTitle());
  229. }
  230. @Test
  231. void testStartAnalysis_difyFailure() {
  232. List<Long> relationIds = Collections.singletonList(10L);
  233. when(userService.getById(userId)).thenReturn(mockUser);
  234. when(relationRepository.findAllById(relationIds)).thenReturn(Collections.singletonList(mockRelation));
  235. when(recordRepository.save(any(RelationAnalysisRecord.class))).thenAnswer(invocation -> {
  236. RelationAnalysisRecord saved = invocation.getArgument(0);
  237. saved.setId(recordId);
  238. return saved;
  239. });
  240. when(difyService.invokeRelationChatflow(anyString(), anyString(), anyString()))
  241. .thenThrow(new RuntimeException("Dify error"));
  242. // Should still return the record (not throw)
  243. RelationAnalysisRecord result = relationAnalysisService.startAnalysis(userId, relationIds, "提问");
  244. assertNotNull(result);
  245. assertEquals(recordId, result.getId());
  246. // No chat messages should be saved when Dify fails
  247. verify(chatMessageRepository, never()).save(any(ChatMessage.class));
  248. }
  249. @Test
  250. void testStartAnalysis_userNotFound() {
  251. when(userService.getById(userId)).thenReturn(null);
  252. assertThrows(BizException.class,
  253. () -> relationAnalysisService.startAnalysis(userId, Collections.singletonList(10L), "提问"));
  254. }
  255. @Test
  256. void testStartAnalysis_skipRelationNotOwnedByUser() {
  257. Relation otherUserRelation = new Relation();
  258. otherUserRelation.setId(99L);
  259. otherUserRelation.setUserId(999L); // different user
  260. otherUserRelation.setName("别人的关系人");
  261. when(userService.getById(userId)).thenReturn(mockUser);
  262. when(relationRepository.findAllById(Collections.singletonList(99L))).thenReturn(Collections.singletonList(otherUserRelation));
  263. when(recordRepository.save(any(RelationAnalysisRecord.class))).thenAnswer(invocation -> {
  264. RelationAnalysisRecord saved = invocation.getArgument(0);
  265. saved.setId(recordId);
  266. return saved;
  267. });
  268. when(difyService.invokeRelationChatflow(anyString(), anyString(), anyString()))
  269. .thenReturn("ok");
  270. RelationAnalysisRecord result = relationAnalysisService.startAnalysis(userId, Collections.singletonList(99L), "提问");
  271. assertNotNull(result);
  272. // Title uses the raw relations list (before ownership filtering), so includes "别人的关系人"
  273. assertTrue(result.getTitle().contains("别人的关系人"));
  274. }
  275. // ─── Helpers ────────────────────────────────────────────────────
  276. private ChatMessage createChatMessage(Long id, String role, String content) {
  277. ChatMessage msg = new ChatMessage();
  278. msg.setId(id);
  279. msg.setRelationRecordId(recordId);
  280. msg.setUserId(userId);
  281. msg.setRole(role);
  282. msg.setContent(content);
  283. msg.setCreatedAt(LocalDateTime.now());
  284. return msg;
  285. }
  286. }