4 changed files with 251 additions and 2 deletions
-
16frontend/src/types/models.ts
-
33frontend/src/views/ModelConfigManager.vue
-
27src/main/java/com/wok/supportbot/config/ChatModelFactory.java
-
177src/test/java/com/wok/supportbot/ChatModelFactoryTests.java
@ -0,0 +1,177 @@ |
|||
package com.wok.supportbot; |
|||
|
|||
import com.fasterxml.jackson.databind.JsonNode; |
|||
import com.fasterxml.jackson.databind.ObjectMapper; |
|||
import com.sun.net.httpserver.HttpServer; |
|||
import com.wok.supportbot.config.ChatModelFactory; |
|||
import com.wok.supportbot.entity.AiModelConfig; |
|||
import com.wok.supportbot.service.AiModelConfigService; |
|||
import org.junit.jupiter.api.AfterEach; |
|||
import org.junit.jupiter.api.BeforeEach; |
|||
import org.junit.jupiter.api.Test; |
|||
import org.junit.jupiter.api.extension.ExtendWith; |
|||
import org.junit.jupiter.params.ParameterizedTest; |
|||
import org.junit.jupiter.params.provider.CsvSource; |
|||
import org.junit.jupiter.params.provider.NullAndEmptySource; |
|||
import org.junit.jupiter.params.provider.ValueSource; |
|||
import org.mockito.InjectMocks; |
|||
import org.mockito.Mock; |
|||
import org.mockito.junit.jupiter.MockitoExtension; |
|||
import org.springframework.ai.chat.model.ChatModel; |
|||
import org.springframework.ai.chat.prompt.Prompt; |
|||
import org.springframework.ai.openai.OpenAiChatOptions; |
|||
|
|||
import java.io.IOException; |
|||
import java.net.InetSocketAddress; |
|||
import java.nio.charset.StandardCharsets; |
|||
import java.time.Duration; |
|||
import java.util.Map; |
|||
import java.util.concurrent.BlockingQueue; |
|||
import java.util.concurrent.LinkedBlockingQueue; |
|||
import java.util.concurrent.TimeUnit; |
|||
|
|||
import static org.junit.jupiter.api.Assertions.*; |
|||
import static org.mockito.Mockito.when; |
|||
|
|||
@ExtendWith(MockitoExtension.class) |
|||
class ChatModelFactoryTests { |
|||
|
|||
private static final ObjectMapper JSON = new ObjectMapper(); |
|||
private final BlockingQueue<JsonNode> requests = new LinkedBlockingQueue<>(); |
|||
private HttpServer server; |
|||
|
|||
@Mock |
|||
private AiModelConfigService configService; |
|||
@InjectMocks |
|||
private ChatModelFactory factory; |
|||
|
|||
@BeforeEach |
|||
void startLocalEndpoint() throws IOException { |
|||
server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); |
|||
server.createContext("/", exchange -> { |
|||
try (exchange) { |
|||
JsonNode request = JSON.readTree(exchange.getRequestBody()); |
|||
requests.add(request); |
|||
boolean streaming = request.path("stream").asBoolean(); |
|||
String response = streaming |
|||
? "data: {\"id\":\"test\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"test\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"ok\"},\"finish_reason\":null}]}\n\n" |
|||
+ "data: {\"id\":\"test\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"test\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n" |
|||
+ "data: [DONE]\n\n" |
|||
: "{\"id\":\"test\",\"object\":\"chat.completion\",\"created\":1,\"model\":\"test\",\"choices\":[{\"index\":0,\"message\":{\"role\":\"assistant\",\"content\":\"ok\"},\"finish_reason\":\"stop\"}]}"; |
|||
byte[] bytes = response.getBytes(StandardCharsets.UTF_8); |
|||
exchange.getResponseHeaders().set("Content-Type", streaming ? "text/event-stream" : "application/json"); |
|||
exchange.sendResponseHeaders(200, bytes.length); |
|||
exchange.getResponseBody().write(bytes); |
|||
} |
|||
}); |
|||
server.start(); |
|||
} |
|||
|
|||
@AfterEach |
|||
void stopLocalEndpoint() { |
|||
if (server != null) server.stop(0); |
|||
} |
|||
|
|||
@ParameterizedTest |
|||
@ValueSource(strings = {"doubao-seed-2-0-pro-260215", "doubao-seed-2-0-lite-260215", |
|||
"doubao-seed-2-0-mini-260215", "doubao-seed-2-0-code-preview-260215"}) |
|||
void seedChatDefaultsToMinimalOnTheWire(String modelName) throws Exception { |
|||
AiModelConfig config = config("CHAT", "volcengine", modelName); |
|||
assertRequest(model(config), "minimal"); |
|||
} |
|||
|
|||
@ParameterizedTest |
|||
@ValueSource(strings = {"minimal", "low", "medium", "high"}) |
|||
void explicitSeedEffortOverridesFastDefault(String effort) throws Exception { |
|||
AiModelConfig config = config("CHAT", "volcengine", "doubao-seed-2-0-pro-260215"); |
|||
config.setExtraConfig(Map.of("reasoningEffort", effort, "topP", 0.8)); |
|||
JsonNode request = assertRequest(model(config), effort); |
|||
assertEquals(0.8, request.path("top_p").asDouble()); |
|||
} |
|||
|
|||
@ParameterizedTest |
|||
@NullAndEmptySource |
|||
@ValueSource(strings = {" "}) |
|||
void clearedSeedEffortRestoresFastDefault(String cleared) throws Exception { |
|||
AiModelConfig config = config("CHAT", "volcengine", "doubao-seed-2-0-pro-260215"); |
|||
config.setExtraConfig(cleared == null ? Map.of() : Map.of("reasoningEffort", cleared)); |
|||
assertRequest(model(config), "minimal"); |
|||
} |
|||
|
|||
@ParameterizedTest |
|||
@CsvSource({"CHAT,openai,gpt-4o", "CHAT,deepseek,deepseek-chat", "CHAT,moonshot,moonshot-v1-8k", |
|||
"CHAT,volcengine,doubao-1-5-pro-32k", "CHAT,openai,doubao-seed-2-0-pro-260215", |
|||
"RAG_REWRITE,volcengine,doubao-seed-2-0-pro-260215", "RAG_REWRITE,volcengine,doubao-1-5-pro-32k"}) |
|||
void otherModelsAndRewriteRetainProviderDefaults(String appType, String provider, String name) throws Exception { |
|||
assertRequest(model(config(appType, provider, name)), null); |
|||
} |
|||
|
|||
@Test |
|||
void unsupportedRewriteModelDoesNotReceiveStaleReasoningSetting() throws Exception { |
|||
AiModelConfig config = config("RAG_REWRITE", "volcengine", "doubao-1-5-pro-32k"); |
|||
config.setExtraConfig(Map.of("reasoningEffort", "medium")); |
|||
assertRequest(model(config), null); |
|||
} |
|||
|
|||
@Test |
|||
void supportedRewriteModelHonorsExplicitEffort() throws Exception { |
|||
AiModelConfig config = config("RAG_REWRITE", "volcengine", "doubao-seed-2-0-pro-260215"); |
|||
config.setExtraConfig(Map.of("reasoningEffort", "low")); |
|||
assertRequest(model(config), "low"); |
|||
} |
|||
|
|||
@Test |
|||
void cacheRefreshAppliesChangedAndClearedEffort() throws Exception { |
|||
AiModelConfig config = config("CHAT", "volcengine", "doubao-seed-2-0-pro-260215"); |
|||
ChatModel initial = model(config); |
|||
assertSame(initial, factory.getChatModel("CHAT")); |
|||
assertRequest(initial, "minimal"); |
|||
|
|||
config.setExtraConfig(Map.of("reasoningEffort", "medium")); |
|||
factory.clearCache(); |
|||
ChatModel changed = factory.getChatModel("CHAT"); |
|||
assertNotSame(initial, changed); |
|||
assertRequest(changed, "medium"); |
|||
|
|||
config.setExtraConfig(Map.of()); |
|||
factory.clearCache(); |
|||
assertRequest(factory.getChatModel("CHAT"), "minimal"); |
|||
} |
|||
|
|||
@Test |
|||
void streamingUsesTheSameFastDefault() throws Exception { |
|||
ChatModel model = model(config("CHAT", "volcengine", "doubao-seed-2-0-pro-260215")); |
|||
var responses = model.stream(new Prompt("测试")).collectList().block(Duration.ofSeconds(5)); |
|||
assertNotNull(responses); |
|||
assertFalse(responses.isEmpty()); |
|||
JsonNode request = requests.poll(5, TimeUnit.SECONDS); |
|||
assertNotNull(request); |
|||
assertTrue(request.path("stream").asBoolean()); |
|||
assertEquals("minimal", request.path("reasoning_effort").asText()); |
|||
} |
|||
|
|||
private AiModelConfig config(String appType, String provider, String modelName) { |
|||
return AiModelConfig.builder() |
|||
.id(1L).appType(appType).provider(provider).modelName(modelName) |
|||
.apiKey("local-test-only").baseUrl("http://127.0.0.1:" + server.getAddress().getPort()) |
|||
.temperature(0.7).maxTokens(4096).build(); |
|||
} |
|||
|
|||
private ChatModel model(AiModelConfig config) { |
|||
when(configService.getActiveConfigWithFullKey(config.getAppType())).thenReturn(config); |
|||
return factory.getChatModel(config.getAppType()); |
|||
} |
|||
|
|||
private JsonNode assertRequest(ChatModel model, String expectedEffort) throws Exception { |
|||
assertEquals(expectedEffort, ((OpenAiChatOptions) model.getDefaultOptions()).getReasoningEffort()); |
|||
assertEquals("ok", model.call(new Prompt("测试")).getResult().getOutput().getText()); |
|||
JsonNode request = requests.poll(5, TimeUnit.SECONDS); |
|||
assertNotNull(request); |
|||
if (expectedEffort == null) assertFalse(request.has("reasoning_effort")); |
|||
else assertEquals(expectedEffort, request.path("reasoning_effort").asText()); |
|||
assertEquals(4096, request.path("max_tokens").asInt()); |
|||
assertEquals(0.7, request.path("temperature").asDouble()); |
|||
assertFalse(request.has("thinking")); |
|||
return request; |
|||
} |
|||
} |
|||
Write
Preview
Loading…
Cancel
Save
Reference in new issue