This commit is contained in:
aruo
2026-05-31 21:45:14 +08:00
parent d4b5015beb
commit ac08345369
67 changed files with 11120 additions and 387 deletions
@@ -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);