fix(personal): persist honest model evidence

This commit is contained in:
2026-07-12 10:38:50 +08:00
parent 1429a76b69
commit 961b339efd
4 changed files with 148 additions and 30 deletions
@@ -12,6 +12,7 @@ import org.dromara.aihr.personal.domain.PersonalAssistantDto.SearchHitResponse;
import org.dromara.aihr.personal.domain.PersonalAssistantDto.SearchScope;
import org.dromara.aihr.personal.support.PersonalOwner;
import org.dromara.aihr.service.AihrModelSeedService;
import org.dromara.aihr.service.AihrModelSeedService.ChatCallResult;
import org.dromara.aihr.service.AihrSopSeedService;
import org.dromara.common.core.exception.ServiceException;
import org.springframework.beans.factory.annotation.Autowired;
@@ -37,7 +38,7 @@ public class PersonalAnswerService {
private static final String NO_EVIDENCE = "当前资料中没有足够依据";
private static final String MODEL_UNAVAILABLE = "AI 服务暂不可用,请查看引用资料";
private static final String POSITION = "生活顾问";
private static final int MAX_QUERY_LENGTH = 2000;
private static final int MAX_QUERY_LENGTH = 1000;
private static final int MAX_ITEM_IDS = 100;
private static final int PER_DOMAIN_LIMIT = 8;
private static final int TOTAL_CITATION_LIMIT = 12;
@@ -58,7 +59,7 @@ public class PersonalAnswerService {
PlatformTransactionManager transactionManager,
ObjectMapper objectMapper) {
this(personalRetrievalService::search, sopSeedService::searchAuthorized,
modelSeedService::tryChat,
modelSeedService::tryChatDetailed,
new JdbcChatPersistence(jdbcTemplate, new TransactionTemplate(transactionManager), objectMapper));
}
@@ -93,20 +94,31 @@ public class PersonalAnswerService {
List<CitationResponse> citations = retrieve(owner, validated);
String answer;
String model = null;
int inputTokens = 0;
int outputTokens = 0;
if (citations.isEmpty()) {
answer = NO_EVIDENCE;
} else {
Optional<String> generated;
Optional<ChatCallResult> generated;
try {
generated = chatRuntime.answer(systemPrompt(), userPrompt(validated, citations), 0.1D);
} catch (RuntimeException ex) {
generated = Optional.empty();
}
answer = generated.filter(value -> !value.isBlank()).map(String::trim).orElse(MODEL_UNAVAILABLE);
if (generated.isPresent() && !generated.get().content().isBlank()) {
ChatCallResult result = generated.get();
answer = result.content().trim();
model = clean(result.modelName());
model = model.isEmpty() ? null : model;
inputTokens = Math.max(0, result.inputTokens());
outputTokens = Math.max(0, result.outputTokens());
} else {
answer = MODEL_UNAVAILABLE;
}
}
long latencyMs = Math.max(0L, (System.nanoTime() - started) / 1_000_000L);
long sessionId = persistence.save(owner, validated.sessionId(), validated.query(), answer,
validated.scopes(), citations, model, PROMPT_VERSION, latencyMs);
validated.scopes(), citations, model, PROMPT_VERSION, inputTokens, outputTokens, latencyMs);
return new AskResponse(sessionId, answer, citations, model, PROMPT_VERSION);
}
@@ -262,14 +274,15 @@ public class PersonalAnswerService {
}
public interface ChatRuntime {
Optional<String> answer(String systemPrompt, String userPrompt, double temperature);
Optional<ChatCallResult> answer(String systemPrompt, String userPrompt, double temperature);
}
public interface ChatPersistence {
boolean sessionAccessible(PersonalOwner owner, long sessionId);
long save(PersonalOwner owner, Long sessionId, String query, String answer, List<SearchScope> scope,
List<CitationResponse> citations, String model, String promptVersion, long latencyMs);
List<CitationResponse> citations, String model, String promptVersion,
int inputTokens, int outputTokens, long latencyMs);
}
private record ValidatedAsk(Long sessionId, String query, List<SearchScope> scopes, LocalDate dateFrom,
@@ -299,7 +312,7 @@ public class PersonalAnswerService {
@Override
public long save(PersonalOwner owner, Long requestedSessionId, String query, String answer,
List<SearchScope> scope, List<CitationResponse> citations, String model,
String promptVersion, long latencyMs) {
String promptVersion, int inputTokens, int outputTokens, long latencyMs) {
try {
Long saved = transaction.execute(status -> {
long sessionId = requestedSessionId == null ? createSession(owner, query, scope) : requestedSessionId;
@@ -310,8 +323,9 @@ public class PersonalAnswerService {
""", owner.tenantId(), owner.userId(), sessionId);
if (touched != 1) throw new ServiceException("PERSONAL_SESSION_NOT_FOUND");
}
insertMessage(owner, sessionId, "user", query, scope, List.of(), null, null, 0L);
insertMessage(owner, sessionId, "assistant", answer, scope, citations, model, promptVersion, latencyMs);
insertMessage(owner, sessionId, "user", query, scope, List.of(), null, null, 0, 0, 0L);
insertMessage(owner, sessionId, "assistant", answer, scope, citations, model, promptVersion,
inputTokens, outputTokens, latencyMs);
return sessionId;
});
if (saved == null) throw new ServiceException("PERSONAL_CHAT_PERSIST_FAILED");
@@ -335,14 +349,14 @@ public class PersonalAnswerService {
private void insertMessage(PersonalOwner owner, long sessionId, String role, String content,
List<SearchScope> scope, List<CitationResponse> citations, String model,
String promptVersion, long latencyMs) {
String promptVersion, int inputTokens, int outputTokens, long latencyMs) {
jdbc.update("""
insert into aihr_personal_chat_message
(id, tenant_id, owner_user_id, session_id, role, content, scope_json, citations_json,
model_name, prompt_version, input_tokens, output_tokens, latency_ms, create_time)
values (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 0, 0, ?, now())
values (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, now())
""", IdWorker.getId(), owner.tenantId(), owner.userId(), sessionId, role, content,
json(scope), json(citations), model, promptVersion, latencyMs);
json(scope), json(citations), model, promptVersion, inputTokens, outputTokens, latencyMs);
}
private String json(Object value) {
@@ -219,8 +219,10 @@ public class AihrModelSeedService {
}
try {
String content = callOpenAiCompatible(runtime, modelName, prompt, request == null ? null : request.systemPrompt(), 0.2);
return new ChatResponse(true, runtime.providerCode(), modelName, content, "openai-compatible", null, List.of());
ChatCallResult result = callOpenAiCompatible(runtime, modelName, prompt,
request == null ? null : request.systemPrompt(), 0.2);
return new ChatResponse(true, runtime.providerCode(), result.modelName(), result.content(),
"openai-compatible", null, List.of());
} catch (Exception e) {
return new ChatResponse(
true,
@@ -274,6 +276,10 @@ public class AihrModelSeedService {
* 供其他模块(如三角色对练)复用的 chat 调用:模型未配置或调用失败返回 empty,由调用方决定兜底。
*/
public Optional<String> tryChat(String systemPrompt, String userPrompt, double temperature) {
return tryChatDetailed(systemPrompt, userPrompt, temperature).map(ChatCallResult::content);
}
public Optional<ChatCallResult> tryChatDetailed(String systemPrompt, String userPrompt, double temperature) {
if (!chatAllowed()) {
// ponytail: manual cost breaker; replace with metered monthly billing guard when vendor usage data is wired.
return Optional.empty();
@@ -290,6 +296,9 @@ public class AihrModelSeedService {
}
}
public record ChatCallResult(String content, String modelName, int inputTokens, int outputTokens) {
}
private boolean chatAllowed() {
return aiEnabled && chatEnabled;
}
@@ -375,7 +384,8 @@ public class AihrModelSeedService {
}
}
private String callOpenAiCompatible(RuntimeConfig runtime, String modelName, String prompt, String systemPrompt, double temperature) throws Exception {
private ChatCallResult callOpenAiCompatible(RuntimeConfig runtime, String modelName, String prompt,
String systemPrompt, double temperature) throws Exception {
ObjectNode body = objectMapper.createObjectNode();
body.put("model", modelName);
body.put("temperature", temperature);
@@ -410,16 +420,23 @@ public class AihrModelSeedService {
}
JsonNode root = objectMapper.readTree(response.body());
return parseChatCallResult(root, modelName);
}
static ChatCallResult parseChatCallResult(JsonNode root, String fallbackModelName) {
JsonNode choices = root.path("choices");
if (!choices.isArray() || choices.size() == 0) {
if (!choices.isArray() || choices.isEmpty()) {
throw new IllegalStateException("LLM response missing choices");
}
String content = choices.get(0).path("message").path("content").asText();
if (isBlank(content)) {
throw new IllegalStateException("LLM response missing message.content");
}
return content;
String responseModel = root.path("model").asText();
String actualModel = isBlank(responseModel) ? fallbackModelName : responseModel;
int inputTokens = Math.max(0, root.path("usage").path("prompt_tokens").asInt(0));
int outputTokens = Math.max(0, root.path("usage").path("completion_tokens").asInt(0));
return new ChatCallResult(content, actualModel, inputTokens, outputTokens);
}
private RuntimeConfig runtimeConfig(String requestedModel) {
@@ -8,6 +8,7 @@ import org.dromara.aihr.personal.domain.PersonalAssistantDto.SearchScope;
import org.dromara.aihr.personal.service.PersonalAnswerService;
import org.dromara.aihr.personal.support.PersonalOwner;
import org.dromara.aihr.domain.AihrSopDto.AuthorizedKnowledgeHit;
import org.dromara.aihr.service.AihrModelSeedService.ChatCallResult;
import org.dromara.common.core.exception.ServiceException;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.Tag;
@@ -45,7 +46,7 @@ class PersonalAnswerServiceTest {
PersonalAnswerService service = service(
List.of(personalHit("9", "个人记录", "先联系业主"), personalHit("9", "重复", "重复")),
List.of(new AuthorizedKnowledgeHit(7L, "企业 SOP", "再登记工单")),
Optional.of("应先联系业主,再登记工单"), persistence, new AtomicInteger()
result("应先联系业主,再登记工单"), persistence, new AtomicInteger()
);
AskResponse response = service.ask(OWNER, request(null, List.of(SearchScope.ENTERPRISE, SearchScope.PERSONAL)));
@@ -57,6 +58,10 @@ class PersonalAnswerServiceTest {
assertEquals(List.of(SearchScope.PERSONAL, SearchScope.ENTERPRISE), persistence.scope);
assertEquals(response.citations(), persistence.citations);
assertEquals(OWNER, persistence.owner);
assertEquals("test-model", response.model());
assertEquals("test-model", persistence.model);
assertEquals(11, persistence.inputTokens);
assertEquals(5, persistence.outputTokens);
}
@Test
@@ -72,7 +77,7 @@ class PersonalAnswerServiceTest {
enterpriseCalls.incrementAndGet();
return List.of(new AuthorizedKnowledgeHit(2L, "企业", "企业内容"));
},
(system, user, temperature) -> Optional.of("答案"),
(system, user, temperature) -> result("答案"),
new RecordingPersistence()
);
@@ -90,7 +95,7 @@ class PersonalAnswerServiceTest {
@Test
void refusesWithoutAuthorizedCitationsAndDoesNotCallModel() {
AtomicInteger modelCalls = new AtomicInteger();
PersonalAnswerService service = service(List.of(), List.of(), Optional.of("不应调用"),
PersonalAnswerService service = service(List.of(), List.of(), result("不应调用"),
new RecordingPersistence(), modelCalls);
AskResponse response = service.ask(OWNER, request(null, List.of(SearchScope.PERSONAL)));
@@ -109,7 +114,7 @@ class PersonalAnswerServiceTest {
(system, user, temperature) -> {
prompts.add(system);
prompts.add(user);
return Optional.of("仅引用回答");
return result("仅引用回答");
},
new RecordingPersistence()
);
@@ -137,7 +142,7 @@ class PersonalAnswerServiceTest {
(query, position, limit) -> List.of(),
(system, user, temperature) -> {
modelCalls.incrementAndGet();
return Optional.of("答案");
return result("答案");
},
persistence
);
@@ -196,8 +201,8 @@ class PersonalAnswerServiceTest {
assertTrue(persistence.sessionAccessible(OWNER, 88L));
assertEquals(88L, persistence.save(OWNER, 88L, "问题", "答案",
List.of(SearchScope.PERSONAL, SearchScope.ENTERPRISE), citations, null,
"personal_assistant_answer_v1", 9L));
List.of(SearchScope.PERSONAL, SearchScope.ENTERPRISE), citations, "provider-model",
"personal_assistant_answer_v1", 17, 8, 9L));
List<Invocation> invocations = new ArrayList<>(mockingDetails(jdbc).getInvocations());
String allSql = invocations.stream().map(invocation -> invocation.getArguments()[0].toString())
@@ -210,6 +215,15 @@ class PersonalAnswerServiceTest {
assertTrue(allArguments.contains("PERSONAL"));
assertTrue(allArguments.contains("ENTERPRISE"));
assertTrue(allArguments.contains("摘录"));
Object[] assistantArgs = invocations.stream()
.filter(invocation -> invocation.getMethod().getName().equals("update"))
.filter(invocation -> invocation.getArguments()[0].toString().contains("aihr_personal_chat_message"))
.map(PersonalAnswerServiceTest::jdbcArguments)
.filter(args -> "assistant".equals(args[4]))
.findFirst().orElseThrow();
assertEquals("provider-model", assistantArgs[8]);
assertEquals(17, assistantArgs[10]);
assertEquals(8, assistantArgs[11]);
}
@Test
@@ -219,7 +233,7 @@ class PersonalAnswerServiceTest {
for (int i = 12; i >= 1; i--) {
hits.add(personalHit(Integer.toString(i), "标题" + i, longText));
}
PersonalAnswerService service = service(hits, List.of(), Optional.of("答案"),
PersonalAnswerService service = service(hits, List.of(), result("答案"),
new RecordingPersistence(), new AtomicInteger());
AskResponse response = service.ask(OWNER, request(null, List.of(SearchScope.PERSONAL)));
@@ -232,23 +246,29 @@ class PersonalAnswerServiceTest {
@Test
void validatesRequestBeforeAnyDependencyInteraction() {
AtomicInteger calls = new AtomicInteger();
RecordingPersistence persistence = new RecordingPersistence();
PersonalAnswerService service = PersonalAnswerService.forTest(
(owner, request) -> { calls.incrementAndGet(); return List.of(); },
(query, position, limit) -> { calls.incrementAndGet(); return List.of(); },
(system, user, temperature) -> { calls.incrementAndGet(); return Optional.empty(); },
new RecordingPersistence()
persistence
);
assertThrows(ServiceException.class, () -> service.ask(OWNER,
new AskRequest(null, "问题", List.of(SearchScope.PERSONAL), null, null, List.of(-1L), "ANSWER")));
assertThrows(ServiceException.class, () -> service.ask(OWNER,
new AskRequest(null, "问题", List.of(SearchScope.PERSONAL), null, null, List.of(), "UNSUPPORTED")));
ServiceException tooLong = assertThrows(ServiceException.class, () -> service.ask(OWNER,
new AskRequest(999L, "问".repeat(1001), List.of(SearchScope.PERSONAL, SearchScope.ENTERPRISE),
null, null, List.of(), "ANSWER")));
assertEquals("PERSONAL_ASK_QUERY_INVALID", tooLong.getMessage());
assertEquals(0, calls.get());
assertEquals(0, persistence.interactions);
}
private static PersonalAnswerService service(List<SearchHitResponse> personal,
List<AuthorizedKnowledgeHit> enterprise,
Optional<String> answer,
Optional<ChatCallResult> answer,
RecordingPersistence persistence,
AtomicInteger modelCalls) {
return PersonalAnswerService.forTest(
@@ -271,14 +291,31 @@ class PersonalAnswerServiceTest {
LocalDateTime.of(2026, 7, 12, 9, 0), 1D);
}
private static Optional<ChatCallResult> result(String content) {
return Optional.of(new ChatCallResult(content, "test-model", 11, 5));
}
private static Object[] jdbcArguments(Invocation invocation) {
Object[] arguments = invocation.getArguments();
if (arguments.length == 2 && arguments[1] instanceof Object[] values) {
return values;
}
return java.util.Arrays.copyOfRange(arguments, 1, arguments.length);
}
private static final class RecordingPersistence implements PersonalAnswerService.ChatPersistence {
private boolean sessionAccessible = true;
private PersonalOwner owner;
private List<SearchScope> scope;
private List<org.dromara.aihr.personal.domain.PersonalAssistantDto.CitationResponse> citations;
private String model;
private int inputTokens;
private int outputTokens;
private int interactions;
@Override
public boolean sessionAccessible(PersonalOwner owner, long sessionId) {
interactions++;
this.owner = owner;
return sessionAccessible;
}
@@ -287,10 +324,14 @@ class PersonalAnswerServiceTest {
public long save(PersonalOwner owner, Long sessionId, String query, String answer,
List<SearchScope> scope,
List<org.dromara.aihr.personal.domain.PersonalAssistantDto.CitationResponse> citations,
String model, String promptVersion, long latencyMs) {
String model, String promptVersion, int inputTokens, int outputTokens, long latencyMs) {
interactions++;
this.owner = owner;
this.scope = scope;
this.citations = citations;
this.model = model;
this.inputTokens = inputTokens;
this.outputTokens = outputTokens;
return sessionId == null ? 500L : sessionId;
}
}
@@ -0,0 +1,46 @@
package org.dromara.aihr.service;
import com.fasterxml.jackson.databind.JsonNode;
import com.fasterxml.jackson.databind.ObjectMapper;
import org.dromara.aihr.service.AihrModelSeedService.ChatCallResult;
import org.junit.jupiter.api.Tag;
import org.junit.jupiter.api.Test;
import static org.junit.jupiter.api.Assertions.assertEquals;
@Tag("dev")
class AihrModelSeedServiceTest {
private final ObjectMapper objectMapper = new ObjectMapper();
@Test
void parsesActualModelAndUsageWithoutExposingRawResponse() throws Exception {
JsonNode response = objectMapper.readTree("""
{
"model": "provider-model-v2",
"choices": [{"message": {"content": "引用回答"}}],
"usage": {"prompt_tokens": 31, "completion_tokens": 12}
}
""");
ChatCallResult result = AihrModelSeedService.parseChatCallResult(response, "configured-model");
assertEquals("引用回答", result.content());
assertEquals("provider-model-v2", result.modelName());
assertEquals(31, result.inputTokens());
assertEquals(12, result.outputTokens());
}
@Test
void missingProviderModelAndUsageUseRuntimeModelAndHonestZeroTokens() throws Exception {
JsonNode response = objectMapper.readTree("""
{"choices": [{"message": {"content": "回答"}}]}
""");
ChatCallResult result = AihrModelSeedService.parseChatCallResult(response, "configured-model");
assertEquals("configured-model", result.modelName());
assertEquals(0, result.inputTokens());
assertEquals(0, result.outputTokens());
}
}