7 changed files with 560 additions and 62 deletions
-
50src/main/java/com/wok/supportbot/app/AssistantApp.java
-
10src/main/java/com/wok/supportbot/app/ChatResult.java
-
59src/main/java/com/wok/supportbot/app/SourceReference.java
-
73src/main/java/com/wok/supportbot/controller/AiController.java
-
14src/main/java/com/wok/supportbot/controller/OpenApiController.java
-
250src/test/java/com/wok/supportbot/AnswerTransportTests.java
-
166src/test/java/com/wok/supportbot/ChatResultEndpointTests.java
@ -0,0 +1,59 @@ |
|||||
|
package com.wok.supportbot.app; |
||||
|
|
||||
|
import com.fasterxml.jackson.annotation.JsonInclude; |
||||
|
import org.springframework.ai.document.Document; |
||||
|
|
||||
|
import java.util.List; |
||||
|
import java.util.Map; |
||||
|
|
||||
|
/** Public citation metadata from the documents actually used by this answer. */ |
||||
|
@JsonInclude(JsonInclude.Include.ALWAYS) |
||||
|
public record SourceReference(String documentId, String title, String sourceName, |
||||
|
Integer chunkIndex, Double score, String snippet) { |
||||
|
|
||||
|
public SourceReference { |
||||
|
if (snippet != null && snippet.length() > 160) { |
||||
|
int end = Character.isHighSurrogate(snippet.charAt(158)) ? 158 : 159; |
||||
|
snippet = snippet.substring(0, end) + "…"; |
||||
|
} |
||||
|
} |
||||
|
|
||||
|
public static List<SourceReference> fromDocuments(List<Document> documents) { |
||||
|
if (documents == null || documents.isEmpty()) { |
||||
|
return List.of(); |
||||
|
} |
||||
|
return documents.stream().map(SourceReference::fromDocument).toList(); |
||||
|
} |
||||
|
|
||||
|
private static SourceReference fromDocument(Document document) { |
||||
|
Map<String, Object> metadata = document.getMetadata(); |
||||
|
// Preserve the published source endpoint's distance semantics; other retrievers expose score. |
||||
|
Object score = metadata.get("distance") != null ? metadata.get("distance") : metadata.get("score"); |
||||
|
return new SourceReference(stringValue(metadata.get("documentId")), |
||||
|
stringValue(metadata.get("title")), stringValue(metadata.get("sourceName")), |
||||
|
integerValue(metadata.get("chunkIndex")), doubleValue(score), document.getText()); |
||||
|
} |
||||
|
|
||||
|
private static String stringValue(Object value) { |
||||
|
return value == null ? null : value.toString(); |
||||
|
} |
||||
|
|
||||
|
private static Integer integerValue(Object value) { |
||||
|
if (value == null) return null; |
||||
|
try { |
||||
|
return Integer.valueOf(value.toString()); |
||||
|
} catch (NumberFormatException ignored) { |
||||
|
return null; |
||||
|
} |
||||
|
} |
||||
|
|
||||
|
private static Double doubleValue(Object value) { |
||||
|
if (value == null) return null; |
||||
|
try { |
||||
|
double number = Double.parseDouble(value.toString()); |
||||
|
return Double.isFinite(number) ? number : null; |
||||
|
} catch (NumberFormatException ignored) { |
||||
|
return null; |
||||
|
} |
||||
|
} |
||||
|
} |
||||
@ -0,0 +1,250 @@ |
|||||
|
package com.wok.supportbot; |
||||
|
|
||||
|
import com.fasterxml.jackson.databind.JsonNode; |
||||
|
import com.fasterxml.jackson.databind.ObjectMapper; |
||||
|
import com.wok.supportbot.app.AssistantApp; |
||||
|
import com.wok.supportbot.app.ChatContext; |
||||
|
import com.wok.supportbot.app.ChatPipeline; |
||||
|
import com.wok.supportbot.app.ChatRequest; |
||||
|
import com.wok.supportbot.app.ChatResult; |
||||
|
import com.wok.supportbot.app.SourceReference; |
||||
|
import com.wok.supportbot.chatmemory.DatabaseChatMemory; |
||||
|
import com.wok.supportbot.config.ChatModelFactory; |
||||
|
import com.wok.supportbot.config.SimpleCircuitBreaker; |
||||
|
import com.wok.supportbot.entity.LlmCallTrace; |
||||
|
import com.wok.supportbot.service.AiModelConfigService; |
||||
|
import com.wok.supportbot.service.ContentSafetyService; |
||||
|
import com.wok.supportbot.service.LlmCallTraceService; |
||||
|
import org.junit.jupiter.api.BeforeEach; |
||||
|
import org.junit.jupiter.api.Test; |
||||
|
import org.junit.jupiter.params.ParameterizedTest; |
||||
|
import org.junit.jupiter.params.provider.ValueSource; |
||||
|
import org.mockito.ArgumentCaptor; |
||||
|
import org.springframework.ai.chat.client.ChatClient; |
||||
|
import org.springframework.ai.chat.messages.AssistantMessage; |
||||
|
import org.springframework.ai.chat.model.ChatResponse; |
||||
|
import org.springframework.ai.chat.model.Generation; |
||||
|
import org.springframework.ai.document.Document; |
||||
|
import org.springframework.test.util.ReflectionTestUtils; |
||||
|
import reactor.core.publisher.Flux; |
||||
|
|
||||
|
import java.time.Duration; |
||||
|
import java.util.ArrayList; |
||||
|
import java.util.List; |
||||
|
import java.util.Map; |
||||
|
import java.util.Optional; |
||||
|
import java.util.concurrent.atomic.AtomicBoolean; |
||||
|
|
||||
|
import static org.junit.jupiter.api.Assertions.*; |
||||
|
import static org.mockito.ArgumentMatchers.*; |
||||
|
import static org.mockito.Mockito.*; |
||||
|
|
||||
|
class AnswerTransportTests { |
||||
|
private static final ObjectMapper JSON = new ObjectMapper(); |
||||
|
private AssistantApp app; |
||||
|
private ChatPipeline pipeline; |
||||
|
private ChatClient.ChatClientRequestSpec spec; |
||||
|
private ChatClient.CallResponseSpec call; |
||||
|
private ChatClient.StreamResponseSpec stream; |
||||
|
private LlmCallTraceService traces; |
||||
|
|
||||
|
@BeforeEach |
||||
|
@SuppressWarnings("unchecked") |
||||
|
void setUp() { |
||||
|
app = new AssistantApp(mock(ChatModelFactory.class), mock(DatabaseChatMemory.class)); |
||||
|
pipeline = mock(ChatPipeline.class); |
||||
|
traces = mock(LlmCallTraceService.class); |
||||
|
ContentSafetyService safety = mock(ContentSafetyService.class); |
||||
|
when(safety.mask(any())).thenAnswer(invocation -> invocation.getArgument(0)); |
||||
|
ReflectionTestUtils.setField(app, "chatPipeline", pipeline); |
||||
|
ReflectionTestUtils.setField(app, "llmCallTraceService", traces); |
||||
|
ReflectionTestUtils.setField(app, "contentSafetyService", safety); |
||||
|
ReflectionTestUtils.setField(app, "aiModelConfigService", mock(AiModelConfigService.class)); |
||||
|
ChatClient client = mock(ChatClient.class); |
||||
|
spec = mock(ChatClient.ChatClientRequestSpec.class, RETURNS_SELF); |
||||
|
call = mock(ChatClient.CallResponseSpec.class); |
||||
|
stream = mock(ChatClient.StreamResponseSpec.class); |
||||
|
when(client.prompt()).thenReturn(spec); |
||||
|
when(spec.call()).thenReturn(call); |
||||
|
when(spec.stream()).thenReturn(stream); |
||||
|
Map<String, ChatClient> cache = (Map<String, ChatClient>) ReflectionTestUtils.getField(app, "chatClientCache"); |
||||
|
cache.put("CHAT:none", client); |
||||
|
} |
||||
|
|
||||
|
@ParameterizedTest |
||||
|
@ValueSource(booleans = {true, false}) |
||||
|
void synchronousAnswerUsesOnlyThisBuildsDocuments(boolean rag) throws Exception { |
||||
|
ChatContext ctx = context(rag); |
||||
|
List<Document> documents = rag ? List.of(document()) : List.of(); |
||||
|
when(pipeline.buildRequest(ctx)).thenReturn(request(ctx, documents, null)); |
||||
|
when(call.chatResponse()).thenReturn(response("回答")); |
||||
|
|
||||
|
ChatResult result = app.chatWithEvents(ctx); |
||||
|
|
||||
|
assertEquals("回答", result.text()); |
||||
|
assertEquals(SourceReference.fromDocuments(documents), result.sources()); |
||||
|
assertTrue(result.suggestions().isEmpty()); |
||||
|
JsonNode json = JSON.valueToTree(result); |
||||
|
if (rag) { |
||||
|
assertEquals("9223372036854775807", json.at("/sources/0/documentId").asText()); |
||||
|
assertTrue(json.at("/sources/0/documentId").isTextual()); |
||||
|
} |
||||
|
verify(pipeline).buildRequest(ctx); |
||||
|
verifyNoMoreInteractions(pipeline); |
||||
|
verify(call).chatResponse(); |
||||
|
} |
||||
|
|
||||
|
@ParameterizedTest |
||||
|
@ValueSource(booleans = {true, false}) |
||||
|
void completedStreamCarriesExactlyOneMetadataChunkBeforeStop(boolean rag) throws Exception { |
||||
|
ChatContext ctx = context(rag); |
||||
|
List<Document> documents = rag ? List.of(document()) : List.of(); |
||||
|
when(pipeline.buildRequest(ctx)).thenReturn(request(ctx, documents, null)); |
||||
|
when(stream.chatResponse()).thenReturn(Flux.just(response("## "), response("回答\n"))); |
||||
|
|
||||
|
Flux<String> result = app.chatStreamOpenAi(ctx); |
||||
|
verifyNoInteractions(pipeline); |
||||
|
List<String> chunks = result.collectList().block(Duration.ofSeconds(5)); |
||||
|
assertEnvelope(chunks, SourceReference.fromDocuments(documents), "## 回答\n"); |
||||
|
verify(pipeline).buildRequest(ctx); |
||||
|
verifyNoMoreInteractions(pipeline); |
||||
|
verify(stream).chatResponse(); |
||||
|
ArgumentCaptor<LlmCallTrace> trace = ArgumentCaptor.forClass(LlmCallTrace.class); |
||||
|
verify(traces, timeout(1000)).recordAsync(trace.capture()); |
||||
|
assertEquals("## 回答\n", trace.getValue().getAiResponse()); |
||||
|
assertEquals("COMPLETE", trace.getValue().getStatus()); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void faqHasEmptySourcesWithoutModelInvocation() throws Exception { |
||||
|
ChatContext ctx = context(true); |
||||
|
when(pipeline.buildRequest(ctx)).thenReturn(request(ctx, List.of(), "标准答案")); |
||||
|
assertEnvelope(app.chatStreamOpenAi(ctx).collectList().block(Duration.ofSeconds(5)), List.of(), "标准答案"); |
||||
|
verify(pipeline).buildRequest(ctx); |
||||
|
verifyNoMoreInteractions(pipeline); |
||||
|
verifyNoInteractions(call, stream); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void synchronousFaqHasEmptySources() { |
||||
|
ChatContext ctx = context(true); |
||||
|
when(pipeline.buildRequest(ctx)).thenReturn(request(ctx, List.of(), "标准答案")); |
||||
|
ChatResult result = app.chatWithEvents(ctx); |
||||
|
assertEquals("标准答案", result.text()); |
||||
|
assertTrue(result.sources().isEmpty()); |
||||
|
verify(pipeline).buildRequest(ctx); |
||||
|
verifyNoMoreInteractions(pipeline); |
||||
|
verifyNoInteractions(call, stream); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void circuitFallbackDoesNotBuildOrRetrieve() throws Exception { |
||||
|
SimpleCircuitBreaker breaker = (SimpleCircuitBreaker) ReflectionTestUtils.getField(app, "aiCircuitBreaker"); |
||||
|
for (int i = 0; i < 3; i++) breaker.recordFailure(-1L); |
||||
|
ChatContext ctx = context(true); |
||||
|
assertTrue(app.chatWithEvents(ctx).sources().isEmpty()); |
||||
|
assertEnvelope(app.chatStreamOpenAi(ctx).collectList().block(Duration.ofSeconds(5)), |
||||
|
List.of(), "AI 服务暂时不可用,请稍后重试。"); |
||||
|
verifyNoInteractions(pipeline, call, stream); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void streamFailureDoesNotRepeatGenerationOrExposeUnusedSources() throws Exception { |
||||
|
ChatContext ctx = context(true); |
||||
|
when(pipeline.buildRequest(ctx)).thenReturn(request(ctx, List.of(document()), null)); |
||||
|
when(stream.chatResponse()).thenReturn(Flux.concat(Flux.just(response("部分回答")), |
||||
|
Flux.error(new IllegalStateException("模型超时")))); |
||||
|
List<String> chunks = app.chatStreamOpenAi(ctx).collectList().block(Duration.ofSeconds(5)); |
||||
|
assertEnvelope(chunks, List.of(), "部分回答抱歉,AI 服务调用失败:模型超时"); |
||||
|
verify(pipeline).buildRequest(ctx); |
||||
|
verifyNoMoreInteractions(pipeline); |
||||
|
verify(stream).chatResponse(); |
||||
|
verifyNoInteractions(call); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void cancellationPropagatesWithoutMetadataOrSecondRequest() throws Exception { |
||||
|
ChatContext ctx = context(true); |
||||
|
AtomicBoolean cancelled = new AtomicBoolean(); |
||||
|
when(pipeline.buildRequest(ctx)).thenReturn(request(ctx, List.of(document()), null)); |
||||
|
when(stream.chatResponse()).thenReturn(Flux.concat(Flux.just(response("首段")), Flux.<ChatResponse>never()) |
||||
|
.doOnCancel(() -> cancelled.set(true))); |
||||
|
List<String> chunks = app.chatStreamOpenAi(ctx).take(2).collectList().block(Duration.ofSeconds(5)); |
||||
|
assertEquals(2, chunks.size()); |
||||
|
assertEquals("首段", JSON.readTree(chunks.get(1)).at("/choices/0/delta/content").asText()); |
||||
|
assertFalse(chunks.stream().anyMatch(chunk -> chunk.contains("\"sources\"") || chunk.equals("[DONE]"))); |
||||
|
verify(traces, timeout(1000)).recordAsync(argThat(trace -> "CANCEL".equals(trace.getStatus()))); |
||||
|
assertTrue(cancelled.get()); |
||||
|
verify(pipeline).buildRequest(ctx); |
||||
|
verifyNoMoreInteractions(pipeline); |
||||
|
verify(stream).chatResponse(); |
||||
|
verifyNoInteractions(call); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void sourceSerializationPreservesNullableFieldsDistanceAndSnippetBoundaries() { |
||||
|
Document doc = document(); |
||||
|
SourceReference source = SourceReference.fromDocuments(List.of(doc)).get(0); |
||||
|
assertEquals("9223372036854775807", source.documentId()); |
||||
|
assertEquals(0.25, source.score()); |
||||
|
assertEquals(2, source.chunkIndex()); |
||||
|
assertEquals(160, source.snippet().length()); |
||||
|
assertTrue(source.snippet().endsWith("…")); |
||||
|
SourceReference nullable = SourceReference.fromDocuments(List.of(new Document("short"))).get(0); |
||||
|
JsonNode json = JSON.copy().setSerializationInclusion(com.fasterxml.jackson.annotation.JsonInclude.Include.NON_NULL) |
||||
|
.valueToTree(nullable); |
||||
|
assertEquals(6, json.size()); |
||||
|
assertTrue(json.get("documentId").isNull()); |
||||
|
assertTrue(json.get("score").isNull()); |
||||
|
SourceReference unicode = new SourceReference(null, null, null, null, null, "a".repeat(158) + "😀xx"); |
||||
|
assertTrue(unicode.snippet().length() <= 160); |
||||
|
assertFalse(Character.isHighSurrogate(unicode.snippet().charAt(unicode.snippet().length() - 2))); |
||||
|
ArrayList<SourceReference> mutable = new ArrayList<>(List.of(source)); |
||||
|
ChatResult result = new ChatResult("answer", null, null, mutable); |
||||
|
mutable.clear(); |
||||
|
assertEquals(List.of(source), result.sources()); |
||||
|
} |
||||
|
|
||||
|
private static void assertEnvelope(List<String> chunks, List<SourceReference> sources, String text) throws Exception { |
||||
|
assertNotNull(chunks); |
||||
|
assertEquals("[DONE]", chunks.get(chunks.size() - 1)); |
||||
|
JsonNode metadata = JSON.readTree(chunks.get(chunks.size() - 3)); |
||||
|
JsonNode stop = JSON.readTree(chunks.get(chunks.size() - 2)); |
||||
|
assertEquals("stop", stop.at("/choices/0/finish_reason").asText()); |
||||
|
assertEquals(0, metadata.get("choices").size()); |
||||
|
assertEquals(JSON.valueToTree(sources), metadata.get("sources")); |
||||
|
StringBuilder answer = new StringBuilder(); |
||||
|
int metadataCount = 0; |
||||
|
for (String raw : chunks.subList(0, chunks.size() - 1)) { |
||||
|
JsonNode chunk = JSON.readTree(raw); |
||||
|
assertEquals("chat.completion.chunk", chunk.get("object").asText()); |
||||
|
assertEquals(metadata.get("id"), chunk.get("id")); |
||||
|
assertEquals(metadata.get("model"), chunk.get("model")); |
||||
|
assertEquals(metadata.get("created"), chunk.get("created")); |
||||
|
if (chunk.has("sources")) metadataCount++; |
||||
|
answer.append(chunk.at("/choices/0/delta/content").asText("")); |
||||
|
} |
||||
|
assertEquals(1, metadataCount); |
||||
|
assertEquals(text, answer.toString()); |
||||
|
} |
||||
|
|
||||
|
private static Document document() { |
||||
|
return new Document("知识".repeat(100), Map.of("documentId", Long.MAX_VALUE, "title", "授权文档", |
||||
|
"sourceName", "manual.pdf", "chunkIndex", "2", "distance", 0.25, "score", 0.9)); |
||||
|
} |
||||
|
|
||||
|
private static ChatContext context(boolean rag) { |
||||
|
return new ChatContext("退货", "transport-chat", "CHAT", null, null, List.of(7L), |
||||
|
"NONE", rag, false, 11L, "售后", "account", null, null); |
||||
|
} |
||||
|
|
||||
|
private static ChatRequest request(ChatContext ctx, List<Document> documents, String faq) { |
||||
|
return new ChatRequest(ctx, ctx.message(), "system", Optional.ofNullable(faq), "system", |
||||
|
documents.isEmpty() ? null : "资料", documents.size(), faq != null ? "FAQ" : "RAG", |
||||
|
"VECTOR", documents, null); |
||||
|
} |
||||
|
|
||||
|
private static ChatResponse response(String text) { |
||||
|
return new ChatResponse(List.of(new Generation(new AssistantMessage(text)))); |
||||
|
} |
||||
|
} |
||||
@ -0,0 +1,166 @@ |
|||||
|
package com.wok.supportbot; |
||||
|
|
||||
|
import com.wok.supportbot.app.AssistantApp; |
||||
|
import com.wok.supportbot.app.ChatContext; |
||||
|
import com.wok.supportbot.app.ChatResult; |
||||
|
import com.wok.supportbot.app.SourceReference; |
||||
|
import com.wok.supportbot.config.RoleAccessConfig; |
||||
|
import com.wok.supportbot.controller.AiController; |
||||
|
import com.wok.supportbot.controller.OpenApiController; |
||||
|
import com.wok.supportbot.entity.ApiKey; |
||||
|
import com.wok.supportbot.rag.CategoryFilter; |
||||
|
import com.wok.supportbot.security.JwtTokenProvider; |
||||
|
import com.wok.supportbot.security.SdkAuthFilter; |
||||
|
import com.wok.supportbot.security.SdkJwtTokenProvider; |
||||
|
import com.wok.supportbot.service.ConversationService; |
||||
|
import com.wok.supportbot.service.CustomerServiceRoleService; |
||||
|
import com.wok.supportbot.service.CustomerServiceRoleService.RoleScope; |
||||
|
import io.jsonwebtoken.Claims; |
||||
|
import org.junit.jupiter.api.BeforeEach; |
||||
|
import org.junit.jupiter.api.Test; |
||||
|
import org.mockito.ArgumentCaptor; |
||||
|
import org.springframework.jdbc.core.JdbcTemplate; |
||||
|
import org.springframework.mock.web.MockHttpServletRequest; |
||||
|
import org.springframework.test.util.ReflectionTestUtils; |
||||
|
import org.springframework.test.web.servlet.MockMvc; |
||||
|
import org.springframework.test.web.servlet.setup.MockMvcBuilders; |
||||
|
|
||||
|
import java.util.List; |
||||
|
import java.util.Set; |
||||
|
|
||||
|
import static org.junit.jupiter.api.Assertions.*; |
||||
|
import static org.mockito.ArgumentMatchers.any; |
||||
|
import static org.mockito.Mockito.*; |
||||
|
import static org.springframework.test.web.servlet.request.MockMvcRequestBuilders.get; |
||||
|
import static org.springframework.test.web.servlet.result.MockMvcResultMatchers.*; |
||||
|
|
||||
|
class ChatResultEndpointTests { |
||||
|
private AssistantApp assistant; |
||||
|
private CustomerServiceRoleService roles; |
||||
|
private ConversationService conversations; |
||||
|
private RoleAccessConfig access; |
||||
|
private AiController controller; |
||||
|
private MockMvc mvc; |
||||
|
private SdkJwtTokenProvider tokens; |
||||
|
|
||||
|
@BeforeEach |
||||
|
void setUp() { |
||||
|
assistant = mock(AssistantApp.class); |
||||
|
roles = mock(CustomerServiceRoleService.class); |
||||
|
conversations = mock(ConversationService.class); |
||||
|
access = new RoleAccessConfig(); |
||||
|
controller = new AiController(); |
||||
|
ReflectionTestUtils.setField(controller, "assistantApp", assistant); |
||||
|
ReflectionTestUtils.setField(controller, "customerServiceRoleService", roles); |
||||
|
ReflectionTestUtils.setField(controller, "conversationService", conversations); |
||||
|
ReflectionTestUtils.setField(controller, "roleAccessConfig", access); |
||||
|
ReflectionTestUtils.setField(controller, "categoryFilter", mock(CategoryFilter.class)); |
||||
|
tokens = mock(SdkJwtTokenProvider.class); |
||||
|
SdkAuthFilter filter = new SdkAuthFilter(tokens, mock(JwtTokenProvider.class), mock(JdbcTemplate.class)); |
||||
|
mvc = MockMvcBuilders.standaloneSetup(controller).addFilters(filter).build(); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void resultRejectsMissingTokenAndUnauthorizedRoleBeforeGeneration() throws Exception { |
||||
|
mvc.perform(get("/ai/chat/result").servletPath("/ai/chat/result").param("message", "问题")) |
||||
|
.andExpect(status().isUnauthorized()); |
||||
|
authorize(); |
||||
|
mvc.perform(get("/ai/chat/result").servletPath("/ai/chat/result") |
||||
|
.header("Authorization", "Bearer valid").param("message", "问题").param("roleId", "99")) |
||||
|
.andExpect(status().isForbidden()); |
||||
|
verifyNoInteractions(assistant, roles, conversations); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void resultKeepsRoleScopeAccountBindingAndDecodedImages() throws Exception { |
||||
|
authorize(); |
||||
|
when(roles.getRoleScope(11L)).thenReturn(new RoleScope(true, "售后", "授权人设", List.of(7L), List.of("tool"))); |
||||
|
SourceReference source = new SourceReference("9223372036854775807", "授权文档", null, null, null, "片段"); |
||||
|
when(assistant.chatWithEvents(any())).thenReturn(new ChatResult("回答", List.of(), List.of(), List.of(source))); |
||||
|
mvc.perform(get("/ai/chat/result").servletPath("/ai/chat/result") |
||||
|
.header("Authorization", "Bearer valid").param("message", "问题").param("chatId", "chat") |
||||
|
.param("roleId", "11").param("accountId", " account ").param("systemPrompt", "越权人设") |
||||
|
.param("enableRag", "true").param("categoryIds", "99").param("imageUrls", "https%3A%2F%2Fexample.org%2Fa.png")) |
||||
|
.andExpect(status().isOk()).andExpect(content().contentTypeCompatibleWith("application/json")) |
||||
|
.andExpect(jsonPath("$.text").value("回答")) |
||||
|
.andExpect(jsonPath("$.data").doesNotExist()) |
||||
|
.andExpect(jsonPath("$.sources[0].documentId").value("9223372036854775807")) |
||||
|
.andExpect(jsonPath("$.mcpEvents").isArray()).andExpect(jsonPath("$.suggestions").isArray()); |
||||
|
ArgumentCaptor<ChatContext> context = ArgumentCaptor.forClass(ChatContext.class); |
||||
|
verify(assistant).chatWithEvents(context.capture()); |
||||
|
ChatContext ctx = context.getValue(); |
||||
|
assertEquals(List.of(7L), ctx.categoryIds()); |
||||
|
assertEquals("授权人设", ctx.systemPrompt()); |
||||
|
assertEquals(List.of("tool"), ctx.allowedMcpTools()); |
||||
|
assertEquals(List.of("https://example.org/a.png"), ctx.imageUrls()); |
||||
|
assertEquals("NONE", ctx.rewriteStrategy()); |
||||
|
assertTrue(ctx.enableRag()); |
||||
|
verify(conversations).bindConversation("chat", "account", 11L); |
||||
|
verifyNoMoreInteractions(assistant); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void strictIsolationAndExplicitRewriteApplyToBothSyncEndpoints() { |
||||
|
access.setStrictIsolation(true); |
||||
|
when(roles.getRoleScope(11L)).thenReturn(new RoleScope(true, "售后", "人设", List.of(), List.of())); |
||||
|
when(assistant.chatWithEvents(any())).thenReturn(new ChatResult("回答", List.of(), List.of(), List.of())); |
||||
|
assertEquals("回答", controller.chatSync("问题", "chat", 11L, "account", null, true, |
||||
|
"REWRITE", null, "99", null)); |
||||
|
ArgumentCaptor<ChatContext> context = ArgumentCaptor.forClass(ChatContext.class); |
||||
|
verify(assistant).chatWithEvents(context.capture()); |
||||
|
assertFalse(context.getValue().enableRag()); |
||||
|
assertTrue(context.getValue().categoryIds().isEmpty()); |
||||
|
assertEquals("REWRITE", context.getValue().rewriteStrategy()); |
||||
|
verifyNoMoreInteractions(assistant); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void openApiRetainsEnvelopeAndAddsSameAnswerSourcesWithNoneDefault() { |
||||
|
OpenApiController open = new OpenApiController(); |
||||
|
ReflectionTestUtils.setField(open, "assistantApp", assistant); |
||||
|
ReflectionTestUtils.setField(open, "customerServiceRoleService", roles); |
||||
|
ReflectionTestUtils.setField(open, "categoryFilter", mock(CategoryFilter.class)); |
||||
|
ReflectionTestUtils.setField(open, "roleAccessConfig", access); |
||||
|
when(roles.getRoleScope(11L)).thenReturn(new RoleScope(true, "售后", "人设", List.of(7L), List.of())); |
||||
|
SourceReference source = new SourceReference("9223372036854775807", "授权文档", null, null, null, "片段"); |
||||
|
when(assistant.chatWithEvents(any())).thenReturn(new ChatResult("回答", List.of(), List.of(), List.of(source))); |
||||
|
ApiKey key = new ApiKey(); |
||||
|
key.setId(3L); |
||||
|
key.setRoleIds("[11]"); |
||||
|
MockHttpServletRequest request = new MockHttpServletRequest(); |
||||
|
request.setAttribute("apiKey", key); |
||||
|
var result = open.chat("问题", "11", "chat", "99", null, true, request); |
||||
|
assertEquals(200, result.getStatusCode().value()); |
||||
|
assertEquals(true, result.getBody().get("success")); |
||||
|
var data = (java.util.Map<?, ?>) result.getBody().get("data"); |
||||
|
assertEquals("回答", data.get("reply")); |
||||
|
assertEquals(List.of(source), data.get("sources")); |
||||
|
ArgumentCaptor<ChatContext> context = ArgumentCaptor.forClass(ChatContext.class); |
||||
|
verify(assistant).chatWithEvents(context.capture()); |
||||
|
assertEquals("NONE", context.getValue().rewriteStrategy()); |
||||
|
assertEquals(List.of(7L), context.getValue().categoryIds()); |
||||
|
verifyNoMoreInteractions(assistant); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void explicitSourcesApiUsesTheSameCitationSerializer() { |
||||
|
ReflectionTestUtils.setField(controller, "chatPipeline", mock(com.wok.supportbot.app.ChatPipeline.class)); |
||||
|
when(roles.getRoleScope(11L)).thenReturn(new RoleScope(true, "售后", "人设", List.of(7L), List.of())); |
||||
|
var document = new org.springframework.ai.document.Document("片段".repeat(100), |
||||
|
java.util.Map.of("documentId", Long.MAX_VALUE, "distance", 0.2)); |
||||
|
when(assistant.retrieveSources(any())).thenReturn(List.of(document)); |
||||
|
var result = controller.chatSources("问题", "chat", null, 11L, "account", null, "99"); |
||||
|
assertEquals(SourceReference.fromDocuments(List.of(document)), result.get("data")); |
||||
|
ArgumentCaptor<ChatContext> context = ArgumentCaptor.forClass(ChatContext.class); |
||||
|
verify(assistant).retrieveSources(context.capture()); |
||||
|
assertEquals(List.of(7L), context.getValue().categoryIds()); |
||||
|
assertEquals("NONE", context.getValue().rewriteStrategy()); |
||||
|
verifyNoMoreInteractions(assistant); |
||||
|
} |
||||
|
|
||||
|
private void authorize() { |
||||
|
Claims claims = mock(Claims.class); |
||||
|
when(tokens.parseToken("valid")).thenReturn(claims); |
||||
|
when(tokens.getAllowedRoleIds(claims)).thenReturn(Set.of(11L)); |
||||
|
} |
||||
|
} |
||||
Write
Preview
Loading…
Cancel
Save
Reference in new issue