fix(personal): persist honest model evidence
This commit is contained in:
+27
-13
@@ -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) {
|
||||
|
||||
+23
-6
@@ -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) {
|
||||
|
||||
+52
-11
@@ -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;
|
||||
}
|
||||
}
|
||||
|
||||
+46
@@ -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());
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user