feat(web-ai): support qwen web search agent
This commit is contained in:
+188
-6
@@ -9,6 +9,8 @@ import org.dromara.aihr.webai.AihrWebSearchProviderService.RuntimeProvider;
|
||||
import org.dromara.common.core.exception.ServiceException;
|
||||
import org.springframework.stereotype.Service;
|
||||
|
||||
import java.io.BufferedReader;
|
||||
import java.io.StringReader;
|
||||
import java.net.Inet4Address;
|
||||
import java.net.Inet6Address;
|
||||
import java.net.InetAddress;
|
||||
@@ -18,8 +20,12 @@ import java.net.http.HttpRequest;
|
||||
import java.net.http.HttpResponse;
|
||||
import java.time.Duration;
|
||||
import java.util.ArrayList;
|
||||
import java.util.LinkedHashMap;
|
||||
import java.util.List;
|
||||
import java.util.Locale;
|
||||
import java.util.Map;
|
||||
import java.util.regex.Matcher;
|
||||
import java.util.regex.Pattern;
|
||||
|
||||
@Service
|
||||
@RequiredArgsConstructor
|
||||
@@ -27,6 +33,8 @@ public class AihrWebSearchClient {
|
||||
|
||||
private static final Duration CONNECT_TIMEOUT = Duration.ofSeconds(5);
|
||||
private static final Duration REQUEST_TIMEOUT = Duration.ofSeconds(12);
|
||||
private static final Duration QWEN_REQUEST_TIMEOUT = Duration.ofSeconds(40);
|
||||
private static final Pattern PUBLIC_URL = Pattern.compile("https?://[^\\s\\]\\)\\\"'<>]+");
|
||||
|
||||
private final AihrWebAiProperties properties;
|
||||
private final AihrWebSearchProviderService providerService;
|
||||
@@ -41,9 +49,6 @@ public class AihrWebSearchClient {
|
||||
}
|
||||
|
||||
SearchResult search(RuntimeProvider runtime, String question) {
|
||||
if (!"tavily".equalsIgnoreCase(clean(runtime.code()))) {
|
||||
throw new ServiceException("暂不支持当前全网检索提供方");
|
||||
}
|
||||
String endpoint = clean(runtime.endpoint());
|
||||
String apiKey = clean(runtime.apiKey());
|
||||
if (!isPublicHttpsUrl(endpoint)) {
|
||||
@@ -52,6 +57,17 @@ public class AihrWebSearchClient {
|
||||
if (apiKey.isBlank()) {
|
||||
throw new ServiceException("全网检索密钥未配置");
|
||||
}
|
||||
String code = clean(runtime.code()).toLowerCase(Locale.ROOT);
|
||||
if ("qwen".equals(code)) {
|
||||
return searchQwen(runtime, question);
|
||||
}
|
||||
if (!"tavily".equals(code)) {
|
||||
throw new ServiceException("暂不支持当前全网检索提供方");
|
||||
}
|
||||
return searchTavily(runtime, question);
|
||||
}
|
||||
|
||||
private SearchResult searchTavily(RuntimeProvider runtime, String question) {
|
||||
try {
|
||||
ObjectNode body = objectMapper.createObjectNode();
|
||||
body.put("query", question);
|
||||
@@ -59,9 +75,9 @@ public class AihrWebSearchClient {
|
||||
body.put("include_answer", "basic");
|
||||
body.put("include_raw_content", false);
|
||||
body.put("max_results", properties.boundedMaxResults());
|
||||
HttpRequest request = HttpRequest.newBuilder(URI.create(endpoint))
|
||||
HttpRequest request = HttpRequest.newBuilder(URI.create(runtime.endpoint()))
|
||||
.timeout(REQUEST_TIMEOUT)
|
||||
.header("Authorization", "Bearer " + apiKey)
|
||||
.header("Authorization", "Bearer " + runtime.apiKey())
|
||||
.header("Content-Type", "application/json")
|
||||
.POST(HttpRequest.BodyPublishers.ofString(objectMapper.writeValueAsString(body)))
|
||||
.build();
|
||||
@@ -84,6 +100,53 @@ public class AihrWebSearchClient {
|
||||
}
|
||||
}
|
||||
|
||||
private SearchResult searchQwen(RuntimeProvider runtime, String question) {
|
||||
String agentId = clean(runtime.agentId());
|
||||
String agentVersion = clean(runtime.agentVersion());
|
||||
if (agentId.isBlank() || agentVersion.isBlank()) {
|
||||
throw new ServiceException("千问 Agent 配置不完整");
|
||||
}
|
||||
try {
|
||||
ObjectNode body = objectMapper.createObjectNode();
|
||||
ObjectNode input = body.putObject("input");
|
||||
ObjectNode message = input.putArray("messages").addObject();
|
||||
message.put("role", "user");
|
||||
message.put("content", question);
|
||||
ObjectNode options = body.putObject("parameters").putObject("agent_options");
|
||||
options.put("agent_id", agentId);
|
||||
options.put("agent_version", agentVersion);
|
||||
options.put("agent_policy", "standard");
|
||||
options.put("forced_search", true);
|
||||
options.put("enable_citation", true);
|
||||
options.put("enable_text_image_mixed", false);
|
||||
options.put("related_video", false);
|
||||
options.put("enable_rec_question", false);
|
||||
body.put("stream", true);
|
||||
HttpRequest request = HttpRequest.newBuilder(URI.create(runtime.endpoint()))
|
||||
.timeout(QWEN_REQUEST_TIMEOUT)
|
||||
.header("Authorization", "Bearer " + runtime.apiKey())
|
||||
.header("Content-Type", "application/json")
|
||||
.POST(HttpRequest.BodyPublishers.ofString(objectMapper.writeValueAsString(body)))
|
||||
.build();
|
||||
HttpResponse<String> response = HttpClient.newBuilder()
|
||||
.connectTimeout(CONNECT_TIMEOUT)
|
||||
.followRedirects(HttpClient.Redirect.NEVER)
|
||||
.build()
|
||||
.send(request, HttpResponse.BodyHandlers.ofString());
|
||||
if (response.statusCode() < 200 || response.statusCode() >= 300) {
|
||||
throw new ServiceException("全网检索服务暂时不可用");
|
||||
}
|
||||
return parseQwenSseResponse(objectMapper, response.body(), properties.boundedMaxResults());
|
||||
} catch (InterruptedException error) {
|
||||
Thread.currentThread().interrupt();
|
||||
throw new ServiceException("全网检索已中断");
|
||||
} catch (ServiceException error) {
|
||||
throw error;
|
||||
} catch (Exception error) {
|
||||
throw new ServiceException("全网检索服务暂时不可用");
|
||||
}
|
||||
}
|
||||
|
||||
RuntimeProvider runtime(AihrKnowledgePrincipal principal) {
|
||||
RuntimeProvider configured = providerService.active(principal);
|
||||
if (configured != null) {
|
||||
@@ -96,7 +159,7 @@ public class AihrWebSearchClient {
|
||||
return null;
|
||||
}
|
||||
return new RuntimeProvider(-1L, clean(properties.getProvider()), clean(properties.getProvider()),
|
||||
clean(properties.getEndpoint()), clean(properties.getApiKey()), true, true, null);
|
||||
clean(properties.getEndpoint()), clean(properties.getApiKey()), "", "", true, true, null);
|
||||
}
|
||||
|
||||
static SearchResult parseTavilyResponse(ObjectMapper mapper, String body, int limit) throws Exception {
|
||||
@@ -124,6 +187,125 @@ public class AihrWebSearchClient {
|
||||
return new SearchResult(answer, List.copyOf(sources));
|
||||
}
|
||||
|
||||
static SearchResult parseQwenSseResponse(ObjectMapper mapper, String body, int limit) throws Exception {
|
||||
StringBuilder answer = new StringBuilder();
|
||||
Map<String, WebSource> sources = new LinkedHashMap<>();
|
||||
try (BufferedReader reader = new BufferedReader(new StringReader(body == null ? "" : body))) {
|
||||
String line;
|
||||
while ((line = reader.readLine()) != null) {
|
||||
String chunk = line.trim();
|
||||
if (!chunk.startsWith("data:")) {
|
||||
continue;
|
||||
}
|
||||
String json = chunk.substring(5).trim();
|
||||
if (json.isBlank() || "[DONE]".equals(json)) {
|
||||
continue;
|
||||
}
|
||||
JsonNode root = mapper.readTree(json);
|
||||
int statusCode = root.path("status_code").asInt(200);
|
||||
String code = clean(root.path("code").asText(""));
|
||||
if (statusCode >= 400 || (!code.isBlank() && !"200".equals(code) && !"0".equals(code))) {
|
||||
throw new ServiceException("全网检索服务暂时不可用");
|
||||
}
|
||||
JsonNode choices = root.path("output").path("choices");
|
||||
if (!choices.isArray() || choices.size() == 0) {
|
||||
continue;
|
||||
}
|
||||
JsonNode message = choices.get(0).path("message");
|
||||
String role = clean(message.path("role").asText(""));
|
||||
JsonNode content = message.path("content");
|
||||
if ("assistant".equals(role) && content.isTextual()) {
|
||||
appendStreamChunk(answer, content.asText(""));
|
||||
collectTextUrls(content.asText(""), sources, limit);
|
||||
}
|
||||
if ("tool".equals(role)) {
|
||||
collectSources(mapper, content, sources, limit);
|
||||
}
|
||||
collectSources(mapper, message.path("additional_kwargs"), sources, limit);
|
||||
collectSources(mapper, message.path("tool_calls"), sources, limit);
|
||||
}
|
||||
}
|
||||
List<WebSource> safeSources = List.copyOf(sources.values());
|
||||
return new SearchResult(safeSources.isEmpty() ? "" : truncate(clean(answer.toString()), 4000), safeSources);
|
||||
}
|
||||
|
||||
private static void appendStreamChunk(StringBuilder answer, String chunk) {
|
||||
if (chunk == null || chunk.isEmpty()) {
|
||||
return;
|
||||
}
|
||||
String current = answer.toString();
|
||||
if (!current.isEmpty() && chunk.startsWith(current)) {
|
||||
answer.setLength(0);
|
||||
}
|
||||
answer.append(chunk);
|
||||
}
|
||||
|
||||
private static void collectSources(ObjectMapper mapper, JsonNode node, Map<String, WebSource> sources, int limit) {
|
||||
if (node == null || node.isMissingNode() || node.isNull() || sources.size() >= safeLimit(limit)) {
|
||||
return;
|
||||
}
|
||||
if (node.isTextual()) {
|
||||
String value = node.asText("").trim();
|
||||
if ((value.startsWith("{") || value.startsWith("[")) && !value.isBlank()) {
|
||||
try {
|
||||
collectSources(mapper, mapper.readTree(value), sources, limit);
|
||||
} catch (Exception ignored) {
|
||||
collectTextUrls(value, sources, limit);
|
||||
}
|
||||
} else {
|
||||
collectTextUrls(value, sources, limit);
|
||||
}
|
||||
return;
|
||||
}
|
||||
if (node.isArray()) {
|
||||
for (JsonNode item : node) {
|
||||
collectSources(mapper, item, sources, limit);
|
||||
if (sources.size() >= safeLimit(limit)) {
|
||||
break;
|
||||
}
|
||||
}
|
||||
return;
|
||||
}
|
||||
if (!node.isObject()) {
|
||||
return;
|
||||
}
|
||||
String url = firstText(node, "url", "link", "source_url", "sourceUrl");
|
||||
if (isPublicHttpUrl(url)) {
|
||||
String title = firstText(node, "title", "name", "site_name", "siteName");
|
||||
sources.putIfAbsent(url, new WebSource(
|
||||
title.isBlank() ? "公开来源" : title,
|
||||
url,
|
||||
truncate(firstText(node, "snippet", "summary", "description", "content"), 600),
|
||||
node.path("score").asDouble(0D)
|
||||
));
|
||||
}
|
||||
node.fields().forEachRemaining(entry -> collectSources(mapper, entry.getValue(), sources, limit));
|
||||
}
|
||||
|
||||
private static void collectTextUrls(String text, Map<String, WebSource> sources, int limit) {
|
||||
Matcher matcher = PUBLIC_URL.matcher(text == null ? "" : text);
|
||||
while (matcher.find() && sources.size() < safeLimit(limit)) {
|
||||
String url = matcher.group().replaceAll("[.,;:!?,。;:!?]+$", "");
|
||||
if (isPublicHttpUrl(url)) {
|
||||
sources.putIfAbsent(url, new WebSource("公开来源", url, "", 0D));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private static String firstText(JsonNode node, String... fields) {
|
||||
for (String field : fields) {
|
||||
JsonNode value = node.path(field);
|
||||
if (value.isTextual() && !value.asText().isBlank()) {
|
||||
return clean(value.asText());
|
||||
}
|
||||
}
|
||||
return "";
|
||||
}
|
||||
|
||||
private static int safeLimit(int limit) {
|
||||
return Math.max(1, Math.min(limit, 8));
|
||||
}
|
||||
|
||||
static boolean isPublicHttpUrl(String value) {
|
||||
try {
|
||||
URI uri = URI.create(clean(value));
|
||||
|
||||
+53
-18
@@ -1,5 +1,6 @@
|
||||
package org.dromara.aihr.webai;
|
||||
|
||||
import com.fasterxml.jackson.annotation.JsonProperty;
|
||||
import lombok.RequiredArgsConstructor;
|
||||
import org.dromara.aihr.knowledge.domain.AihrKnowledgePrincipal;
|
||||
import org.dromara.common.core.exception.ServiceException;
|
||||
@@ -14,6 +15,7 @@ import java.sql.Statement;
|
||||
import java.time.LocalDateTime;
|
||||
import java.util.List;
|
||||
import java.util.Locale;
|
||||
import java.util.Set;
|
||||
|
||||
@Service
|
||||
@RequiredArgsConstructor
|
||||
@@ -26,11 +28,13 @@ public class AihrWebSearchProviderService {
|
||||
public List<ProviderResponse> list(AihrKnowledgePrincipal principal) {
|
||||
ensureTable();
|
||||
return jdbcTemplate.query("""
|
||||
select id, provider_name, provider_code, endpoint, api_key, enabled, last_test_ok, last_test_time
|
||||
select id, provider_name, provider_code, endpoint, api_key, agent_id, agent_version,
|
||||
enabled, last_test_ok, last_test_time
|
||||
from aihr_web_search_provider where tenant_id = ? order by id
|
||||
""", (rs, rowNum) -> new ProviderResponse(
|
||||
rs.getLong("id"), rs.getString("provider_name"), rs.getString("provider_code"),
|
||||
rs.getString("endpoint"), hasText(rs.getString("api_key")), rs.getBoolean("enabled"),
|
||||
rs.getString("endpoint"), rs.getString("agent_id"), rs.getString("agent_version"),
|
||||
hasText(rs.getString("api_key")), rs.getBoolean("enabled"),
|
||||
rs.getBoolean("last_test_ok"), rs.getObject("last_test_time", LocalDateTime.class)
|
||||
), principal.tenantId());
|
||||
}
|
||||
@@ -42,14 +46,17 @@ public class AihrWebSearchProviderService {
|
||||
jdbcTemplate.update(connection -> {
|
||||
PreparedStatement statement = connection.prepareStatement("""
|
||||
insert into aihr_web_search_provider
|
||||
(tenant_id, provider_name, provider_code, endpoint, api_key, enabled, create_time, update_time)
|
||||
values (?, ?, ?, ?, ?, 0, now(), now())
|
||||
(tenant_id, provider_name, provider_code, endpoint, api_key, agent_id, agent_version,
|
||||
enabled, create_time, update_time)
|
||||
values (?, ?, ?, ?, ?, ?, ?, 0, now(), now())
|
||||
""", Statement.RETURN_GENERATED_KEYS);
|
||||
statement.setString(1, principal.tenantId());
|
||||
statement.setString(2, data.name());
|
||||
statement.setString(3, data.code());
|
||||
statement.setString(4, data.endpoint());
|
||||
statement.setString(5, secretCodec.encrypt(data.apiKey()));
|
||||
statement.setString(6, data.agentId());
|
||||
statement.setString(7, data.agentVersion());
|
||||
return statement;
|
||||
}, keyHolder);
|
||||
Number key = keyHolder.getKey();
|
||||
@@ -63,11 +70,13 @@ public class AihrWebSearchProviderService {
|
||||
int updated = jdbcTemplate.update("""
|
||||
update aihr_web_search_provider
|
||||
set provider_name = ?, provider_code = ?, endpoint = ?,
|
||||
agent_id = ?, agent_version = ?,
|
||||
api_key = case when ? is null then api_key else ? end,
|
||||
enabled = 0, last_test_ok = 0, last_test_time = null,
|
||||
update_time = now()
|
||||
where tenant_id = ? and id = ?
|
||||
""", data.name(), data.code(), data.endpoint(), apiKey, apiKey, principal.tenantId(), request.id());
|
||||
""", data.name(), data.code(), data.endpoint(), data.agentId(), data.agentVersion(),
|
||||
apiKey, apiKey, principal.tenantId(), request.id());
|
||||
if (updated == 0) {
|
||||
throw new IllegalArgumentException("全网检索提供方不存在");
|
||||
}
|
||||
@@ -79,7 +88,9 @@ public class AihrWebSearchProviderService {
|
||||
ensureTable();
|
||||
if (enabled) {
|
||||
RuntimeProvider selected = provider(principal, id);
|
||||
if (!AihrWebSearchClient.isPublicHttpsUrl(selected.endpoint()) || !hasText(selected.apiKey())) {
|
||||
if (!AihrWebSearchClient.isPublicHttpsUrl(selected.endpoint()) || !hasText(selected.apiKey())
|
||||
|| ("qwen".equals(selected.code())
|
||||
&& (!hasText(selected.agentId()) || !hasText(selected.agentVersion())))) {
|
||||
throw new IllegalArgumentException("请先配置安全的公开接口和访问密钥");
|
||||
}
|
||||
if (!selected.lastTestOk() || selected.lastTestTime() == null) {
|
||||
@@ -112,7 +123,8 @@ public class AihrWebSearchProviderService {
|
||||
ensureTable();
|
||||
try {
|
||||
List<RuntimeProvider> rows = jdbcTemplate.query("""
|
||||
select id, provider_name, provider_code, endpoint, api_key, enabled, last_test_ok, last_test_time
|
||||
select id, provider_name, provider_code, endpoint, api_key, agent_id, agent_version,
|
||||
enabled, last_test_ok, last_test_time
|
||||
from aihr_web_search_provider
|
||||
where tenant_id = ? and enabled = 1
|
||||
order by id desc limit 1
|
||||
@@ -126,7 +138,8 @@ public class AihrWebSearchProviderService {
|
||||
RuntimeProvider provider(AihrKnowledgePrincipal principal, Long id) {
|
||||
ensureTable();
|
||||
List<RuntimeProvider> rows = jdbcTemplate.query("""
|
||||
select id, provider_name, provider_code, endpoint, api_key, enabled, last_test_ok, last_test_time
|
||||
select id, provider_name, provider_code, endpoint, api_key, agent_id, agent_version,
|
||||
enabled, last_test_ok, last_test_time
|
||||
from aihr_web_search_provider where tenant_id = ? and id = ?
|
||||
""", (rs, rowNum) -> providerRow(rs), principal.tenantId(), id);
|
||||
if (rows.isEmpty()) {
|
||||
@@ -137,17 +150,20 @@ public class AihrWebSearchProviderService {
|
||||
|
||||
private RuntimeProvider providerRow(java.sql.ResultSet rs) throws java.sql.SQLException {
|
||||
return new RuntimeProvider(rs.getLong("id"), rs.getString("provider_name"), rs.getString("provider_code"),
|
||||
rs.getString("endpoint"), secretCodec.decrypt(rs.getString("api_key")), rs.getBoolean("enabled"),
|
||||
rs.getString("endpoint"), secretCodec.decrypt(rs.getString("api_key")),
|
||||
rs.getString("agent_id"), rs.getString("agent_version"), rs.getBoolean("enabled"),
|
||||
rs.getBoolean("last_test_ok"), rs.getObject("last_test_time", LocalDateTime.class));
|
||||
}
|
||||
|
||||
private ProviderResponse responseById(AihrKnowledgePrincipal principal, Long id) {
|
||||
List<ProviderResponse> rows = jdbcTemplate.query("""
|
||||
select id, provider_name, provider_code, endpoint, api_key, enabled, last_test_ok, last_test_time
|
||||
select id, provider_name, provider_code, endpoint, api_key, agent_id, agent_version,
|
||||
enabled, last_test_ok, last_test_time
|
||||
from aihr_web_search_provider where tenant_id = ? and id = ?
|
||||
""", (rs, rowNum) -> new ProviderResponse(
|
||||
rs.getLong("id"), rs.getString("provider_name"), rs.getString("provider_code"),
|
||||
rs.getString("endpoint"), hasText(rs.getString("api_key")), rs.getBoolean("enabled"),
|
||||
rs.getString("endpoint"), rs.getString("agent_id"), rs.getString("agent_version"),
|
||||
hasText(rs.getString("api_key")), rs.getBoolean("enabled"),
|
||||
rs.getBoolean("last_test_ok"), rs.getObject("last_test_time", LocalDateTime.class)
|
||||
), principal.tenantId(), id);
|
||||
if (rows.isEmpty()) {
|
||||
@@ -166,8 +182,8 @@ public class AihrWebSearchProviderService {
|
||||
if (name.isBlank() || name.length() > 100) {
|
||||
throw new IllegalArgumentException("提供方名称不能为空且不能超过100字");
|
||||
}
|
||||
if (!"tavily".equals(code)) {
|
||||
throw new IllegalArgumentException("当前仅支持 tavily");
|
||||
if (!Set.of("tavily", "qwen").contains(code)) {
|
||||
throw new IllegalArgumentException("当前仅支持 Tavily 或千问联网检索");
|
||||
}
|
||||
if (!AihrWebSearchClient.isPublicHttpsUrl(endpoint)) {
|
||||
throw new IllegalArgumentException("检索地址必须是可公开访问的 HTTPS 地址");
|
||||
@@ -179,7 +195,20 @@ public class AihrWebSearchProviderService {
|
||||
if (apiKey.length() > 1000) {
|
||||
throw new IllegalArgumentException("访问密钥长度超限");
|
||||
}
|
||||
return new ProviderData(name, code, endpoint, apiKey);
|
||||
String agentId = clean(request.agentId());
|
||||
String agentVersion = clean(request.agentVersion()).toLowerCase(Locale.ROOT);
|
||||
if ("qwen".equals(code)) {
|
||||
if (!agentId.startsWith("aid-") || agentId.length() > 100) {
|
||||
throw new IllegalArgumentException("千问 Agent ID 格式不正确");
|
||||
}
|
||||
if (!Set.of("beta", "release").contains(agentVersion)) {
|
||||
throw new IllegalArgumentException("千问 Agent Version 仅支持 beta 或 release");
|
||||
}
|
||||
} else {
|
||||
agentId = "";
|
||||
agentVersion = "";
|
||||
}
|
||||
return new ProviderData(name, code, endpoint, apiKey, agentId, agentVersion);
|
||||
}
|
||||
|
||||
private void ensureTable() {
|
||||
@@ -195,6 +224,7 @@ public class AihrWebSearchProviderService {
|
||||
`id` bigint NOT NULL AUTO_INCREMENT, `tenant_id` varchar(20) NOT NULL,
|
||||
`provider_name` varchar(100) NOT NULL, `provider_code` varchar(40) NOT NULL DEFAULT 'tavily',
|
||||
`endpoint` varchar(500) NOT NULL, `api_key` varchar(1000) DEFAULT NULL,
|
||||
`agent_id` varchar(100) DEFAULT NULL, `agent_version` varchar(20) DEFAULT NULL,
|
||||
`enabled` tinyint NOT NULL DEFAULT 0,
|
||||
`last_test_ok` tinyint NOT NULL DEFAULT 0,
|
||||
`last_test_time` datetime DEFAULT NULL,
|
||||
@@ -216,18 +246,23 @@ public class AihrWebSearchProviderService {
|
||||
return value == null ? "" : value.trim();
|
||||
}
|
||||
|
||||
public record ProviderRequest(Long id, String name, String providerCode, String endpoint, String apiKey) {
|
||||
public record ProviderRequest(Long id, String name, String providerCode, String endpoint,
|
||||
@JsonProperty(access = JsonProperty.Access.WRITE_ONLY) String apiKey,
|
||||
String agentId, String agentVersion) {
|
||||
}
|
||||
|
||||
public record ProviderResponse(Long id, String name, String providerCode, String endpoint,
|
||||
String agentId, String agentVersion,
|
||||
boolean configured, boolean enabled, boolean lastTestOk,
|
||||
LocalDateTime lastTestTime) {
|
||||
}
|
||||
|
||||
record RuntimeProvider(Long id, String name, String code, String endpoint, String apiKey, boolean enabled,
|
||||
boolean lastTestOk, LocalDateTime lastTestTime) {
|
||||
record RuntimeProvider(Long id, String name, String code, String endpoint, String apiKey,
|
||||
String agentId, String agentVersion, boolean enabled, boolean lastTestOk,
|
||||
LocalDateTime lastTestTime) {
|
||||
}
|
||||
|
||||
private record ProviderData(String name, String code, String endpoint, String apiKey) {
|
||||
private record ProviderData(String name, String code, String endpoint, String apiKey,
|
||||
String agentId, String agentVersion) {
|
||||
}
|
||||
}
|
||||
|
||||
+37
-2
@@ -6,6 +6,7 @@ import org.dromara.aihr.knowledge.service.AihrKnowledgePrincipalResolver;
|
||||
import org.dromara.aihr.webai.AihrWebAiDto.QueryRequest;
|
||||
import org.dromara.aihr.webai.AihrWebSearchClient.SearchResult;
|
||||
import org.dromara.aihr.webai.AihrWebSearchClient.WebSource;
|
||||
import org.dromara.aihr.webai.AihrWebSearchProviderService.ProviderRequest;
|
||||
import org.dromara.aihr.webai.AihrWebSearchProviderService.RuntimeProvider;
|
||||
import org.dromara.common.core.exception.ServiceException;
|
||||
import org.junit.jupiter.api.Tag;
|
||||
@@ -38,6 +39,21 @@ class AihrWebAiSecurityTest {
|
||||
assertEquals("tavily", properties.getProvider());
|
||||
}
|
||||
|
||||
@Test
|
||||
@Tag("dev")
|
||||
void providerRequestSecretIsAcceptedButNeverSerialized() throws Exception {
|
||||
ObjectMapper mapper = new ObjectMapper();
|
||||
String json = """
|
||||
{"name":"千问","providerCode":"qwen","endpoint":"https://8.8.8.8/search",
|
||||
"apiKey":"must-not-reach-logs","agentId":"aid-test","agentVersion":"release"}
|
||||
""";
|
||||
|
||||
ProviderRequest request = mapper.readValue(json, ProviderRequest.class);
|
||||
|
||||
assertEquals("must-not-reach-logs", request.apiKey());
|
||||
assertFalse(mapper.writeValueAsString(request).contains("must-not-reach-logs"));
|
||||
}
|
||||
|
||||
@Test
|
||||
@Tag("dev")
|
||||
void searchEndpointRejectsLocalAndPrivateNetworks() {
|
||||
@@ -96,6 +112,25 @@ class AihrWebAiSecurityTest {
|
||||
assertEquals("https://8.8.8.8/article", result.sources().get(0).url());
|
||||
}
|
||||
|
||||
@Test
|
||||
@Tag("dev")
|
||||
void qwenSseKeepsGeneratedAnswerAndPublicToolSources() throws Exception {
|
||||
String body = """
|
||||
data: {"status_code":200,"code":"","output":{"choices":[{"message":{"role":"tool","content":"{\\"results\\":[{\\"title\\":\\"公开资料\\",\\"url\\":\\"https://8.8.8.8/qwen\\",\\"snippet\\":\\"公开摘要\\"},{\\"title\\":\\"内网\\",\\"url\\":\\"http://127.0.0.1/secret\\"}]}"}}]}}
|
||||
|
||||
data: {"status_code":200,"code":"","output":{"choices":[{"message":{"role":"assistant","content":"根据","additional_kwargs":{}}}]}}
|
||||
|
||||
data: {"status_code":200,"code":"","output":{"choices":[{"message":{"role":"assistant","content":"公开资料生成答案。","additional_kwargs":{}}}]}}
|
||||
""";
|
||||
|
||||
SearchResult result = AihrWebSearchClient.parseQwenSseResponse(new ObjectMapper(), body, 5);
|
||||
|
||||
assertEquals("根据公开资料生成答案。", result.answer());
|
||||
assertEquals(1, result.sources().size());
|
||||
assertEquals("公开资料", result.sources().get(0).title());
|
||||
assertEquals("https://8.8.8.8/qwen", result.sources().get(0).url());
|
||||
}
|
||||
|
||||
@Test
|
||||
@Tag("dev")
|
||||
void providerRequestsRequireHttpsEvenForPublicHosts() {
|
||||
@@ -103,7 +138,7 @@ class AihrWebAiSecurityTest {
|
||||
new AihrWebAiProperties(), mock(AihrWebSearchProviderService.class), new ObjectMapper()
|
||||
);
|
||||
RuntimeProvider provider = new RuntimeProvider(1L, "test", "tavily",
|
||||
"http://8.8.8.8/search", "secret", true, true, LocalDateTime.now());
|
||||
"http://8.8.8.8/search", "secret", "", "", true, true, LocalDateTime.now());
|
||||
|
||||
assertThrows(ServiceException.class, () -> client.search(provider, "test question"));
|
||||
}
|
||||
@@ -119,7 +154,7 @@ class AihrWebAiSecurityTest {
|
||||
"000000", 7L, "app_user", "staff-7", Set.of("employee"), Set.of(), "mobile"
|
||||
);
|
||||
RuntimeProvider provider = new RuntimeProvider(1L, "Tavily", "tavily",
|
||||
"https://8.8.8.8/search", "secret", true, true, LocalDateTime.now());
|
||||
"https://8.8.8.8/search", "secret", "", "", true, true, LocalDateTime.now());
|
||||
when(resolver.current()).thenReturn(principal);
|
||||
when(searchClient.runtime(principal)).thenReturn(provider);
|
||||
when(searchClient.search(provider, "物业行业新规"))
|
||||
|
||||
Reference in New Issue
Block a user