feat(aihr): complete M3 protected practice replay

This commit is contained in:
2026-07-23 22:21:43 +08:00
parent 83f775faec
commit 067cc30d80
11 changed files with 413 additions and 47 deletions
@@ -131,26 +131,54 @@ public class AihrSpeechController {
if (text.length() > MAX_TTS_CHARS) {
text = text.substring(0, MAX_TTS_CHARS);
}
String ttsText = text;
return speechService.synthesize(
text,
ttsText,
request == null ? null : request.voice(),
request == null ? null : request.voiceProfile())
.map(this::storeTtsAudio)
.map(audio -> storeTtsAudio(audio, request, ttsText))
.orElseGet(() -> R.fail("语音合成失败,已降级为文本展示"));
}
private R<TtsResponse> storeTtsAudio(byte[] audio) {
private R<TtsResponse> storeTtsAudio(byte[] audio, TtsRequest request, String text) {
Path tempFile = null;
try {
tempFile = Files.createTempFile("aihr-tts-", ".mp3");
Files.write(tempFile, audio);
File file = tempFile.toFile();
SysOssVo oss = ossService.upload(file);
if (hasPracticeContext(request)) {
var context = request.practiceContext();
if (!"customer".equalsIgnoreCase(context.role())) {
deleteOssQuietly(oss.getOssId());
return R.fail(HttpStatus.BAD_REQUEST, "训练语音仅支持关联业主回合");
}
String ownerIdentity = currentAppUsername();
if (ownerIdentity.isBlank()) {
deleteOssQuietly(oss.getOssId());
return R.fail(HttpStatus.FORBIDDEN, "请使用员工端账号关联训练语音");
}
try {
Long replacedOssId = practiceSeedService.attachGeneratedCustomerAudio(
context.sessionId(), context.turnIndex(), text, oss.getOssId(), ownerIdentity
);
if (replacedOssId != null && !replacedOssId.equals(oss.getOssId())) {
deleteOssQuietly(replacedOssId);
}
} catch (RuntimeException ex) {
deleteOssQuietly(oss.getOssId());
log.warn("bind generated practice tts audio failed(处理错误已隐藏)");
return R.fail(HttpStatus.BAD_REQUEST, "训练语音关联失败,请刷新训练后重试");
}
}
// Keep the old client contract playable without returning a public OSS URL.
String inlineAudioUrl = "data:audio/mp3;base64," + Base64.getEncoder().encodeToString(audio);
return R.ok(new TtsResponse(inlineAudioUrl, "openai-compatible", oss.getOssId(), inlineAudioUrl));
} catch (Exception e) {
log.warn("tts audio upload failed(处理错误已隐藏)");
if (hasPracticeContext(request)) {
return R.fail("训练录音保存失败,请稍后重试");
}
String inlineAudioUrl = "data:audio/mp3;base64," + Base64.getEncoder().encodeToString(audio);
return R.ok(new TtsResponse(inlineAudioUrl, "openai-compatible"));
} finally {
@@ -163,4 +191,8 @@ public class AihrSpeechController {
}
}
}
private static boolean hasPracticeContext(TtsRequest request) {
return request != null && request.practiceContext() != null;
}
}
@@ -11,7 +11,17 @@ public final class AihrSpeechDto {
}
}
public record TtsRequest(String text, String voice, VoiceProfile voiceProfile) {
/**
* 可选训练上下文只用于把本次服务端新生成的业主语音与当前训练回合关联。
* 客户端不能借此提交或绑定任意已有 OSS 文件。
*/
public record TtsPracticeContext(String sessionId, Integer turnIndex, String role) {
}
public record TtsRequest(String text, String voice, VoiceProfile voiceProfile, TtsPracticeContext practiceContext) {
public TtsRequest(String text, String voice, VoiceProfile voiceProfile) {
this(text, voice, voiceProfile, null);
}
}
public record VoiceProfile(String role, String voice, Double speed, String emotion, String dialect) {
@@ -85,6 +85,7 @@ import java.util.LinkedHashMap;
import java.util.LinkedHashSet;
import java.util.List;
import java.util.Map;
import java.util.Objects;
import java.util.Optional;
import java.util.Set;
import java.util.concurrent.ConcurrentHashMap;
@@ -276,6 +277,7 @@ public class AihrPracticeSeedService {
private volatile boolean annotationTableReady;
private volatile boolean audioTableReady;
private final TransactionTemplate assignmentTransaction;
private final TransactionTemplate practiceAudioTransaction;
@Autowired
public AihrPracticeSeedService(ObjectMapper objectMapper, JdbcTemplate jdbcTemplate, AihrPracticeLlmService practiceLlmService,
@@ -302,6 +304,7 @@ public class AihrPracticeSeedService {
this.jdbcTemplate = jdbcTemplate;
this.practiceLlmService = practiceLlmService;
this.assignmentTransaction = configureAssignmentTransaction(pilotExportTransaction);
this.practiceAudioTransaction = configurePracticeAudioTransaction(pilotExportTransaction);
this.sopSeedService = sopSeedService;
this.pilotExportTransaction = configurePilotExportTransaction(pilotExportTransaction);
this.ossService = ossService;
@@ -585,7 +588,7 @@ public class AihrPracticeSeedService {
tenantId(), scenario.id(), trainee, resolveExtPartyId(request, trainee), request == null ? null : request.assignmentId(),
isMobile(request), LocalDateTime.ofInstant(java.time.Instant.ofEpochMilli(promptPresentedAt), java.time.ZoneId.systemDefault()),
new AtomicInteger(0), new ArrayList<>(), new ArrayList<>(), new ArrayList<>(), customerLines,
new ArrayList<>(), new ArrayList<>(List.of(promptPresentedAt)), new ArrayList<>()
new ArrayList<>(), new ArrayList<>(), new ArrayList<>(List.of(promptPresentedAt)), new ArrayList<>()
));
return new StartResponse(
sessionId,
@@ -620,6 +623,12 @@ public class AihrPracticeSeedService {
return handleTurn(request, null, scenario, style);
}
synchronized (session) {
// A concurrent /finish persists the session and removes it from the active map
// while holding the same monitor. Never let a stale /turn return an answer that
// can no longer be included in the completed training record.
if (activeSessions.get(request.sessionId()) != session) {
throw new ServiceException("训练会话已完成,请查看训练记录");
}
return handleTurn(request, session, scenario, style);
}
}
@@ -800,14 +809,27 @@ public class AihrPracticeSeedService {
public FinishResponse finish(FinishRequest request, String ownerIdentity) {
String sessionId = request == null ? null : request.sessionId();
ActiveSession activeSession;
if (isBlank(ownerIdentity)) {
activeSession = sessionId == null ? null : activeSessions.remove(sessionId);
} else {
ActiveSession existing = sessionId == null ? null : activeSessions.get(sessionId);
requireSessionOwner(existing, ownerIdentity, tenantId());
activeSession = activeSessions.remove(sessionId);
ActiveSession activeSession = sessionId == null ? null : activeSessions.get(sessionId);
if (!isBlank(ownerIdentity)) {
requireSessionOwner(activeSession, ownerIdentity, tenantId());
}
if (activeSession == null) {
return finishActiveSession(sessionId, null);
}
synchronized (activeSession) {
if (activeSessions.get(sessionId) != activeSession) {
throw new ServiceException("训练会话已完成");
}
if (!isBlank(ownerIdentity)) {
requireSessionOwner(activeSession, ownerIdentity, tenantId());
}
FinishResponse response = finishActiveSession(sessionId, activeSession);
activeSessions.remove(sessionId, activeSession);
return response;
}
}
private FinishResponse finishActiveSession(String sessionId, ActiveSession activeSession) {
ScenarioSeed scenario = resolveScenario(activeSession == null ? null : activeSession.scenarioId(), sessionId);
String trainee = activeSession == null ? scenario.trainee() : activeSession.trainee();
PracticeResult result = evaluate(activeSession, scenario);
@@ -1677,6 +1699,15 @@ public class AihrPracticeSeedService {
return result;
}
private static TransactionTemplate configurePracticeAudioTransaction(TransactionTemplate transactionTemplate) {
if (transactionTemplate == null || transactionTemplate.getTransactionManager() == null) {
return null;
}
TransactionTemplate result = new TransactionTemplate(transactionTemplate.getTransactionManager());
result.setIsolationLevel(TransactionDefinition.ISOLATION_READ_COMMITTED);
return result;
}
private int countCompletedPilotPeople() {
ensurePracticeTable();
return count("""
@@ -1822,6 +1853,145 @@ public class AihrPracticeSeedService {
return count != null && count > 0;
}
/**
* Binds only the TTS object produced by the current request to an AI-customer turn.
* The client supplies a session/turn hint, while ownership and the spoken source text
* are always revalidated against the active or already-finished session.
*/
public Long attachGeneratedCustomerAudio(String sessionId, Integer requestedTurnIndex, String customerText,
Long ossId, String ownerIdentity) {
String id = sessionId == null ? "" : sessionId.trim();
String owner = ownerIdentity == null ? "" : ownerIdentity.trim();
if (id.isBlank() || owner.isBlank() || ossId == null || ossId <= 0) {
throw new ServiceException("训练语音关联参数无效");
}
ActiveSession activeSession = activeSessions.get(id);
if (activeSession != null) {
synchronized (activeSession) {
if (activeSessions.get(id) == activeSession) {
requireSessionOwner(activeSession, owner, tenantId());
ScenarioSeed scenario = resolveScenario(activeSession.scenarioId(), id);
int turnIndex = requireCustomerAudioTurnIndex(requestedTurnIndex, scenario);
verifyCustomerAudioText(customerLine(activeSession, turnIndex, scenario.rounds().get(turnIndex).customer()), customerText);
return rememberCustomerAudio(activeSession, turnIndex, ossId);
}
}
}
return attachFinishedCustomerAudio(id, requestedTurnIndex, customerText, ossId, owner);
}
private Long attachFinishedCustomerAudio(String sessionId, Integer requestedTurnIndex, String customerText,
Long ossId, String ownerIdentity) {
ensurePracticeTable();
ensureAudioTable();
ensureAssignmentTable();
return inPracticeAudioTransaction(() -> attachFinishedCustomerAudioLocked(
sessionId, requestedTurnIndex, customerText, ossId, ownerIdentity
));
}
/**
* Keeps the dialogue reference and its protected playback row in the same transaction.
* Locking the session row serializes late TTS retries for the same completed training.
*/
private Long attachFinishedCustomerAudioLocked(String sessionId, Integer requestedTurnIndex, String customerText,
Long ossId, String ownerIdentity) {
List<FinishedSessionAudioSource> rows = jdbcTemplate.query("""
SELECT scenario_id, dialogue_json
FROM aihr_practice_session
WHERE tenant_id = ? AND session_id = ? AND ext_party_id = ?
AND mode = 'mobile' AND finished_time IS NOT NULL
LIMIT 1
FOR UPDATE
""", (rs, rowNum) -> new FinishedSessionAudioSource(
rs.getString("scenario_id"), rs.getString("dialogue_json")
), tenantId(), sessionId, ownerIdentity);
if (rows.isEmpty()) {
throw new ServiceException("训练会话不存在或无权访问");
}
FinishedSessionAudioSource source = rows.get(0);
ScenarioSeed scenario = resolveScenario(source.scenarioId(), sessionId);
int turnIndex = requireCustomerAudioTurnIndex(requestedTurnIndex, scenario);
List<DialogueResponse> dialogue = readDialogue(source.dialogueJson(), scenario);
int dialogueIndex = customerDialogueIndex(dialogue, turnIndex);
if (dialogueIndex < 0) {
throw new ServiceException("训练语音回合无效");
}
DialogueResponse customerTurn = dialogue.get(dialogueIndex);
// Completed dialogue is stored after sensitive-text masking. Compare the same
// canonical form so an otherwise valid late TTS does not fail on phone/room masking.
verifyCustomerAudioText(customerTurn.text(), maskSensitiveText(customerText));
List<Long> storedAudioIds = jdbcTemplate.query("""
SELECT oss_id
FROM aihr_practice_audio
WHERE tenant_id = ? AND session_id = ? AND turn_index = ? AND role = 'customer'
FOR UPDATE
""", (rs, rowNum) -> rs.getLong("oss_id"), tenantId(), sessionId, turnIndex + 1);
Long previousOssId = customerTurn.audioOssId();
if (previousOssId == null && !storedAudioIds.isEmpty()) {
previousOssId = storedAudioIds.get(0);
}
List<DialogueResponse> updatedDialogue = new ArrayList<>(dialogue);
updatedDialogue.set(dialogueIndex, new DialogueResponse(
customerTurn.role(), customerTurn.label(), customerTurn.text(), "", ossId,
customerTurn.emotion(), customerTurn.trust(), customerTurn.redFlag(), customerTurn.coachHint()
));
LocalDateTime now = LocalDateTime.now();
jdbcTemplate.update("""
INSERT INTO aihr_practice_audio
(tenant_id, session_id, turn_index, role, oss_id, audio_url, create_time, update_time)
VALUES (?, ?, ?, 'customer', ?, '', ?, ?)
ON DUPLICATE KEY UPDATE oss_id = VALUES(oss_id), audio_url = VALUES(audio_url), update_time = VALUES(update_time)
""", tenantId(), sessionId, turnIndex + 1, ossId, Timestamp.valueOf(now), Timestamp.valueOf(now));
int updated = jdbcTemplate.update("""
UPDATE aihr_practice_session
SET dialogue_json = ?, update_time = ?
WHERE tenant_id = ? AND session_id = ? AND ext_party_id = ? AND mode = 'mobile'
""", writeJson(maskDialogue(updatedDialogue)), Timestamp.valueOf(now), tenantId(), sessionId, ownerIdentity);
if (updated != 1) {
throw new ServiceException("训练语音关联失败,请刷新训练后重试");
}
if (previousOssId == null || Objects.equals(previousOssId, ossId) || canReadPracticeAudioForTenant(previousOssId)) {
return null;
}
return previousOssId;
}
private static int requireCustomerAudioTurnIndex(Integer requestedTurnIndex, ScenarioSeed scenario) {
int turnIndex = requestedTurnIndex == null ? -1 : requestedTurnIndex;
if (turnIndex < 0 || scenario == null || turnIndex >= scenario.rounds().size()) {
throw new ServiceException("训练语音回合无效");
}
return turnIndex;
}
private static void verifyCustomerAudioText(String expectedText, String requestedText) {
String expected = ttsComparableText(expectedText);
String requested = ttsComparableText(requestedText);
if (expected.isBlank() || !expected.equals(requested)) {
throw new ServiceException("训练语音与当前业主话术不匹配");
}
}
private static String ttsComparableText(String text) {
String value = text == null ? "" : text.trim();
return value.length() > 300 ? value.substring(0, 300) : value;
}
private static int customerDialogueIndex(List<DialogueResponse> dialogue, int turnIndex) {
int customerIndex = 0;
for (int index = 0; index < dialogue.size(); index++) {
if (!"customer".equals(dialogue.get(index).role())) {
continue;
}
if (customerIndex == turnIndex) {
return index;
}
customerIndex += 1;
}
return -1;
}
public void stagePracticeAudio(Long ossId, String audioUrl, String ownerIdentity, Long ownerUserId) {
if (ossId == null || isBlank(ownerIdentity) || ownerUserId == null || ownerUserId <= 0) {
return;
@@ -3081,6 +3251,13 @@ public class AihrPracticeSeedService {
return assignmentTransaction.execute(status -> work.get());
}
private <T> T inPracticeAudioTransaction(Supplier<T> work) {
if (practiceAudioTransaction == null) {
return work.get();
}
return practiceAudioTransaction.execute(status -> work.get());
}
static LocalDate normalizeAssignmentDueDate(String value, LocalDate today) {
LocalDate base = today == null ? businessToday() : today;
if (value == null || value.isBlank()) {
@@ -4423,6 +4600,20 @@ public class AihrPracticeSeedService {
}
}
private Long rememberCustomerAudio(ActiveSession session, int roundIndex, Long ossId) {
if (session == null || ossId == null) {
return null;
}
List<Long> audioOssIds = session.customerAudioOssIds();
synchronized (audioOssIds) {
while (audioOssIds.size() <= roundIndex) {
audioOssIds.add(null);
}
Long previousOssId = audioOssIds.set(roundIndex, ossId);
return Objects.equals(previousOssId, ossId) ? null : previousOssId;
}
}
private void rememberResponseLatency(ActiveSession session, int roundIndex, long submittedAt) {
if (session == null) {
return;
@@ -4656,18 +4847,24 @@ public class AihrPracticeSeedService {
if (isBlank(sessionId) || session == null) {
return;
}
List<String> urls = snapshot(session.traineeAudioUrls());
if (urls.stream().allMatch(this::isBlank)) {
List<String> traineeUrls = snapshot(session.traineeAudioUrls());
List<Long> traineeOssIds = snapshotLong(session.traineeAudioOssIds());
List<Long> customerOssIds = snapshotLong(session.customerAudioOssIds());
boolean hasTraineeAudio = !traineeUrls.stream().allMatch(this::isBlank)
|| traineeOssIds.stream().anyMatch(Objects::nonNull);
boolean hasCustomerAudio = customerOssIds.stream().anyMatch(Objects::nonNull);
if (!hasTraineeAudio && !hasCustomerAudio) {
return;
}
List<Long> ossIds = snapshotLong(session.traineeAudioOssIds());
ensureAudioTable();
String id = sessionId.trim();
LocalDateTime now = LocalDateTime.now();
jdbcTemplate.update("DELETE FROM aihr_practice_audio WHERE tenant_id = ? AND session_id = ?", tenantId(), id);
for (int i = 0; i < urls.size(); i++) {
String url = urls.get(i);
if (isBlank(url)) {
int traineeCount = Math.max(traineeUrls.size(), traineeOssIds.size());
for (int i = 0; i < traineeCount; i++) {
String url = i < traineeUrls.size() ? traineeUrls.get(i) : "";
Long ossId = i < traineeOssIds.size() ? traineeOssIds.get(i) : null;
if (isBlank(url) && ossId == null) {
continue;
}
jdbcTemplate.update("""
@@ -4678,12 +4875,30 @@ public class AihrPracticeSeedService {
tenantId(),
id,
i + 1,
i < ossIds.size() ? ossIds.get(i) : null,
url.trim(),
ossId,
firstNonBlank(url, ""),
Timestamp.valueOf(now),
Timestamp.valueOf(now)
);
bindPracticeAudio(ossId);
}
for (int i = 0; i < customerOssIds.size(); i++) {
Long ossId = customerOssIds.get(i);
if (ossId == null) {
continue;
}
jdbcTemplate.update("""
INSERT INTO aihr_practice_audio
(tenant_id, session_id, turn_index, role, oss_id, audio_url, create_time, update_time)
VALUES (?, ?, ?, 'customer', ?, '', ?, ?)
""",
tenantId(),
id,
i + 1,
ossId,
Timestamp.valueOf(now),
Timestamp.valueOf(now)
);
bindPracticeAudio(i < ossIds.size() ? ossIds.get(i) : null);
}
}
@@ -4762,12 +4977,23 @@ public class AihrPracticeSeedService {
List<String> audioUrls = snapshot(activeSession == null ? null : activeSession.traineeAudioUrls());
List<Long> audioOssIds = snapshotLong(activeSession == null ? null : activeSession.traineeAudioOssIds());
List<String> customerLines = snapshot(activeSession == null ? null : activeSession.customerLines());
List<Long> customerAudioOssIds = snapshotLong(activeSession == null ? null : activeSession.customerAudioOssIds());
List<TurnEvidence> evidenceList = snapshotEvidence(activeSession == null ? null : activeSession.turnEvidence());
for (int i = 0; i < scenario.rounds().size(); i++) {
RoundSeed round = scenario.rounds().get(i);
TurnEvidence evidence = i < evidenceList.size() ? evidenceList.get(i) : null;
String customerLine = i < customerLines.size() && !isBlank(customerLines.get(i)) ? customerLines.get(i) : round.customer();
turns.add(new DialogueResponse("customer", "AI业主", customerLine));
turns.add(new DialogueResponse(
"customer",
"AI业主",
customerLine,
"",
i < customerAudioOssIds.size() ? customerAudioOssIds.get(i) : null,
null,
null,
null,
""
));
if (i < replies.size() && !isBlank(replies.get(i))) {
turns.add(new DialogueResponse(
"trainee",
@@ -5564,6 +5790,9 @@ public class AihrPracticeSeedService {
private record SessionAnnotationSource(String scenarioId, String dialogueJson, String annotationsJson) {
}
private record FinishedSessionAudioSource(String scenarioId, String dialogueJson) {
}
private record ScoreSnapshot(Integer total, Integer taskCompletion, Integer responseTimeliness,
Integer compliance, Integer emotion, Integer communication, Integer marketing) {
}
@@ -5669,7 +5898,7 @@ public class AihrPracticeSeedService {
private record ActiveSession(String tenantId, String scenarioId, String trainee, String extPartyId, Long assignmentId, boolean mobile,
LocalDateTime startedAt, AtomicInteger currentRoundIndex, List<String> traineeReplies,
List<String> traineeAudioUrls, List<Long> traineeAudioOssIds,
List<String> customerLines, List<TurnEvidence> turnEvidence,
List<String> customerLines, List<Long> customerAudioOssIds, List<TurnEvidence> turnEvidence,
List<Long> promptPresentedAtMillis, List<Long> responseLatenciesMs) {
}
@@ -665,6 +665,50 @@ public class AihrPracticeSeedServiceTest {
assertTrue(mobileController.contains("answerDailyDrill(id, request, ownMobileExtPartyId(null), currentAppUserId())"));
}
@Test
@SuppressWarnings("unchecked")
public void m3GeneratedCustomerAudioIsBoundToTheOwnedTurnAndSurvivesIntoReviewStorage() throws Exception {
JdbcTemplate jdbcTemplate = mock(JdbcTemplate.class);
when(jdbcTemplate.query(contains("SELECT enabled"), any(RowMapper.class), eq("000000"), eq("complaint-water")))
.thenReturn(List.of());
when(jdbcTemplate.query(contains("SELECT scenario_code"), any(RowMapper.class), eq("000000"), eq("complaint-water")))
.thenReturn(List.of());
AihrPracticeSeedService service = new AihrPracticeSeedService(new ObjectMapper(), jdbcTemplate, null, null);
var session = service.start(new StartRequest("employee-a", "complaint-water", "mobile", null));
assertNull(service.attachGeneratedCustomerAudio(session.sessionId(), 0, session.customerText(), 9001L, "employee-a"));
Field sessionsField = AihrPracticeSeedService.class.getDeclaredField("activeSessions");
sessionsField.setAccessible(true);
Map<String, Object> sessions = (Map<String, Object>) sessionsField.get(service);
Object activeSession = sessions.get(session.sessionId());
Method customerAudioAccessor = activeSession.getClass().getDeclaredMethod("customerAudioOssIds");
customerAudioAccessor.setAccessible(true);
assertEquals(List.of(9001L), customerAudioAccessor.invoke(activeSession));
assertEquals(9001L, service.attachGeneratedCustomerAudio(session.sessionId(), 0, session.customerText(), 9002L, "employee-a"));
assertEquals(List.of(9002L), customerAudioAccessor.invoke(activeSession));
assertThrows(ServiceException.class, () -> service.attachGeneratedCustomerAudio(
session.sessionId(), 0, session.customerText(), 9002L, "employee-b"
));
assertThrows(ServiceException.class, () -> service.attachGeneratedCustomerAudio(
session.sessionId(), 0, "与当前业主话术不一致", 9002L, "employee-a"
));
String serviceSource = Files.readString(Path.of("src/main/java/org/dromara/aihr/service/AihrPracticeSeedService.java"));
String speechController = Files.readString(Path.of("src/main/java/org/dromara/aihr/controller/AihrSpeechController.java"));
assertTrue(serviceSource.contains("customerAudioOssIds"));
assertTrue(serviceSource.contains("VALUES (?, ?, ?, 'customer', ?, '', ?, ?)"));
assertTrue(serviceSource.contains("ON DUPLICATE KEY UPDATE oss_id = VALUES(oss_id)"));
assertTrue(serviceSource.contains("activeSessions.get(request.sessionId()) != session"));
assertTrue(serviceSource.contains("inPracticeAudioTransaction"));
assertTrue(serviceSource.contains("LIMIT 1\n FOR UPDATE"));
assertTrue(serviceSource.contains("maskSensitiveText(customerText)"));
assertTrue(speechController.contains("attachGeneratedCustomerAudio"));
assertTrue(speechController.contains("Long replacedOssId"));
assertTrue(speechController.contains("训练语音仅支持关联业主回合"));
}
@Test
public void startChecksScenarioEnabledBeforeSeedFallback() throws Exception {
String source = Files.readString(Path.of("src/main/java/org/dromara/aihr/service/AihrPracticeSeedService.java"));