feat(agent): confirm existing domain drafts

This commit is contained in:
2026-07-24 20:33:07 +08:00
parent 775a326ae2
commit 6f8c549752
7 changed files with 205 additions and 5 deletions
@@ -0,0 +1,72 @@
package org.dromara.aihr.agent;
import lombok.RequiredArgsConstructor;
import org.dromara.aihr.agent.AihrAgentDto.ActionDraft;
import org.dromara.aihr.memory.AihrMemoryDto.ConfirmMemoryRequest;
import org.dromara.aihr.memory.AihrMemoryDto.ConfirmMemoryResponse;
import org.dromara.aihr.memory.AihrMemoryDto.MemoryCandidateResponse;
import org.dromara.aihr.memory.AihrMemoryService;
import org.dromara.common.core.constant.GlobalConstants;
import org.dromara.common.core.exception.ServiceException;
import org.redisson.api.RBucket;
import org.redisson.api.RedissonClient;
import org.springframework.stereotype.Service;
import java.security.SecureRandom;
import java.time.Duration;
import java.util.Base64;
@Service
@RequiredArgsConstructor
public class AihrAgentActionService {
private static final Duration DRAFT_TTL = Duration.ofMinutes(30);
private static final String PREFIX = GlobalConstants.GLOBAL_REDIS_KEY + "aihr:agent:action:";
private static final SecureRandom RANDOM = new SecureRandom();
private final AihrMemoryService memoryService;
private final RedissonClient redissonClient;
public ActionDraft register(MemoryCandidateResponse candidate) {
if (candidate == null || candidate.id() == null || candidate.version() == null) {
throw new ServiceException("确认草稿无效", 400);
}
RBucket<String> index = redissonClient.getBucket(PREFIX + "candidate:" + candidate.id());
String draftId = index.get();
if (!valid(draftId)) {
draftId = ticket();
index.set(draftId, DRAFT_TTL);
}
redissonClient.<Long>getBucket(PREFIX + "draft:" + draftId).set(candidate.id(), DRAFT_TTL);
return new ActionDraft(draftId, candidate.targetDomain(), candidate.version(), candidate.draft());
}
public ConfirmMemoryResponse confirm(String draftId, ConfirmMemoryRequest request) {
return memoryService.confirm(resolve(draftId), request);
}
public MemoryCandidateResponse dismiss(String draftId) {
return memoryService.dismiss(resolve(draftId));
}
private Long resolve(String draftId) {
if (!valid(draftId)) {
throw new ServiceException("确认草稿不存在或已过期", 404);
}
Long candidateId = redissonClient.<Long>getBucket(PREFIX + "draft:" + draftId).get();
if (candidateId == null || candidateId <= 0) {
throw new ServiceException("确认草稿不存在或已过期", 404);
}
return candidateId;
}
private static String ticket() {
byte[] bytes = new byte[32];
RANDOM.nextBytes(bytes);
return Base64.getUrlEncoder().withoutPadding().encodeToString(bytes);
}
private static boolean valid(String value) {
return value != null && value.length() == 43 && value.matches("[A-Za-z0-9_-]+");
}
}
@@ -4,9 +4,13 @@ import cn.dev33.satoken.annotation.SaCheckLogin;
import lombok.RequiredArgsConstructor; import lombok.RequiredArgsConstructor;
import org.dromara.aihr.agent.AihrAgentDto.AgentRequest; import org.dromara.aihr.agent.AihrAgentDto.AgentRequest;
import org.dromara.aihr.agent.AihrAgentDto.AgentResponse; import org.dromara.aihr.agent.AihrAgentDto.AgentResponse;
import org.dromara.aihr.memory.AihrMemoryDto.ConfirmMemoryRequest;
import org.dromara.aihr.memory.AihrMemoryDto.ConfirmMemoryResponse;
import org.dromara.aihr.memory.AihrMemoryDto.MemoryCandidateResponse;
import org.dromara.common.core.domain.R; import org.dromara.common.core.domain.R;
import org.springframework.http.MediaType; import org.springframework.http.MediaType;
import org.springframework.web.bind.annotation.PostMapping; import org.springframework.web.bind.annotation.PostMapping;
import org.springframework.web.bind.annotation.PathVariable;
import org.springframework.web.bind.annotation.RequestBody; import org.springframework.web.bind.annotation.RequestBody;
import org.springframework.web.bind.annotation.RequestMapping; import org.springframework.web.bind.annotation.RequestMapping;
import org.springframework.web.bind.annotation.RequestParam; import org.springframework.web.bind.annotation.RequestParam;
@@ -20,6 +24,7 @@ import org.springframework.web.multipart.MultipartFile;
public class AihrAgentController { public class AihrAgentController {
private final AihrAgentOrchestrator orchestrator; private final AihrAgentOrchestrator orchestrator;
private final AihrAgentActionService actionService;
@SaCheckLogin @SaCheckLogin
@PostMapping("/messages") @PostMapping("/messages")
@@ -43,4 +48,17 @@ public class AihrAgentController {
file file
)); ));
} }
@SaCheckLogin
@PostMapping("/actions/{draftId}/confirm")
public R<ConfirmMemoryResponse> confirm(@PathVariable String draftId,
@RequestBody ConfirmMemoryRequest request) {
return R.ok(actionService.confirm(draftId, request));
}
@SaCheckLogin
@PostMapping("/actions/{draftId}/dismiss")
public R<MemoryCandidateResponse> dismiss(@PathVariable String draftId) {
return R.ok(actionService.dismiss(draftId));
}
} }
@@ -79,7 +79,7 @@ public final class AihrAgentDto {
public record SourceSummary(String type, String label, String verifiedAt) { public record SourceSummary(String type, String label, String verifiedAt) {
} }
public record ActionDraft(String draftId, String type, Long version, Object payload) { public record ActionDraft(String draftId, String type, Integer version, Object payload) {
} }
public record Clarification(String prompt, List<String> fields) { public record Clarification(String prompt, List<String> fields) {
@@ -5,6 +5,7 @@ import org.dromara.aihr.agent.AihrAgentDto.AgentPlan;
import org.dromara.aihr.agent.AihrAgentDto.AgentRequest; import org.dromara.aihr.agent.AihrAgentDto.AgentRequest;
import org.dromara.aihr.agent.AihrAgentDto.AgentResponse; import org.dromara.aihr.agent.AihrAgentDto.AgentResponse;
import org.dromara.aihr.agent.AihrAgentDto.AgentStatus; import org.dromara.aihr.agent.AihrAgentDto.AgentStatus;
import org.dromara.aihr.agent.AihrAgentDto.ActionDraft;
import org.dromara.aihr.agent.AihrAgentDto.Clarification; import org.dromara.aihr.agent.AihrAgentDto.Clarification;
import org.dromara.aihr.agent.AihrAgentDto.Intent; import org.dromara.aihr.agent.AihrAgentDto.Intent;
import org.dromara.aihr.agent.AihrAgentDto.SourceSummary; import org.dromara.aihr.agent.AihrAgentDto.SourceSummary;
@@ -30,6 +31,7 @@ public class AihrAgentOrchestrator {
private final AihrAgentPolicy policy; private final AihrAgentPolicy policy;
private final AihrKnowledgePrincipalResolver principalResolver; private final AihrKnowledgePrincipalResolver principalResolver;
private final AihrKnowledgeQueryService queryService; private final AihrKnowledgeQueryService queryService;
private final AihrAgentActionService actionService;
public AgentResponse handle(AgentRequest request) { public AgentResponse handle(AgentRequest request) {
requireQuestion(request); requireQuestion(request);
@@ -78,12 +80,22 @@ public class AihrAgentOrchestrator {
private AgentResponse fromResponse(AgentPlan plan, QueryResponse response, String sourceType) { private AgentResponse fromResponse(AgentPlan plan, QueryResponse response, String sourceType) {
AgentStatus status = status(response); AgentStatus status = status(response);
ActionDraft actionDraft = null;
if (response.memoryCandidate() != null) {
actionDraft = actionService.register(response.memoryCandidate());
status = response.memoryCandidate().missingFields().isEmpty()
? AgentStatus.NEEDS_CONFIRMATION
: AgentStatus.NEEDS_INPUT;
sourceType = "MEMORY_DRAFT";
}
String sourceLabel = response.citations().isEmpty() ? sourceType String sourceLabel = response.citations().isEmpty() ? sourceType
: response.citations().get(0).title(); : response.citations().get(0).title();
String answer = actionDraft == null ? response.answer()
: "已整理为待确认记录,请核对后保存。";
return new AgentResponse( return new AgentResponse(
runId(), response.conversationId(), response.contextVersion(), plan.intent(), status, response.answer(), runId(), response.conversationId(), response.contextVersion(), plan.intent(), status, answer,
List.of(new SourceSummary(sourceType, sourceLabel, OffsetDateTime.now().toString())), List.of(new SourceSummary(sourceType, sourceLabel, OffsetDateTime.now().toString())),
response.citations(), response.resources(), response.data(), null, null, List.of() response.citations(), response.resources(), response.data(), actionDraft, null, List.of()
); );
} }
@@ -0,0 +1,69 @@
package org.dromara.aihr.agent;
import org.dromara.aihr.memory.AihrMemoryDto.ConfirmMemoryRequest;
import org.dromara.aihr.memory.AihrMemoryDto.ConfirmMemoryResponse;
import org.dromara.aihr.memory.AihrMemoryDto.MemoryCandidateResponse;
import org.dromara.aihr.memory.AihrMemoryDto.MemoryDraft;
import org.dromara.aihr.memory.AihrMemoryService;
import org.junit.jupiter.api.Tag;
import org.junit.jupiter.api.Test;
import org.redisson.api.RBucket;
import org.redisson.api.RedissonClient;
import java.util.List;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.mockito.ArgumentMatchers.anyString;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;
@Tag("dev")
class AihrAgentActionServiceTest {
@SuppressWarnings({"rawtypes", "unchecked"})
@Test
void opaqueDraftDelegatesConfirmationToExistingMemoryStateMachine() {
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(null, 7L);
var service = new AihrAgentActionService(memory, redis);
var candidate = candidate();
var draft = service.register(candidate);
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(draft.draftId(), request);
assertEquals("ASSISTANT_CAPTURE", result.targetDomain());
assertEquals(88L, result.targetId());
verify(memory).confirm(7L, request);
}
@SuppressWarnings({"rawtypes", "unchecked"})
@Test
void dismissAlsoUsesTheExistingCandidateReference() {
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(7L);
when(memory.dismiss(7L)).thenReturn(candidate());
var service = new AihrAgentActionService(memory, redis);
var result = service.dismiss("0123456789012345678901234567890123456789012");
assertEquals(7L, result.id());
verify(memory).dismiss(7L);
}
private static MemoryCandidateResponse candidate() {
var draft = new MemoryDraft("P1", "3栋", null, "1201", "回访", "业主回访",
"3栋1201需要回访", null, "待跟进", null, null, null, null);
return new MemoryCandidateResponse(7L, 1, "DRAFT", "FOLLOW_UP", "ASSISTANT_CAPTURE",
draft, List.of(), "2026-07-24T21:00:00");
}
}
@@ -53,7 +53,8 @@ class AihrAgentMediaTest {
Set.of("employee"), Set.of("P1"), "app"); Set.of("employee"), Set.of("P1"), "app");
} }
}, },
query query,
null
); );
} }
@@ -9,6 +9,8 @@ import org.dromara.aihr.knowledge.domain.AihrKnowledgeQueryDto.QueryResponse;
import org.dromara.aihr.knowledge.service.AihrKnowledgePrincipalResolver; import org.dromara.aihr.knowledge.service.AihrKnowledgePrincipalResolver;
import org.dromara.aihr.knowledge.service.AihrKnowledgeQueryService; import org.dromara.aihr.knowledge.service.AihrKnowledgeQueryService;
import org.dromara.aihr.knowledge.service.AihrKnowledgeDataToolService.CurrentTaskSummary; import org.dromara.aihr.knowledge.service.AihrKnowledgeDataToolService.CurrentTaskSummary;
import org.dromara.aihr.memory.AihrMemoryDto.MemoryCandidateResponse;
import org.dromara.aihr.memory.AihrMemoryDto.MemoryDraft;
import org.junit.jupiter.api.Tag; import org.junit.jupiter.api.Tag;
import org.junit.jupiter.api.Test; import org.junit.jupiter.api.Test;
@@ -71,7 +73,32 @@ class AihrAgentOrchestratorTest {
assertEquals(AgentStatus.UNAVAILABLE, result.status()); assertEquals(AgentStatus.UNAVAILABLE, result.status());
} }
@Test
void captureCandidateBecomesAnOpaqueConfirmationDraft() {
var candidate = new MemoryCandidateResponse(7L, 1, "DRAFT", "FOLLOW_UP", "ASSISTANT_CAPTURE",
new MemoryDraft("P1", "3栋", null, "1201", "回访", "业主回访",
"3栋1201需要回访", null, "待跟进", null, null, null, null),
List.of(), "2026-07-24T21:00:00");
var query = new CapturingQueryService(response("", null, true).withMemoryCandidate(candidate));
var action = new AihrAgentActionService(null, null) {
@Override
public AihrAgentDto.ActionDraft register(MemoryCandidateResponse ignored) {
return new AihrAgentDto.ActionDraft("opaque-draft", "ASSISTANT_CAPTURE", 1, candidate.draft());
}
};
var result = orchestrator(query, action).handle(request("记一下,3栋1201需要回访"));
assertEquals(AgentStatus.NEEDS_CONFIRMATION, result.status());
assertEquals("opaque-draft", result.actionDraft().draftId());
}
private static AihrAgentOrchestrator orchestrator(CapturingQueryService query) { private static AihrAgentOrchestrator orchestrator(CapturingQueryService query) {
return orchestrator(query, null);
}
private static AihrAgentOrchestrator orchestrator(CapturingQueryService query,
AihrAgentActionService actionService) {
return new AihrAgentOrchestrator( return new AihrAgentOrchestrator(
new AihrAgentPlanner(null, new ObjectMapper()), new AihrAgentPlanner(null, new ObjectMapper()),
new AihrAgentPolicy(), new AihrAgentPolicy(),
@@ -82,7 +109,8 @@ class AihrAgentOrchestratorTest {
Set.of("employee"), Set.of("P1"), "app"); Set.of("employee"), Set.of("P1"), "app");
} }
}, },
query query,
actionService
); );
} }