commit
This commit is contained in:
@@ -1,5 +1,7 @@
|
||||
package org.example.agent.tool;
|
||||
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
import org.springframework.ai.tool.annotation.Tool;
|
||||
import org.springframework.context.i18n.LocaleContextHolder;
|
||||
import org.springframework.stereotype.Component;
|
||||
@@ -8,12 +10,18 @@ import java.time.LocalDateTime;
|
||||
|
||||
@Component
|
||||
public class DateTimeTools {
|
||||
|
||||
private static final Logger logger = LoggerFactory.getLogger(DateTimeTools.class);
|
||||
|
||||
/** 工具名常量,用于动态构建提示词 */
|
||||
public static final String TOOL_GET_CURRENT_DATETIME = "getCurrentDateTime";
|
||||
|
||||
@Tool(description = "Get the current date and time in the user's timezone")
|
||||
@Tool(description = "Get the current date and time in the user's timezone. " +
|
||||
"IMPORTANT: Time changes constantly. Always call this tool when user asks about time, " +
|
||||
"even if there's a recent time query in the conversation history.")
|
||||
public String getCurrentDateTime() {
|
||||
return LocalDateTime.now().atZone(LocaleContextHolder.getTimeZone().toZoneId()).toString();
|
||||
String currentTime = LocalDateTime.now().atZone(LocaleContextHolder.getTimeZone().toZoneId()).toString();
|
||||
logger.debug("🕐 getCurrentDateTime 调用 - 返回时间: {}", currentTime);
|
||||
return currentTime;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -60,6 +60,17 @@ public class MilvusClientFactory {
|
||||
logger.info("collection '{}' 已存在", MilvusConstants.MILVUS_COLLECTION_NAME);
|
||||
}
|
||||
|
||||
// 3. 加载 collection 到内存(搜索必须)
|
||||
logger.info("正在加载 collection '{}' 到内存...", MilvusConstants.MILVUS_COLLECTION_NAME);
|
||||
R<RpcStatus> loadResp = client.loadCollection(LoadCollectionParam.newBuilder()
|
||||
.withCollectionName(MilvusConstants.MILVUS_COLLECTION_NAME)
|
||||
.build());
|
||||
if (loadResp.getStatus() == 0) {
|
||||
logger.info("collection '{}' 已加载", MilvusConstants.MILVUS_COLLECTION_NAME);
|
||||
} else {
|
||||
logger.warn("collection '{}' 加载失败: {}", MilvusConstants.MILVUS_COLLECTION_NAME, loadResp.getMessage());
|
||||
}
|
||||
|
||||
return client;
|
||||
|
||||
} catch (Exception e) {
|
||||
@@ -124,7 +135,7 @@ public class MilvusClientFactory {
|
||||
FieldType vectorField = FieldType.newBuilder()
|
||||
.withName("vector")
|
||||
.withDataType(DataType.FloatVector) // 改为 FloatVector
|
||||
.withDimension(MilvusConstants.VECTOR_DIM)
|
||||
.withDimension(milvusProperties.getVectorDim())
|
||||
.build();
|
||||
|
||||
FieldType contentField = FieldType.newBuilder()
|
||||
|
||||
@@ -15,6 +15,7 @@ public class MilvusProperties {
|
||||
private Long timeout = 10000L;
|
||||
private String token = "";
|
||||
private boolean secure = false;
|
||||
private int vectorDim = 1024;
|
||||
|
||||
public String getHost() {
|
||||
return host;
|
||||
@@ -80,6 +81,14 @@ public class MilvusProperties {
|
||||
this.secure = secure;
|
||||
}
|
||||
|
||||
public int getVectorDim() {
|
||||
return vectorDim;
|
||||
}
|
||||
|
||||
public void setVectorDim(int vectorDim) {
|
||||
this.vectorDim = vectorDim;
|
||||
}
|
||||
|
||||
public String getAddress() {
|
||||
return host + ":" + port;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,91 @@
|
||||
package org.example.config;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.embedding.EmbeddingModel;
|
||||
import org.springframework.beans.factory.annotation.Value;
|
||||
import org.springframework.context.annotation.Bean;
|
||||
import org.springframework.context.annotation.Configuration;
|
||||
import org.springframework.context.annotation.Primary;
|
||||
|
||||
/**
|
||||
* 模型路由配置 — 由 yml 驱动,不硬编码模型名。
|
||||
* <p>
|
||||
* 配置示例:
|
||||
* <pre>{@code
|
||||
* model-routing:
|
||||
* chat: deepseek
|
||||
* embedding: siliconflow
|
||||
* }</pre>
|
||||
* <p>
|
||||
* 匹配优先级:Bean 名 > 类名(均不区分大小写)。
|
||||
* 切换模型只改 yml + pom + 对应 api-key,Java 代码不动。
|
||||
*/
|
||||
@Configuration
|
||||
public class ModelRoutingConfig {
|
||||
|
||||
private static final Logger log = LoggerFactory.getLogger(ModelRoutingConfig.class);
|
||||
|
||||
@Value("${model-routing.chat:deepseek}")
|
||||
private String chatKeyword;
|
||||
|
||||
@Value("${model-routing.embedding:siliconflow}")
|
||||
private String embeddingKeyword;
|
||||
|
||||
@Bean
|
||||
@Primary
|
||||
public ChatModel chatModel(List<ChatModel> chatModels) {
|
||||
log.info("Chat 路由: keyword='{}', 可用: {}", chatKeyword,
|
||||
chatModels.stream().map(c -> c.getClass().getSimpleName()).toList());
|
||||
|
||||
for (ChatModel cm : chatModels) {
|
||||
if (matches(cm.getClass(), chatKeyword)) {
|
||||
log.info(" → 选中 {}", cm.getClass().getSimpleName());
|
||||
return cm;
|
||||
}
|
||||
}
|
||||
|
||||
log.warn(" → 未匹配, 回退到 {}", chatModels.get(0).getClass().getSimpleName());
|
||||
return chatModels.get(0);
|
||||
}
|
||||
|
||||
@Bean
|
||||
@Primary
|
||||
public EmbeddingModel embeddingModel(Map<String, EmbeddingModel> embeddingBeans) {
|
||||
log.info("Embedding 路由: keyword='{}', 可用: {}", embeddingKeyword, embeddingBeans.keySet());
|
||||
|
||||
// 先按 Bean 名匹配
|
||||
for (Map.Entry<String, EmbeddingModel> entry : embeddingBeans.entrySet()) {
|
||||
if (containsIgnoreCase(entry.getKey(), embeddingKeyword)) {
|
||||
log.info(" → Bean 名匹配: {} → {}", entry.getKey(),
|
||||
entry.getValue().getClass().getSimpleName());
|
||||
return entry.getValue();
|
||||
}
|
||||
}
|
||||
|
||||
// 再按类名匹配
|
||||
for (EmbeddingModel em : embeddingBeans.values()) {
|
||||
if (matches(em.getClass(), embeddingKeyword)) {
|
||||
log.info(" → 类名匹配: {}", em.getClass().getSimpleName());
|
||||
return em;
|
||||
}
|
||||
}
|
||||
|
||||
var first = embeddingBeans.values().iterator().next();
|
||||
log.warn(" → 未匹配, 回退到 {}", first.getClass().getSimpleName());
|
||||
return first;
|
||||
}
|
||||
|
||||
private boolean matches(Class<?> clazz, String keyword) {
|
||||
return containsIgnoreCase(clazz.getName(), keyword)
|
||||
|| containsIgnoreCase(clazz.getSimpleName(), keyword);
|
||||
}
|
||||
|
||||
private boolean containsIgnoreCase(String text, String keyword) {
|
||||
return text.toLowerCase().contains(keyword.toLowerCase());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
package org.example.config;
|
||||
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
import org.springframework.ai.document.MetadataMode;
|
||||
import org.springframework.ai.embedding.EmbeddingModel;
|
||||
import org.springframework.ai.openai.OpenAiEmbeddingModel;
|
||||
import org.springframework.ai.openai.OpenAiEmbeddingOptions;
|
||||
import org.springframework.ai.openai.api.OpenAiApi;
|
||||
import org.springframework.beans.factory.annotation.Value;
|
||||
import org.springframework.context.annotation.Bean;
|
||||
import org.springframework.context.annotation.Configuration;
|
||||
import org.springframework.web.client.RestClient;
|
||||
import org.springframework.web.reactive.function.client.WebClient;
|
||||
|
||||
/**
|
||||
* SiliconFlow Embedding 配置(BGE-M3, OpenAI 兼容协议, 1024维)
|
||||
* <p>
|
||||
* Chat 走 DeepSeek、Embedding 走 SiliconFlow,两者都是 OpenAI 兼容但地址不同,
|
||||
* 因此单独为 SiliconFlow 创建 OpenAiApi + EmbeddingModel Bean。
|
||||
*/
|
||||
@Configuration
|
||||
public class SiliconFlowEmbeddingConfig {
|
||||
|
||||
private static final Logger log = LoggerFactory.getLogger(SiliconFlowEmbeddingConfig.class);
|
||||
|
||||
@Value("${siliconflow.api-key}")
|
||||
private String apiKey;
|
||||
|
||||
@Value("${siliconflow.base-url}")
|
||||
private String baseUrl;
|
||||
|
||||
@Value("${siliconflow.embedding.model}")
|
||||
private String model;
|
||||
|
||||
@Bean
|
||||
public OpenAiApi siliconFlowApi(RestClient.Builder restClientBuilder, WebClient.Builder webClientBuilder) {
|
||||
log.info("创建 SiliconFlow OpenAiApi: {}", baseUrl);
|
||||
return OpenAiApi.builder()
|
||||
.baseUrl(baseUrl)
|
||||
.apiKey(apiKey)
|
||||
.restClientBuilder(restClientBuilder)
|
||||
.build();
|
||||
}
|
||||
|
||||
@Bean
|
||||
public EmbeddingModel siliconFlowEmbeddingModel(OpenAiApi siliconFlowApi) {
|
||||
log.info("创建 SiliconFlow EmbeddingModel, model: {}", model);
|
||||
return new OpenAiEmbeddingModel(siliconFlowApi, MetadataMode.EMBED,
|
||||
OpenAiEmbeddingOptions.builder()
|
||||
.model(model)
|
||||
.build());
|
||||
}
|
||||
}
|
||||
@@ -1,8 +1,5 @@
|
||||
package org.example.controller;
|
||||
|
||||
import com.alibaba.cloud.ai.dashscope.api.DashScopeApi;
|
||||
import com.alibaba.cloud.ai.dashscope.chat.DashScopeChatModel;
|
||||
import com.alibaba.cloud.ai.dashscope.chat.DashScopeChatOptions;
|
||||
import com.alibaba.cloud.ai.graph.NodeOutput;
|
||||
import com.alibaba.cloud.ai.graph.OverAllState;
|
||||
import com.alibaba.cloud.ai.graph.agent.ReactAgent;
|
||||
@@ -14,6 +11,7 @@ import org.example.service.AiOpsService;
|
||||
import org.example.service.ChatService;
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.tool.ToolCallback;
|
||||
import org.springframework.ai.tool.ToolCallbackProvider;
|
||||
import org.springframework.beans.factory.annotation.Autowired;
|
||||
@@ -46,7 +44,7 @@ public class ChatController {
|
||||
@Autowired
|
||||
private ChatService chatService;
|
||||
|
||||
@Autowired
|
||||
@Autowired(required = false)
|
||||
private ToolCallbackProvider tools;
|
||||
|
||||
private final ExecutorService executor = Executors.newCachedThreadPool();
|
||||
@@ -79,9 +77,8 @@ public class ChatController {
|
||||
List<Map<String, String>> history = session.getHistory();
|
||||
logger.info("会话历史消息对数: {}", history.size() / 2);
|
||||
|
||||
// 创建 DashScope API 和 ChatModel
|
||||
DashScopeApi dashScopeApi = chatService.createDashScopeApi();
|
||||
DashScopeChatModel chatModel = chatService.createStandardChatModel(dashScopeApi);
|
||||
// 获取注入的 ChatModel
|
||||
ChatModel chatModel = chatService.getChatModel();
|
||||
|
||||
// 记录可用工具
|
||||
chatService.logAvailableTools();
|
||||
@@ -167,9 +164,8 @@ public class ChatController {
|
||||
List<Map<String, String>> history = session.getHistory();
|
||||
logger.info("ReactAgent 会话历史消息对数: {}", history.size() / 2);
|
||||
|
||||
// 创建 DashScope API 和 ChatModel
|
||||
DashScopeApi dashScopeApi = chatService.createDashScopeApi();
|
||||
DashScopeChatModel chatModel = chatService.createStandardChatModel(dashScopeApi);
|
||||
// 获取注入的 ChatModel
|
||||
ChatModel chatModel = chatService.getChatModel();
|
||||
|
||||
// 记录可用工具
|
||||
chatService.logAvailableTools();
|
||||
@@ -289,18 +285,9 @@ public class ChatController {
|
||||
try {
|
||||
logger.info("收到 AI 智能运维请求 - 启动多 Agent 协作流程");
|
||||
|
||||
DashScopeApi dashScopeApi = chatService.createDashScopeApi();
|
||||
DashScopeChatModel chatModel = DashScopeChatModel.builder()
|
||||
.dashScopeApi(dashScopeApi)
|
||||
.defaultOptions(DashScopeChatOptions.builder()
|
||||
.withModel(DashScopeChatModel.DEFAULT_MODEL_NAME)
|
||||
.withTemperature(0.3)
|
||||
.withMaxToken(8000)
|
||||
.withTopP(0.9)
|
||||
.build())
|
||||
.build();
|
||||
ChatModel chatModel = chatService.getChatModel();
|
||||
|
||||
ToolCallback[] toolCallbacks = tools.getToolCallbacks();
|
||||
ToolCallback[] toolCallbacks = tools != null ? tools.getToolCallbacks() : new ToolCallback[0];
|
||||
|
||||
emitter.send(SseEmitter.event().name("message").data(SseMessage.content("正在读取告警并拆解任务...\n")));
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
package org.example.service;
|
||||
|
||||
import com.alibaba.cloud.ai.dashscope.chat.DashScopeChatModel;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import com.alibaba.cloud.ai.graph.OverAllState;
|
||||
import com.alibaba.cloud.ai.graph.agent.ReactAgent;
|
||||
import com.alibaba.cloud.ai.graph.agent.flow.agent.SupervisorAgent;
|
||||
@@ -48,7 +48,7 @@ public class AiOpsService {
|
||||
* @return 分析结果状态
|
||||
* @throws GraphRunnerException 如果 Agent 执行失败
|
||||
*/
|
||||
public Optional<OverAllState> executeAiOpsAnalysis(DashScopeChatModel chatModel, ToolCallback[] toolCallbacks) throws GraphRunnerException {
|
||||
public Optional<OverAllState> executeAiOpsAnalysis(ChatModel chatModel, ToolCallback[] toolCallbacks) throws GraphRunnerException {
|
||||
logger.info("开始执行 AI Ops 多 Agent 协作流程");
|
||||
|
||||
// 构建 Planner 和 Executor Agent
|
||||
@@ -67,7 +67,18 @@ public class AiOpsService {
|
||||
String taskPrompt = "你是企业级 SRE,接到了自动化告警排查任务。请结合工具调用,执行**规划→执行→再规划**的闭环,并最终按照固定模板输出《告警分析报告》。禁止编造虚假数据,如连续多次查询失败需诚实反馈无法完成的原因。";
|
||||
|
||||
logger.info("调用 Supervisor Agent 开始编排...");
|
||||
return supervisorAgent.invoke(taskPrompt);
|
||||
|
||||
Optional<OverAllState> stateOptional = supervisorAgent.invoke(taskPrompt);
|
||||
|
||||
// 添加调试代码
|
||||
if (stateOptional.isPresent()) {
|
||||
OverAllState state = stateOptional.get();
|
||||
logger.debug("Final State Keys: {}", state.data().keySet()); // 打印所有 key
|
||||
logger.debug("Planner Plan: {}", state.value("planner_plan"));
|
||||
logger.debug("Executor Feedback: {}", state.value("executor_feedback"));
|
||||
}
|
||||
|
||||
return stateOptional;
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -97,7 +108,7 @@ public class AiOpsService {
|
||||
/**
|
||||
* 构建 Planner Agent
|
||||
*/
|
||||
private ReactAgent buildPlannerAgent(DashScopeChatModel chatModel, ToolCallback[] toolCallbacks) {
|
||||
private ReactAgent buildPlannerAgent(ChatModel chatModel, ToolCallback[] toolCallbacks) {
|
||||
return ReactAgent.builder()
|
||||
.name("planner_agent")
|
||||
.description("负责拆解告警、规划与再规划步骤")
|
||||
@@ -112,7 +123,7 @@ public class AiOpsService {
|
||||
/**
|
||||
* 构建 Executor Agent
|
||||
*/
|
||||
private ReactAgent buildExecutorAgent(DashScopeChatModel chatModel, ToolCallback[] toolCallbacks) {
|
||||
private ReactAgent buildExecutorAgent(ChatModel chatModel, ToolCallback[] toolCallbacks) {
|
||||
return ReactAgent.builder()
|
||||
.name("executor_agent")
|
||||
.description("负责执行 Planner 的首个步骤并及时反馈")
|
||||
|
||||
@@ -1,8 +1,5 @@
|
||||
package org.example.service;
|
||||
|
||||
import com.alibaba.cloud.ai.dashscope.api.DashScopeApi;
|
||||
import com.alibaba.cloud.ai.dashscope.chat.DashScopeChatModel;
|
||||
import com.alibaba.cloud.ai.dashscope.chat.DashScopeChatOptions;
|
||||
import com.alibaba.cloud.ai.graph.agent.ReactAgent;
|
||||
import com.alibaba.cloud.ai.graph.exception.GraphRunnerException;
|
||||
import org.example.agent.tool.DateTimeTools;
|
||||
@@ -11,10 +8,10 @@ import org.example.agent.tool.QueryLogsTools;
|
||||
import org.example.agent.tool.QueryMetricsTools;
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.tool.ToolCallback;
|
||||
import org.springframework.ai.tool.ToolCallbackProvider;
|
||||
import org.springframework.beans.factory.annotation.Autowired;
|
||||
import org.springframework.beans.factory.annotation.Value;
|
||||
import org.springframework.stereotype.Service;
|
||||
|
||||
import java.util.List;
|
||||
@@ -41,44 +38,17 @@ public class ChatService {
|
||||
@Autowired(required = false) // Mock 模式下才注册,所以设置为 optional,真实环境通过mcp配置注入
|
||||
private QueryLogsTools queryLogsTools;
|
||||
|
||||
@Autowired
|
||||
@Autowired(required = false)
|
||||
private ToolCallbackProvider tools;
|
||||
|
||||
@Value("${spring.ai.dashscope.api-key}")
|
||||
private String dashScopeApiKey;
|
||||
@Autowired
|
||||
private ChatModel chatModel;
|
||||
|
||||
/**
|
||||
* 创建 DashScope API 实例
|
||||
* 获取注入的 ChatModel
|
||||
*/
|
||||
public DashScopeApi createDashScopeApi() {
|
||||
return DashScopeApi.builder()
|
||||
.apiKey(dashScopeApiKey)
|
||||
.build();
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建 ChatModel
|
||||
* @param temperature 控制随机性 (0.0-1.0)
|
||||
* @param maxToken 最大输出长度
|
||||
* @param topP 核采样参数
|
||||
*/
|
||||
public DashScopeChatModel createChatModel(DashScopeApi dashScopeApi, double temperature, int maxToken, double topP) {
|
||||
return DashScopeChatModel.builder()
|
||||
.dashScopeApi(dashScopeApi)
|
||||
.defaultOptions(DashScopeChatOptions.builder()
|
||||
.withModel(DashScopeChatModel.DEFAULT_MODEL_NAME)
|
||||
.withTemperature(temperature)
|
||||
.withMaxToken(maxToken)
|
||||
.withTopP(topP)
|
||||
.build())
|
||||
.build();
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建标准对话 ChatModel(默认参数)
|
||||
*/
|
||||
public DashScopeChatModel createStandardChatModel(DashScopeApi dashScopeApi) {
|
||||
return createChatModel(dashScopeApi, 0.7, 2000, 0.9);
|
||||
public ChatModel getChatModel() {
|
||||
return chatModel;
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -88,20 +58,29 @@ public class ChatService {
|
||||
*/
|
||||
public String buildSystemPrompt(List<Map<String, String>> history) {
|
||||
StringBuilder systemPromptBuilder = new StringBuilder();
|
||||
|
||||
|
||||
// 基础系统提示
|
||||
systemPromptBuilder.append("你是一个专业的智能助手,可以获取当前时间、查询天气信息、搜索内部文档知识库,以及查询 Prometheus 告警信息。\n");
|
||||
systemPromptBuilder.append("当用户询问时间相关问题时,使用 getCurrentDateTime 工具。\n");
|
||||
systemPromptBuilder.append("当用户询问时间相关问题时,**必须每次都调用 getCurrentDateTime 工具**,因为时间会不断变化。即使历史消息中有时间信息,也不要直接复用,必须重新查询最新时间。\n");
|
||||
systemPromptBuilder.append("当用户需要查询公司内部文档、流程、最佳实践或技术指南时,使用 queryInternalDocs 工具。\n");
|
||||
systemPromptBuilder.append("当用户需要查询 Prometheus 告警、监控指标或系统告警状态时,使用 queryPrometheusAlerts 工具。\n");
|
||||
systemPromptBuilder.append("当用户需要查询腾讯云日志时,请调用腾讯云mcp服务查询,默认查询地域ap-guangzhou,查询时间范围为近一个月。\n\n");
|
||||
|
||||
// 添加历史消息
|
||||
|
||||
// 添加历史消息(过滤时间查询相关内容)
|
||||
if (!history.isEmpty()) {
|
||||
systemPromptBuilder.append("--- 对话历史 ---\n");
|
||||
for (Map<String, String> msg : history) {
|
||||
String role = msg.get("role");
|
||||
String content = msg.get("content");
|
||||
|
||||
// 🔧 过滤时间查询相关的历史消息,避免 LLM 复用旧的时间信息
|
||||
if ("user".equals(role) && isTimeQuery(content)) {
|
||||
continue; // 跳过时间查询问题
|
||||
}
|
||||
if ("assistant".equals(role) && containsTimeInfo(content)) {
|
||||
continue; // 跳过包含时间信息的回答
|
||||
}
|
||||
|
||||
if ("user".equals(role)) {
|
||||
systemPromptBuilder.append("用户: ").append(content).append("\n");
|
||||
} else if ("assistant".equals(role)) {
|
||||
@@ -110,12 +89,36 @@ public class ChatService {
|
||||
}
|
||||
systemPromptBuilder.append("--- 对话历史结束 ---\n\n");
|
||||
}
|
||||
|
||||
|
||||
systemPromptBuilder.append("请基于以上对话历史,回答用户的新问题。");
|
||||
|
||||
|
||||
return systemPromptBuilder.toString();
|
||||
}
|
||||
|
||||
/**
|
||||
* 判断是否为时间查询问题
|
||||
*/
|
||||
private boolean isTimeQuery(String content) {
|
||||
if (content == null) {
|
||||
return false;
|
||||
}
|
||||
// 匹配常见的时间查询模式
|
||||
return content.matches(".*(现在|当前|此时).*(几点|时间).*") ||
|
||||
content.matches(".*(几点|时间).*(了|呢|[??]).*") ||
|
||||
content.toLowerCase().matches(".*(what.*time|current.*time).*");
|
||||
}
|
||||
|
||||
/**
|
||||
* 判断是否包含时间信息
|
||||
*/
|
||||
private boolean containsTimeInfo(String content) {
|
||||
if (content == null) {
|
||||
return false;
|
||||
}
|
||||
// 匹配日期时间格式:2026年5月31日、15:57、下午3点 等
|
||||
return content.matches(".*(\\d{4}年\\d{1,2}月\\d{1,2}日|\\d{1,2}:\\d{2}|[上下午]+\\d{1,2}[点时]).*");
|
||||
}
|
||||
|
||||
/**
|
||||
* 动态构建方法工具数组
|
||||
* 根据 cls.mock-enabled 决定是否包含 QueryLogsTools
|
||||
@@ -134,6 +137,9 @@ public class ChatService {
|
||||
* 获取工具回调列表,mcp服务提供的工具
|
||||
*/
|
||||
public ToolCallback[] getToolCallbacks() {
|
||||
if (tools == null) {
|
||||
return new ToolCallback[0];
|
||||
}
|
||||
return tools.getToolCallbacks();
|
||||
}
|
||||
|
||||
@@ -141,6 +147,10 @@ public class ChatService {
|
||||
* 记录可用工具列表:mcp服务提供的工具
|
||||
*/
|
||||
public void logAvailableTools() {
|
||||
if (tools == null) {
|
||||
logger.info("MCP 未启用,无远程工具");
|
||||
return;
|
||||
}
|
||||
ToolCallback[] toolCallbacks = tools.getToolCallbacks();
|
||||
logger.info("可用工具列表:");
|
||||
for (ToolCallback toolCallback : toolCallbacks) {
|
||||
@@ -154,7 +164,7 @@ public class ChatService {
|
||||
* @param systemPrompt 系统提示词
|
||||
* @return 配置好的 ReactAgent
|
||||
*/
|
||||
public ReactAgent createReactAgent(DashScopeChatModel chatModel, String systemPrompt) {
|
||||
public ReactAgent createReactAgent(ChatModel chatModel, String systemPrompt) {
|
||||
return ReactAgent.builder()
|
||||
.name("intelligent_assistant")
|
||||
.model(chatModel)
|
||||
|
||||
@@ -1,22 +1,18 @@
|
||||
package org.example.service;
|
||||
|
||||
import com.alibaba.dashscope.aigc.generation.Generation;
|
||||
import com.alibaba.dashscope.aigc.generation.GenerationParam;
|
||||
import com.alibaba.dashscope.aigc.generation.GenerationResult;
|
||||
import com.alibaba.dashscope.common.Message;
|
||||
import com.alibaba.dashscope.common.Role;
|
||||
import com.alibaba.dashscope.exception.ApiException;
|
||||
import com.alibaba.dashscope.exception.InputRequiredException;
|
||||
import com.alibaba.dashscope.exception.NoApiKeyException;
|
||||
import com.alibaba.dashscope.utils.Constants;
|
||||
import io.reactivex.Flowable;
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
import org.springframework.ai.chat.messages.AssistantMessage;
|
||||
import org.springframework.ai.chat.messages.Message;
|
||||
import org.springframework.ai.chat.messages.UserMessage;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.chat.model.ChatResponse;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.beans.factory.annotation.Autowired;
|
||||
import org.springframework.beans.factory.annotation.Value;
|
||||
import org.springframework.stereotype.Service;
|
||||
import reactor.core.publisher.Flux;
|
||||
|
||||
import jakarta.annotation.PostConstruct;
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
@@ -33,29 +29,12 @@ public class RagService {
|
||||
@Autowired
|
||||
private VectorSearchService vectorSearchService;
|
||||
|
||||
@Value("${dashscope.api.key}")
|
||||
private String apiKey;
|
||||
@Autowired
|
||||
private ChatModel chatModel;
|
||||
|
||||
@Value("${rag.top-k:3}")
|
||||
private int topK;
|
||||
|
||||
@Value("${rag.model:qwen3-30b-a3b-thinking-2507}")
|
||||
private String model;
|
||||
|
||||
private Generation generation;
|
||||
|
||||
@PostConstruct
|
||||
public void init() {
|
||||
// 设置 API Key 和 Base URL
|
||||
Constants.apiKey = apiKey;
|
||||
Constants.baseHttpApiUrl = "https://dashscope.aliyuncs.com/api/v1";
|
||||
|
||||
// 创建 Generation 实例
|
||||
generation = new Generation();
|
||||
|
||||
logger.info("RAG 服务初始化完成,model: {}, topK: {}", model, topK);
|
||||
}
|
||||
|
||||
/**
|
||||
* 流式处理用户问题(不带历史消息)
|
||||
*
|
||||
@@ -138,86 +117,64 @@ public class RagService {
|
||||
* @param history 历史消息列表
|
||||
* @param callback 流式回调接口
|
||||
*/
|
||||
private void generateAnswerStream(String prompt, List<Map<String, String>> history, StreamCallback callback)
|
||||
throws NoApiKeyException, ApiException, InputRequiredException {
|
||||
|
||||
private void generateAnswerStream(String prompt, List<Map<String, String>> history, StreamCallback callback) {
|
||||
// 构建消息列表:历史消息 + 当前问题
|
||||
List<Message> messages = new ArrayList<>();
|
||||
|
||||
|
||||
// 添加历史消息
|
||||
for (Map<String, String> historyMsg : history) {
|
||||
String role = historyMsg.get("role");
|
||||
String content = historyMsg.get("content");
|
||||
|
||||
|
||||
if ("user".equals(role)) {
|
||||
messages.add(Message.builder()
|
||||
.role(Role.USER.getValue())
|
||||
.content(content)
|
||||
.build());
|
||||
messages.add(new UserMessage(content));
|
||||
} else if ("assistant".equals(role)) {
|
||||
messages.add(Message.builder()
|
||||
.role(Role.ASSISTANT.getValue())
|
||||
.content(content)
|
||||
.build());
|
||||
messages.add(new AssistantMessage(content));
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// 添加当前用户问题
|
||||
Message userMsg = Message.builder()
|
||||
.role(Role.USER.getValue())
|
||||
.content(prompt)
|
||||
.build();
|
||||
messages.add(userMsg);
|
||||
|
||||
logger.debug("发送给AI模型的消息数量: {}(包含 {} 条历史消息)",
|
||||
messages.add(new UserMessage(prompt));
|
||||
|
||||
logger.debug("发送给AI模型的消息数量: {}(包含 {} 条历史消息)",
|
||||
messages.size(), history.size());
|
||||
|
||||
GenerationParam param = GenerationParam.builder()
|
||||
.apiKey(apiKey)
|
||||
.model(model)
|
||||
.incrementalOutput(true)
|
||||
.resultFormat("message")
|
||||
.messages(messages)
|
||||
.build();
|
||||
|
||||
logger.info("开始调用AI模型流式接口...");
|
||||
|
||||
Flowable<GenerationResult> result = generation.streamCall(param);
|
||||
|
||||
|
||||
StringBuilder reasoningContent = new StringBuilder();
|
||||
StringBuilder finalContent = new StringBuilder();
|
||||
|
||||
|
||||
Flux<ChatResponse> flux = chatModel.stream(new Prompt(messages));
|
||||
|
||||
logger.info("开始接收AI模型流式响应...");
|
||||
|
||||
result.blockingForEach(message -> {
|
||||
if (message.getOutput() != null &&
|
||||
message.getOutput().getChoices() != null &&
|
||||
!message.getOutput().getChoices().isEmpty()) {
|
||||
|
||||
// 获取消息内容
|
||||
// 注意:qwen3-30b-a3b-thinking-2507 模型会在 content 中返回完整内容
|
||||
// reasoning 部分可能需要通过特殊方式提取或者直接包含在 content 中
|
||||
String content = message.getOutput().getChoices().get(0).getMessage().getContent();
|
||||
flux.subscribe(
|
||||
response -> {
|
||||
if (response.getResults() != null && !response.getResults().isEmpty()) {
|
||||
String content = response.getResults().get(0).getOutput().getText();
|
||||
|
||||
if (content != null && !content.isEmpty()) {
|
||||
logger.debug("收到AI模型内容块: {}", content);
|
||||
|
||||
// 对于 thinking 模型,content 可能包含思考过程和最终答案
|
||||
// 这里我们将所有内容都作为答案返回
|
||||
finalContent.append(content);
|
||||
callback.onContentChunk(content);
|
||||
|
||||
logger.debug("已调用 onContentChunk 回调");
|
||||
} else {
|
||||
logger.debug("收到空内容块,跳过");
|
||||
if (content != null && !content.isEmpty()) {
|
||||
logger.debug("收到AI模型内容块: {}", content);
|
||||
|
||||
finalContent.append(content);
|
||||
callback.onContentChunk(content);
|
||||
|
||||
logger.debug("已调用 onContentChunk 回调");
|
||||
} else {
|
||||
logger.debug("收到空内容块,跳过");
|
||||
}
|
||||
}
|
||||
},
|
||||
error -> {
|
||||
logger.error("AI模型流式响应失败", error);
|
||||
callback.onError(new Exception("AI模型流式响应失败: " + error.getMessage(), error));
|
||||
},
|
||||
() -> {
|
||||
logger.info("AI模型流式响应完成,总内容长度: {}", finalContent.length());
|
||||
callback.onComplete(finalContent.toString(), reasoningContent.toString());
|
||||
logger.info("已调用 onComplete 回调");
|
||||
}
|
||||
});
|
||||
|
||||
logger.info("AI模型流式响应完成,总内容长度: {}", finalContent.length());
|
||||
|
||||
callback.onComplete(finalContent.toString(), reasoningContent.toString());
|
||||
logger.info("已调用 onComplete 回调");
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
@@ -1,19 +1,11 @@
|
||||
package org.example.service;
|
||||
|
||||
import com.alibaba.dashscope.embeddings.TextEmbedding;
|
||||
import com.alibaba.dashscope.embeddings.TextEmbeddingParam;
|
||||
import com.alibaba.dashscope.embeddings.TextEmbeddingResult;
|
||||
import com.alibaba.dashscope.embeddings.TextEmbeddingOutput;
|
||||
import com.alibaba.dashscope.embeddings.TextEmbeddingResultItem;
|
||||
import com.alibaba.dashscope.exception.NoApiKeyException;
|
||||
import com.alibaba.dashscope.utils.Constants;
|
||||
import org.jetbrains.annotations.NotNull;
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
import org.springframework.beans.factory.annotation.Value;
|
||||
import org.springframework.ai.embedding.EmbeddingModel;
|
||||
import org.springframework.beans.factory.annotation.Autowired;
|
||||
import org.springframework.stereotype.Service;
|
||||
|
||||
import jakarta.annotation.PostConstruct;
|
||||
import java.util.ArrayList;
|
||||
import java.util.Collections;
|
||||
import java.util.List;
|
||||
@@ -27,44 +19,8 @@ public class VectorEmbeddingService {
|
||||
|
||||
private static final Logger logger = LoggerFactory.getLogger(VectorEmbeddingService.class);
|
||||
|
||||
@Value("${dashscope.api.key}")
|
||||
private String apiKey;
|
||||
|
||||
@Value("${dashscope.embedding.model}")
|
||||
private String model;
|
||||
|
||||
private TextEmbedding textEmbedding;
|
||||
|
||||
@PostConstruct
|
||||
public void init() {
|
||||
// 验证 API Key
|
||||
if (apiKey == null || apiKey.trim().isEmpty() || apiKey.equals("your-api-key-here")) {
|
||||
logger.error("API Key 未正确配置!当前值: {}", apiKey);
|
||||
throw new IllegalStateException("请设置环境变量 DASHSCOPE_API_KEY 或在 application.yml 中配置正确的 API Key");
|
||||
}
|
||||
|
||||
// 打印 API Key 前缀用于调试(不打印完整 Key 保证安全)
|
||||
String maskedKey = apiKey.length() > 8 ?
|
||||
apiKey.substring(0, 8) + "..." + apiKey.substring(apiKey.length() - 4) :
|
||||
"***";
|
||||
logger.info("API Key 已加载: {}", maskedKey);
|
||||
|
||||
// 设置全局 API Key(确保设置成功)
|
||||
Constants.apiKey = apiKey;
|
||||
|
||||
// 验证 API Key 是否设置成功
|
||||
if (Constants.apiKey == null || Constants.apiKey.isEmpty()) {
|
||||
logger.error("Constants.apiKey 设置失败!");
|
||||
throw new IllegalStateException("API Key 设置到 Constants 失败");
|
||||
}
|
||||
|
||||
logger.info("Constants.apiKey 已设置: {}", Constants.apiKey.substring(0, Math.min(8, Constants.apiKey.length())) + "...");
|
||||
|
||||
// 创建 TextEmbedding 实例
|
||||
textEmbedding = new TextEmbedding();
|
||||
|
||||
logger.info("阿里云 DashScope Embedding 服务初始化完成,模型: {}", model);
|
||||
}
|
||||
@Autowired
|
||||
private EmbeddingModel embeddingModel;
|
||||
|
||||
/**
|
||||
* 生成向量嵌入
|
||||
@@ -81,73 +37,25 @@ public class VectorEmbeddingService {
|
||||
}
|
||||
|
||||
logger.debug("开始生成向量嵌入, 内容长度: {} 字符", content.length());
|
||||
|
||||
// 确保 API Key 已设置(防止被其他地方覆盖)
|
||||
if (Constants.apiKey == null || Constants.apiKey.isEmpty()) {
|
||||
logger.warn("检测到 Constants.apiKey 为空,重新设置");
|
||||
Constants.apiKey = apiKey;
|
||||
|
||||
float[] embedding = embeddingModel.embed(content);
|
||||
|
||||
List<Float> floatEmbedding = new ArrayList<>(embedding.length);
|
||||
for (float v : embedding) {
|
||||
floatEmbedding.add(v);
|
||||
}
|
||||
|
||||
logger.debug("调用 API 前 Constants.apiKey: {}",
|
||||
Constants.apiKey != null ? Constants.apiKey.substring(0, Math.min(8, Constants.apiKey.length())) + "..." : "null");
|
||||
|
||||
// 构建请求参数
|
||||
TextEmbeddingParam param = TextEmbeddingParam
|
||||
.builder()
|
||||
.model(model)
|
||||
.texts(Collections.singletonList(content))
|
||||
.build();
|
||||
|
||||
// 调用 API
|
||||
TextEmbeddingResult result = textEmbedding.call(param);
|
||||
|
||||
// 检查结果
|
||||
List<Float> floatEmbedding = getFloats(result);
|
||||
|
||||
logger.info("成功生成向量嵌入, 内容长度: {} 字符, 向量维度: {}",
|
||||
logger.info("成功生成向量嵌入, 内容长度: {} 字符, 向量维度: {}",
|
||||
content.length(), floatEmbedding.size());
|
||||
|
||||
return floatEmbedding;
|
||||
|
||||
} catch (NoApiKeyException e) {
|
||||
logger.error("API Key 未设置或无效", e);
|
||||
throw new RuntimeException("API Key 未设置,请配置 dashscope.api.key", e);
|
||||
} catch (Exception e) {
|
||||
logger.error("生成向量嵌入失败, 内容长度: {}", content != null ? content.length() : 0, e);
|
||||
throw new RuntimeException("生成向量嵌入失败: " + e.getMessage(), e);
|
||||
}
|
||||
}
|
||||
|
||||
@NotNull
|
||||
private static List<Float> getFloats(TextEmbeddingResult result) {
|
||||
if (result == null || result.getOutput() == null || result.getOutput().getEmbeddings() == null) {
|
||||
throw new RuntimeException("DashScope API 返回空结果");
|
||||
}
|
||||
|
||||
TextEmbeddingOutput output = result.getOutput();
|
||||
List<TextEmbeddingResultItem> embeddings = output.getEmbeddings();
|
||||
|
||||
if (embeddings.isEmpty()) {
|
||||
throw new RuntimeException("DashScope API 返回空向量列表");
|
||||
}
|
||||
|
||||
// 获取第一个文本的向量
|
||||
List<Double> embeddingDoubles = embeddings.get(0).getEmbedding();
|
||||
|
||||
// 转换为 List<Float>
|
||||
List<Float> floatEmbedding = new ArrayList<>(embeddingDoubles.size());
|
||||
for (Double value : embeddingDoubles) {
|
||||
floatEmbedding.add(value.floatValue());
|
||||
}
|
||||
return floatEmbedding;
|
||||
}
|
||||
|
||||
/**
|
||||
* 批量生成向量嵌入
|
||||
*
|
||||
* @param contents 文本内容列表
|
||||
* @return 向量嵌入列表
|
||||
*/
|
||||
public List<List<Float>> generateEmbeddings(List<String> contents) {
|
||||
try {
|
||||
if (contents == null || contents.isEmpty()) {
|
||||
@@ -156,54 +64,24 @@ public class VectorEmbeddingService {
|
||||
}
|
||||
|
||||
logger.info("开始批量生成向量嵌入, 数量: {}", contents.size());
|
||||
|
||||
// 确保 API Key 已设置
|
||||
if (Constants.apiKey == null || Constants.apiKey.isEmpty()) {
|
||||
logger.warn("检测到 Constants.apiKey 为空,重新设置");
|
||||
Constants.apiKey = apiKey;
|
||||
}
|
||||
|
||||
// 构建请求参数 - 批量输入
|
||||
TextEmbeddingParam param = TextEmbeddingParam
|
||||
.builder()
|
||||
.model(model)
|
||||
.texts(contents)
|
||||
.build();
|
||||
List<float[]> embeddings = embeddingModel.embed(contents);
|
||||
|
||||
// 调用 API
|
||||
TextEmbeddingResult result = textEmbedding.call(param);
|
||||
|
||||
// 检查结果
|
||||
if (result == null || result.getOutput() == null || result.getOutput().getEmbeddings() == null) {
|
||||
throw new RuntimeException("批量 DashScope API 返回空结果");
|
||||
}
|
||||
|
||||
List<TextEmbeddingResultItem> embeddingItems = result.getOutput().getEmbeddings();
|
||||
|
||||
if (embeddingItems.isEmpty()) {
|
||||
throw new RuntimeException("批量 DashScope API 返回空向量列表");
|
||||
}
|
||||
|
||||
// 转换结果
|
||||
List<List<Float>> embeddings = new ArrayList<>();
|
||||
for (TextEmbeddingResultItem item : embeddingItems) {
|
||||
List<Double> embeddingDoubles = item.getEmbedding();
|
||||
List<Float> embedding = new ArrayList<>(embeddingDoubles.size());
|
||||
for (Double value : embeddingDoubles) {
|
||||
embedding.add(value.floatValue());
|
||||
List<List<Float>> result = new ArrayList<>();
|
||||
for (float[] embedding : embeddings) {
|
||||
List<Float> floatEmbedding = new ArrayList<>(embedding.length);
|
||||
for (float v : embedding) {
|
||||
floatEmbedding.add(v);
|
||||
}
|
||||
embeddings.add(embedding);
|
||||
result.add(floatEmbedding);
|
||||
}
|
||||
|
||||
logger.info("成功批量生成向量嵌入, 数量: {}, 维度: {}",
|
||||
embeddings.size(),
|
||||
embeddings.isEmpty() ? 0 : embeddings.get(0).size());
|
||||
logger.info("成功批量生成向量嵌入, 数量: {}, 维度: {}",
|
||||
result.size(),
|
||||
result.isEmpty() ? 0 : result.get(0).size());
|
||||
|
||||
return embeddings;
|
||||
return result;
|
||||
|
||||
} catch (NoApiKeyException e) {
|
||||
logger.error("批量调用时 API Key 未设置或无效", e);
|
||||
throw new RuntimeException("API Key 未设置,请配置 dashscope.api.key", e);
|
||||
} catch (Exception e) {
|
||||
logger.error("批量生成向量嵌入失败", e);
|
||||
throw new RuntimeException("批量生成向量嵌入失败: " + e.getMessage(), e);
|
||||
|
||||
Reference in New Issue
Block a user