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