This commit is contained in:
zhuyongxin
2026-04-30 11:19:46 +08:00
parent 0671ad8bab
commit 80eada415d
42 changed files with 9135 additions and 0 deletions
+11
View File
@@ -0,0 +1,11 @@
package org.example;
import org.springframework.boot.SpringApplication;
import org.springframework.boot.autoconfigure.SpringBootApplication;
@SpringBootApplication
public class Main {
public static void main(String[] args) {
SpringApplication.run(Main.class, args);
}
}
@@ -0,0 +1,19 @@
package org.example.agent.tool;
import org.springframework.ai.tool.annotation.Tool;
import org.springframework.context.i18n.LocaleContextHolder;
import org.springframework.stereotype.Component;
import java.time.LocalDateTime;
@Component
public class DateTimeTools {
/** 工具名常量,用于动态构建提示词 */
public static final String TOOL_GET_CURRENT_DATETIME = "getCurrentDateTime";
@Tool(description = "Get the current date and time in the user's timezone")
public String getCurrentDateTime() {
return LocalDateTime.now().atZone(LocaleContextHolder.getTimeZone().toZoneId()).toString();
}
}
@@ -0,0 +1,79 @@
package org.example.agent.tool;
import com.fasterxml.jackson.databind.ObjectMapper;
import org.example.service.VectorSearchService;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.tool.annotation.Tool;
import org.springframework.ai.tool.annotation.ToolParam;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.stereotype.Component;
import java.util.List;
/**
* 内部文档查询工具
* 使用 RAG (Retrieval-Augmented Generation) 从内部知识库检索相关文档
*/
@Component
public class InternalDocsTools {
private static final Logger logger = LoggerFactory.getLogger(InternalDocsTools.class);
/** 工具名常量,用于动态构建提示词 */
public static final String TOOL_QUERY_INTERNAL_DOCS = "queryInternalDocs";
private final VectorSearchService vectorSearchService;
@Value("${rag.top-k:3}")
private int topK = 3; // 默认值
private final ObjectMapper objectMapper = new ObjectMapper();
/**
* 构造函数注入依赖
* Spring 会自动注入 VectorSearchService
*/
@Autowired
public InternalDocsTools(VectorSearchService vectorSearchService) {
this.vectorSearchService = vectorSearchService;
}
/**
* 查询内部文档工具
*
* @param query 搜索查询,描述您要查找的信息
* @return JSON 格式的搜索结果,包含相关文档内容、相似度分数和元数据
*/
@Tool(description = "Use this tool to search internal documentation and knowledge base for relevant information. " +
"It performs RAG (Retrieval-Augmented Generation) to find similar documents and extract processing steps. " +
"This is useful when you need to understand internal procedures, best practices, or step-by-step guides " +
"stored in the company's documentation.")
public String queryInternalDocs(
@ToolParam(description = "Search query describing what information you are looking for")
String query) {
try {
// 使用向量搜索服务检索相关文档
List<VectorSearchService.SearchResult> searchResults =
vectorSearchService.searchSimilarDocuments(query, topK);
if (searchResults.isEmpty()) {
return "{\"status\": \"no_results\", \"message\": \"No relevant documents found in the knowledge base.\"}";
}
// 将搜索结果转换为 JSON 格式
String resultJson = objectMapper.writeValueAsString(searchResults);
return resultJson;
} catch (Exception e) {
logger.error("[工具错误] queryInternalDocs 执行失败", e);
return String.format("{\"status\": \"error\", \"message\": \"Failed to query internal docs: %s\"}",
e.getMessage());
}
}
}
@@ -0,0 +1,694 @@
package org.example.agent.tool;
import com.fasterxml.jackson.annotation.JsonProperty;
import com.fasterxml.jackson.databind.ObjectMapper;
import lombok.Data;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.tool.annotation.Tool;
import org.springframework.ai.tool.annotation.ToolParam;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.stereotype.Component;
import java.time.Instant;
import java.time.ZoneId;
import java.time.format.DateTimeFormatter;
import java.time.temporal.ChronoUnit;
import java.util.ArrayList;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
/**
* 日志查询工具
* 用于查询 CLS(云日志服务)的日志信息
* 支持 Mock 模式,提供与告警关联的模拟日志数据
*/
@Component
public class QueryLogsTools {
private static final Logger logger = LoggerFactory.getLogger(QueryLogsTools.class);
/** 工具名常量,用于动态构建提示词 */
public static final String TOOL_QUERY_LOGS = "queryLogs";
public static final String TOOL_GET_AVAILABLE_LOG_TOPICS = "getAvailableLogTopics";
private final ObjectMapper objectMapper = new ObjectMapper();
@Value("${cls.mock-enabled:false}")
private boolean mockEnabled;
private static final DateTimeFormatter FORMATTER = DateTimeFormatter
.ofPattern("yyyy-MM-dd HH:mm:ss")
.withZone(ZoneId.of("Asia/Shanghai"));
@jakarta.annotation.PostConstruct
public void init() {
logger.info("✅ QueryLogsTools 初始化成功, Mock模式: {}", mockEnabled);
}
/**
* 获取可用的日志主题列表
* 用于查询前先了解有哪些日志主题可供查询
*/
@Tool(description = "Get all available log topics and their descriptions. " +
"Call this tool first before querying logs to understand what log topics are available. " +
"Returns a list of log topics with their names, descriptions, and example queries.")
public String getAvailableLogTopics() {
logger.info("获取可用的日志主题列表");
try {
List<LogTopicInfo> topics = new ArrayList<>();
// 系统指标日志
LogTopicInfo systemMetrics = new LogTopicInfo();
systemMetrics.setTopicName("system-metrics");
systemMetrics.setDescription("系统指标日志,包含 CPU、内存、磁盘使用率等系统资源监控数据");
systemMetrics.setExampleQueries(List.of(
"cpu_usage:>80",
"memory_usage:>85",
"disk_usage:>90",
"level:WARN AND service:payment-service"
));
systemMetrics.setRelatedAlerts(List.of("HighCPUUsage", "HighMemoryUsage", "HighDiskUsage"));
topics.add(systemMetrics);
// 应用日志
LogTopicInfo applicationLogs = new LogTopicInfo();
applicationLogs.setTopicName("application-logs");
applicationLogs.setDescription("应用日志,包含应用程序的错误日志、警告日志、慢请求日志、下游依赖调用日志等");
applicationLogs.setExampleQueries(List.of(
"level:ERROR",
"level:FATAL",
"http_status:500",
"response_time:>3000",
"slow",
"downstream OR redis OR database OR mq"
));
applicationLogs.setRelatedAlerts(List.of("ServiceUnavailable", "SlowResponse", "HighMemoryUsage"));
topics.add(applicationLogs);
// 数据库慢查询日志
LogTopicInfo dbSlowQuery = new LogTopicInfo();
dbSlowQuery.setTopicName("database-slow-query");
dbSlowQuery.setDescription("数据库慢查询日志,包含执行时间较长的 SQL 查询,可用于分析数据库性能问题");
dbSlowQuery.setExampleQueries(List.of(
"query_time:>2",
"table:orders",
"query_type:SELECT",
"*" // 查询所有慢查询
));
dbSlowQuery.setRelatedAlerts(List.of("SlowResponse", "ServiceUnavailable"));
topics.add(dbSlowQuery);
// 系统事件日志
LogTopicInfo systemEvents = new LogTopicInfo();
systemEvents.setTopicName("system-events");
systemEvents.setDescription("系统事件日志,包含 Kubernetes Pod 重启、OOM Kill、容器崩溃等系统级事件");
systemEvents.setExampleQueries(List.of(
"restart OR crash",
"oom_kill",
"event_type:PodRestart",
"reason:OOMKilled"
));
systemEvents.setRelatedAlerts(List.of("ServiceUnavailable", "HighMemoryUsage"));
topics.add(systemEvents);
// 构建输出
LogTopicsOutput output = new LogTopicsOutput();
output.setSuccess(true);
output.setTopics(topics);
output.setAvailableRegions(List.of("ap-guangzhou", "ap-shanghai", "ap-beijing", "ap-chengdu"));
output.setDefaultRegion("ap-guangzhou");
output.setMessage(String.format("共有 %d 个可用的日志主题。建议使用默认地域 'ap-guangzhou' 或省略 region 参数", topics.size()));
return objectMapper.writerWithDefaultPrettyPrinter().writeValueAsString(output);
} catch (Exception e) {
logger.error("获取日志主题列表失败", e);
return "{\"success\":false,\"message\":\"获取日志主题列表失败: " + e.getMessage() + "\"}";
}
}
/**
* 查询日志
* 从云日志服务查询指定条件的日志
*
* @param region 地域,如 ap-guangzhou
* @param logTopic 日志主题,如 system-metrics, application-logs
* @param query 查询条件,如 level:ERROR OR cpu_usage:>80
* @param limit 返回的日志条数,默认20条
*/
// 有效地域列表
private static final List<String> VALID_REGIONS = List.of(
"ap-guangzhou", "ap-shanghai", "ap-beijing", "ap-chengdu"
);
private static final String DEFAULT_REGION = "ap-guangzhou";
@Tool(description = "Query logs from Cloud Log Service (CLS). " +
"Use this tool to search application logs, system metrics, and other log data. " +
"IMPORTANT: Before calling this tool, you should call getAvailableLogTopics to understand what log topics are available. " +
"Available log topics: " +
"1) 'system-metrics' - System metrics logs (CPU, memory, disk usage, etc. Related to HighCPUUsage, HighMemoryUsage, HighDiskUsage alerts); " +
"2) 'application-logs' - Application logs (error logs, slow request logs, downstream dependency logs. Related to ServiceUnavailable, SlowResponse alerts); " +
"3) 'database-slow-query' - Database slow query logs (SQL queries with long execution time. Related to SlowResponse alerts); " +
"4) 'system-events' - System event logs (Pod restart, OOM Kill, container crash. Related to ServiceUnavailable, HighMemoryUsage alerts). " +
"logTopic (required, one of the above topics or their CLS topicId), " +
"query (optional, defaults to a curated search if empty), " +
"limit (optional, default 20, max 100).")
public String queryLogs(
@ToolParam(description = "地域,可选值: ap-guangzhou, ap-shanghai, ap-beijing, ap-chengdu。默认 ap-guangzhou") String region,
@ToolParam(description = "日志主题,如 system-metrics, application-logs, database-slow-query, system-events,也支持 CLS TopicId") String logTopic,
@ToolParam(description = "查询条件,支持 Lucene 语法,如 level:ERROR OR cpu_usage:>80;为空时返回该主题近 5 条核心日志") String query,
@ToolParam(description = "返回日志条数,默认20,最大100") Integer limit) {
int actualLimit = (limit == null || limit <= 0) ? 20 : Math.min(limit, 100);
String safeQuery = query == null ? "" : query;
try {
List<LogEntry> logEntries;
if (mockEnabled) {
// Mock 模式:返回与告警关联的模拟日志数据
logEntries = buildMockLogs(region, logTopic, safeQuery, actualLimit);
logger.info("使用 Mock 数据,返回 {} 条日志", logEntries.size());
} else {
// 真实模式:调用 CLS API(这里预留接口,后续实现)
return buildErrorResponse("CLS 真实查询尚未实现,请启用 mock 模式进行测试");
}
// 构建成功响应
QueryLogsOutput output = new QueryLogsOutput();
output.setSuccess(!logEntries.isEmpty());
output.setRegion(region);
output.setLogTopic(logTopic);
output.setQuery(safeQuery.isBlank() ? "DEFAULT_QUERY" : safeQuery);
output.setLogs(logEntries);
output.setTotal(logEntries.size());
output.setMessage(logEntries.isEmpty() ? "未找到匹配的日志" : String.format("成功查询到 %d 条日志", logEntries.size()));
String jsonResult = objectMapper.writerWithDefaultPrettyPrinter().writeValueAsString(output);
logger.info("日志查询完成: 找到 {} 条日志", logEntries.size());
return jsonResult;
} catch (Exception e) {
logger.error("查询日志失败", e);
return buildErrorResponse("查询失败: " + e.getMessage());
}
}
/**
* 构建 Mock 日志数据
* 根据日志主题和查询条件返回与告警关联的模拟数据
*/
private List<LogEntry> buildMockLogs(String region, String logTopic, String query, int limit) {
List<LogEntry> logs = new ArrayList<>();
Instant now = Instant.now();
String safeTopic = logTopic == null ? "system-metrics" : logTopic.toLowerCase();
String normalizedQuery = query == null ? "" : query.toLowerCase();
// 根据日志主题和查询条件生成对应的 mock 数据
switch (safeTopic) {
case "system-metrics":
logs.addAll(buildSystemMetricsLogs(now, normalizedQuery, limit));
break;
case "application-logs":
logs.addAll(buildApplicationLogs(now, normalizedQuery, limit));
break;
case "database-slow-query":
logs.addAll(buildDatabaseSlowQueryLogs(now, normalizedQuery, limit));
break;
case "system-events":
logs.addAll(buildSystemEventsLogs(now, normalizedQuery, limit));
break;
default:
logs.addAll(buildGenericLogs(now, normalizedQuery, limit));
}
if (logs.isEmpty()) {
logs.addAll(buildGenericLogs(now, normalizedQuery, limit));
}
// 限制返回条数
if (logs.size() > limit) {
logs = logs.subList(0, limit);
}
return logs;
}
/**
* 构建系统指标日志(与 CPU、内存、磁盘告警关联)
*/
private List<LogEntry> buildSystemMetricsLogs(Instant now, String query, int limit) {
List<LogEntry> logs = new ArrayList<>();
// CPU 相关日志
if (query.contains("cpu") || query.contains(">80")) {
for (int i = 0; i < 5; i++) {
LogEntry log = new LogEntry();
log.setTimestamp(FORMATTER.format(now.minus(i * 2, ChronoUnit.MINUTES)));
log.setLevel("WARN");
log.setService("payment-service");
log.setInstance("pod-payment-service-7d8f9c6b5-x2k4m");
log.setMessage(String.format("CPU使用率过高: %.1f%%, 进程: java (PID: 1), 线程数: 245", 92.0 - i * 1.5));
log.setMetrics(Map.of(
"cpu_usage", String.format("%.1f", 92.0 - i * 1.5),
"cpu_cores", "4",
"load_average_1m", "3.82",
"load_average_5m", "3.65",
"top_process", "java",
"process_threads", "245"
));
logs.add(log);
}
}
// 内存相关日志
if (query.contains("memory") || query.contains(">85") || query.contains("oom")) {
for (int i = 0; i < 5; i++) {
LogEntry log = new LogEntry();
log.setTimestamp(FORMATTER.format(now.minus(i * 3, ChronoUnit.MINUTES)));
log.setLevel("WARN");
log.setService("order-service");
log.setInstance("pod-order-service-5c7d8e9f1-m3n2p");
log.setMessage(String.format("内存使用率过高: %.1f%%, JVM堆内存: %.1fGB/4GB, GC次数: %d",
91.0 - i * 1.2, 3.8 - i * 0.1, 128 - i * 5));
log.setMetrics(Map.of(
"memory_usage", String.format("%.1f", 91.0 - i * 1.2),
"jvm_heap_used", String.format("%.1fGB", 3.8 - i * 0.1),
"jvm_heap_max", "4GB",
"gc_count", String.valueOf(128 - i * 5),
"gc_time_ms", String.valueOf(1250 + i * 50)
));
logs.add(log);
}
// 添加 GC 警告日志
LogEntry gcLog = new LogEntry();
gcLog.setTimestamp(FORMATTER.format(now.minus(8, ChronoUnit.MINUTES)));
gcLog.setLevel("WARN");
gcLog.setService("order-service");
gcLog.setInstance("pod-order-service-5c7d8e9f1-m3n2p");
gcLog.setMessage("频繁 Full GC 警告: 过去10分钟内发生 15 次 Full GC, 平均耗时 850ms, 建议检查内存泄漏");
gcLog.setMetrics(Map.of(
"full_gc_count", "15",
"avg_gc_time_ms", "850",
"survivor_space", "95%",
"old_gen", "89%"
));
logs.add(gcLog);
}
// 磁盘相关日志
if (query.contains("disk") || query.contains("filesystem")) {
for (int i = 0; i < 3; i++) {
LogEntry log = new LogEntry();
log.setTimestamp(FORMATTER.format(now.minus(i * 5, ChronoUnit.MINUTES)));
log.setLevel("WARN");
log.setService("log-collector");
log.setInstance("node-worker-01");
log.setMessage(String.format("磁盘使用率告警: /data 分区使用率 %.1f%%, 可用空间: %.1fGB",
85.0 + i * 2, 15.0 - i * 2));
log.setMetrics(Map.of(
"disk_usage", String.format("%.1f%%", 85.0 + i * 2),
"disk_available", String.format("%.1fGB", 15.0 - i * 2),
"disk_total", "100GB",
"mount_point", "/data",
"largest_dir", "/data/logs"
));
logs.add(log);
}
}
return logs;
}
/**
* 构建应用日志(与服务不可用、慢响应告警关联)
*/
private List<LogEntry> buildApplicationLogs(Instant now, String query, int limit) {
List<LogEntry> logs = new ArrayList<>();
// ERROR 级别日志
if (query.contains("error") || query.contains("fatal") || query.contains("500")) {
// 数据库连接错误
LogEntry dbError = new LogEntry();
dbError.setTimestamp(FORMATTER.format(now.minus(5, ChronoUnit.MINUTES)));
dbError.setLevel("ERROR");
dbError.setService("order-service");
dbError.setInstance("pod-order-service-5c7d8e9f1-m3n2p");
dbError.setMessage("数据库连接池耗尽: Cannot acquire connection from pool, " +
"active: 50/50, waiting: 23, timeout: 30000ms");
dbError.setMetrics(Map.of(
"error_type", "ConnectionPoolExhaustedException",
"pool_active", "50",
"pool_max", "50",
"waiting_threads", "23"
));
logs.add(dbError);
// OOM 错误
LogEntry oomError = new LogEntry();
oomError.setTimestamp(FORMATTER.format(now.minus(12, ChronoUnit.MINUTES)));
oomError.setLevel("FATAL");
oomError.setService("order-service");
oomError.setInstance("pod-order-service-5c7d8e9f1-m3n2p");
oomError.setMessage("java.lang.OutOfMemoryError: Java heap space at " +
"com.example.order.service.OrderService.processLargeOrder(OrderService.java:156)");
oomError.setMetrics(Map.of(
"error_type", "OutOfMemoryError",
"heap_used", "3.9GB",
"heap_max", "4GB",
"stack_trace", "OrderService.processLargeOrder -> OrderRepository.findByCondition -> HikariPool.getConnection"
));
logs.add(oomError);
// HTTP 500 错误
for (int i = 0; i < 3; i++) {
LogEntry httpError = new LogEntry();
httpError.setTimestamp(FORMATTER.format(now.minus(3 + i, ChronoUnit.MINUTES)));
httpError.setLevel("ERROR");
httpError.setService("user-service");
httpError.setInstance("pod-user-service-8e9f0a1b2-k5j6h");
httpError.setMessage(String.format("HTTP 500 Internal Server Error: /api/v1/users/profile, " +
"耗时: %dms, 错误: Database query timeout", 5200 + i * 300));
httpError.setMetrics(Map.of(
"http_status", "500",
"uri", "/api/v1/users/profile",
"method", "GET",
"duration_ms", String.valueOf(5200 + i * 300),
"error_cause", "QueryTimeoutException"
));
logs.add(httpError);
}
}
// 慢响应相关日志
if (query.contains("response_time") || query.contains("slow") || query.contains(">3000")) {
for (int i = 0; i < 5; i++) {
LogEntry slowLog = new LogEntry();
slowLog.setTimestamp(FORMATTER.format(now.minus(i * 2, ChronoUnit.MINUTES)));
slowLog.setLevel("WARN");
slowLog.setService("user-service");
slowLog.setInstance("pod-user-service-8e9f0a1b2-k5j6h");
slowLog.setMessage(String.format("慢请求警告: %s, 响应时间: %dms, 阈值: 3000ms",
i % 2 == 0 ? "/api/v1/users/profile" : "/api/v1/users/orders",
4200 - i * 150));
slowLog.setMetrics(Map.of(
"uri", i % 2 == 0 ? "/api/v1/users/profile" : "/api/v1/users/orders",
"response_time_ms", String.valueOf(4200 - i * 150),
"threshold_ms", "3000",
"db_time_ms", String.valueOf(3800 - i * 100),
"cache_hit", "false"
));
logs.add(slowLog);
}
}
// 下游服务依赖相关日志
if (query.contains("downstream") || query.contains("redis") ||
query.contains("database") || query.contains("mq")) {
LogEntry redisError = new LogEntry();
redisError.setTimestamp(FORMATTER.format(now.minus(7, ChronoUnit.MINUTES)));
redisError.setLevel("ERROR");
redisError.setService("payment-service");
redisError.setInstance("pod-payment-service-7d8f9c6b5-x2k4m");
redisError.setMessage("Redis 连接超时: 无法连接到 Redis 集群, 节点: redis-cluster-01:6379, 超时: 3000ms");
redisError.setMetrics(Map.of(
"dependency", "redis",
"host", "redis-cluster-01:6379",
"timeout_ms", "3000",
"retry_count", "3"
));
logs.add(redisError);
LogEntry mqError = new LogEntry();
mqError.setTimestamp(FORMATTER.format(now.minus(9, ChronoUnit.MINUTES)));
mqError.setLevel("WARN");
mqError.setService("order-service");
mqError.setInstance("pod-order-service-5c7d8e9f1-m3n2p");
mqError.setMessage("消息队列积压警告: 队列 order-process-queue 积压消息数: 15823, 消费速率下降");
mqError.setMetrics(Map.of(
"dependency", "rabbitmq",
"queue", "order-process-queue",
"pending_messages", "15823",
"consumer_count", "3"
));
logs.add(mqError);
}
return logs;
}
/**
* 构建数据库慢查询日志(与慢响应告警关联)
*/
private List<LogEntry> buildDatabaseSlowQueryLogs(Instant now, String query, int limit) {
List<LogEntry> logs = new ArrayList<>();
// 慢查询日志
LogEntry slowQuery1 = new LogEntry();
slowQuery1.setTimestamp(FORMATTER.format(now.minus(3, ChronoUnit.MINUTES)));
slowQuery1.setLevel("WARN");
slowQuery1.setService("mysql");
slowQuery1.setInstance("mysql-primary-01");
slowQuery1.setMessage("慢查询: SELECT * FROM orders WHERE user_id = ? AND status IN (?, ?, ?) " +
"ORDER BY created_at DESC LIMIT 100, 执行时间: 3.2s, 扫描行数: 1,245,678");
slowQuery1.setMetrics(Map.of(
"query_time_sec", "3.2",
"rows_examined", "1245678",
"rows_returned", "100",
"index_used", "idx_user_id",
"table", "orders",
"query_type", "SELECT"
));
logs.add(slowQuery1);
LogEntry slowQuery2 = new LogEntry();
slowQuery2.setTimestamp(FORMATTER.format(now.minus(6, ChronoUnit.MINUTES)));
slowQuery2.setLevel("WARN");
slowQuery2.setService("mysql");
slowQuery2.setInstance("mysql-primary-01");
slowQuery2.setMessage("慢查询: SELECT u.*, p.* FROM users u LEFT JOIN user_profiles p ON u.id = p.user_id " +
"WHERE u.last_login > ?, 执行时间: 2.8s, 全表扫描");
slowQuery2.setMetrics(Map.of(
"query_time_sec", "2.8",
"rows_examined", "856234",
"rows_returned", "45678",
"index_used", "NONE",
"table", "users, user_profiles",
"query_type", "SELECT",
"warning", "Full table scan detected"
));
logs.add(slowQuery2);
LogEntry slowQuery3 = new LogEntry();
slowQuery3.setTimestamp(FORMATTER.format(now.minus(8, ChronoUnit.MINUTES)));
slowQuery3.setLevel("WARN");
slowQuery3.setService("mysql");
slowQuery3.setInstance("mysql-primary-01");
slowQuery3.setMessage("慢查询: UPDATE orders SET status = ? WHERE created_at < ? AND status = ?, " +
"执行时间: 4.5s, 锁等待时间: 2.1s");
slowQuery3.setMetrics(Map.of(
"query_time_sec", "4.5",
"lock_time_sec", "2.1",
"rows_affected", "23456",
"table", "orders",
"query_type", "UPDATE",
"warning", "High lock contention"
));
logs.add(slowQuery3);
return logs;
}
/**
* 构建系统事件日志(与服务不可用告警关联)
*/
private List<LogEntry> buildSystemEventsLogs(Instant now, String query, int limit) {
List<LogEntry> logs = new ArrayList<>();
// 服务重启事件
if (query.contains("restart") || query.contains("crash") ||
query.contains("oom_kill")) {
LogEntry restartEvent = new LogEntry();
restartEvent.setTimestamp(FORMATTER.format(now.minus(15, ChronoUnit.MINUTES)));
restartEvent.setLevel("WARN");
restartEvent.setService("kubernetes");
restartEvent.setInstance("kube-controller-manager");
restartEvent.setMessage("Pod 重启事件: pod-order-service-5c7d8e9f1-m3n2p, 原因: OOMKilled, " +
"容器退出码: 137, 重启次数: 3");
restartEvent.setMetrics(Map.of(
"event_type", "PodRestart",
"pod", "pod-order-service-5c7d8e9f1-m3n2p",
"reason", "OOMKilled",
"exit_code", "137",
"restart_count", "3",
"namespace", "production"
));
logs.add(restartEvent);
LogEntry oomKillEvent = new LogEntry();
oomKillEvent.setTimestamp(FORMATTER.format(now.minus(16, ChronoUnit.MINUTES)));
oomKillEvent.setLevel("ERROR");
oomKillEvent.setService("kernel");
oomKillEvent.setInstance("node-worker-02");
oomKillEvent.setMessage("OOM Killer 触发: 进程 java (PID: 12345) 被杀死, " +
"内存使用: 3.9GB, 内存限制: 4GB");
oomKillEvent.setMetrics(Map.of(
"event_type", "OOMKill",
"process", "java",
"pid", "12345",
"memory_used", "3.9GB",
"memory_limit", "4GB",
"cgroup", "/kubepods/pod-order-service"
));
logs.add(oomKillEvent);
}
return logs;
}
/**
* 构建通用日志
*/
private List<LogEntry> buildGenericLogs(Instant now, String query, int limit) {
List<LogEntry> logs = new ArrayList<>();
for (int i = 0; i < Math.min(limit, 10); i++) {
LogEntry log = new LogEntry();
log.setTimestamp(FORMATTER.format(now.minus(i, ChronoUnit.MINUTES)));
log.setLevel(i % 3 == 0 ? "ERROR" : (i % 3 == 1 ? "WARN" : "INFO"));
log.setService("generic-service");
log.setInstance("instance-" + i);
log.setMessage("日志消息 #" + i + ", 查询条件: " + query);
log.setMetrics(new HashMap<>());
logs.add(log);
}
return logs;
}
/**
* 构建错误响应
*/
private String buildErrorResponse(String message) {
try {
QueryLogsOutput output = new QueryLogsOutput();
output.setSuccess(false);
output.setMessage(message);
return objectMapper.writerWithDefaultPrettyPrinter().writeValueAsString(output);
} catch (Exception e) {
return String.format("{\"success\":false,\"message\":\"%s\"}", message);
}
}
// ==================== 数据模型 ====================
/**
* 日志条目
*/
@Data
public static class LogEntry {
@JsonProperty("timestamp")
private String timestamp;
@JsonProperty("level")
private String level;
@JsonProperty("service")
private String service;
@JsonProperty("instance")
private String instance;
@JsonProperty("message")
private String message;
@JsonProperty("metrics")
private Map<String, String> metrics;
}
/**
* 日志查询输出
*/
@Data
public static class QueryLogsOutput {
@JsonProperty("success")
private boolean success;
@JsonProperty("region")
private String region;
@JsonProperty("log_topic")
private String logTopic;
@JsonProperty("query")
private String query;
@JsonProperty("logs")
private List<LogEntry> logs;
@JsonProperty("total")
private int total;
@JsonProperty("message")
private String message;
}
/**
* 日志主题信息
*/
@Data
public static class LogTopicInfo {
@JsonProperty("topic_name")
private String topicName;
@JsonProperty("description")
private String description;
@JsonProperty("example_queries")
private List<String> exampleQueries;
@JsonProperty("related_alerts")
private List<String> relatedAlerts;
}
/**
* 日志主题列表输出
*/
@Data
public static class LogTopicsOutput {
@JsonProperty("success")
private boolean success;
@JsonProperty("topics")
private List<LogTopicInfo> topics;
@JsonProperty("available_regions")
private List<String> availableRegions;
@JsonProperty("default_region")
private String defaultRegion;
@JsonProperty("message")
private String message;
}
}
@@ -0,0 +1,303 @@
package org.example.agent.tool;
import com.fasterxml.jackson.annotation.JsonProperty;
import com.fasterxml.jackson.databind.ObjectMapper;
import lombok.Data;
import okhttp3.OkHttpClient;
import okhttp3.Request;
import okhttp3.Response;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.tool.annotation.Tool;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.stereotype.Component;
import java.time.Duration;
import java.time.Instant;
import java.time.temporal.ChronoUnit;
import java.util.*;
/**
* Prometheus 告警查询工具
* 用于查询 Prometheus 的活动告警信息
*/
@Component
public class QueryMetricsTools {
private static final Logger logger = LoggerFactory.getLogger(QueryMetricsTools.class);
/** 工具名常量,用于动态构建提示词 */
public static final String TOOL_QUERY_PROMETHEUS_ALERTS = "queryPrometheusAlerts";
private final ObjectMapper objectMapper = new ObjectMapper();
@Value("${prometheus.base-url}")
private String prometheusBaseUrl;
@Value("${prometheus.timeout:10}")
private int timeout;
@Value("${prometheus.mock-enabled:false}")
private boolean mockEnabled;
private OkHttpClient httpClient;
@jakarta.annotation.PostConstruct
public void init() {
this.httpClient = new OkHttpClient.Builder()
.connectTimeout(Duration.ofSeconds(timeout))
.readTimeout(Duration.ofSeconds(timeout))
.build();
logger.info("✅ QueryMetricsTools 初始化成功, Prometheus URL: {}, Mock模式: {}", prometheusBaseUrl, mockEnabled);
}
/**
* 查询 Prometheus 活动告警
* 该工具从 Prometheus 告警系统检索所有当前活动/触发的告警,包括标签、注释、状态和值
*/
@Tool(description = "Query active alerts from Prometheus alerting system. " +
"This tool retrieves all currently active/firing alerts including their labels, annotations, state, and values. " +
"Use this tool when you need to check what alerts are currently firing, investigate alert conditions, or monitor alert status.")
public String queryPrometheusAlerts() {
logger.info("开始查询 Prometheus 活动告警, Mock模式: {}", mockEnabled);
try {
List<SimplifiedAlert> simplifiedAlerts;
if (mockEnabled) {
// Mock 模式:返回与文档关联的模拟告警数据
simplifiedAlerts = buildMockAlerts();
logger.info("使用 Mock 数据,返回 {} 个模拟告警", simplifiedAlerts.size());
} else {
// 真实模式:调用 Prometheus Alerts API
PrometheusAlertsResult result = fetchPrometheusAlerts();
if (!"success".equals(result.getStatus())) {
return buildErrorResponse("Prometheus API 返回非成功状态: " + result.getStatus(), result.getError());
}
// 转换为简化格式,对于相同的 alertname,只保留第一个
Set<String> seenAlertNames = new HashSet<>();
simplifiedAlerts = new ArrayList<>();
for (PrometheusAlert alert : result.getData().getAlerts()) {
String alertName = alert.getLabels().get("alertname");
// 如果这个 alertname 已经存在,跳过
if (seenAlertNames.contains(alertName)) {
continue;
}
// 标记为已见过
seenAlertNames.add(alertName);
SimplifiedAlert simplified = new SimplifiedAlert();
simplified.setAlertName(alertName);
simplified.setDescription(alert.getAnnotations().getOrDefault("description", ""));
simplified.setState(alert.getState());
simplified.setActiveAt(alert.getActiveAt());
simplified.setDuration(calculateDuration(alert.getActiveAt()));
simplifiedAlerts.add(simplified);
}
}
// 构建成功响应
PrometheusAlertsOutput output = new PrometheusAlertsOutput();
output.setSuccess(true);
output.setAlerts(simplifiedAlerts);
output.setMessage(String.format("成功检索到 %d 个活动告警", simplifiedAlerts.size()));
String jsonResult = objectMapper.writerWithDefaultPrettyPrinter().writeValueAsString(output);
logger.info("Prometheus 告警查询完成: 找到 {} 个告警", simplifiedAlerts.size());
return jsonResult;
} catch (Exception e) {
logger.error("查询 Prometheus 告警失败", e);
return buildErrorResponse("查询失败", e.getMessage());
}
}
/**
* 构建 Mock 告警数据
* 与 aiops-docs 文档中的告警类型对应:
* - HighCPUUsage: CPU使用率过高
* - HighMemoryUsage: 内存使用率过高
* - HighDiskUsage: 磁盘使用率过高
* - ServiceUnavailable: 服务不可用
* - SlowResponse: 响应时间过长
*/
private List<SimplifiedAlert> buildMockAlerts() {
List<SimplifiedAlert> alerts = new ArrayList<>();
Instant now = Instant.now();
// 告警1: CPU使用率过高 - 持续约25分钟
SimplifiedAlert cpuAlert = new SimplifiedAlert();
cpuAlert.setAlertName("HighCPUUsage");
cpuAlert.setDescription("服务 payment-service 的 CPU 使用率持续超过 80%,当前值为 92%。" +
"实例: pod-payment-service-7d8f9c6b5-x2k4m,命名空间: production");
cpuAlert.setState("firing");
Instant cpuActiveAt = now.minus(25, ChronoUnit.MINUTES);
cpuAlert.setActiveAt(cpuActiveAt.toString());
cpuAlert.setDuration(calculateDuration(cpuActiveAt.toString()));
alerts.add(cpuAlert);
// 告警2: 内存使用率过高 - 持续约15分钟
SimplifiedAlert memoryAlert = new SimplifiedAlert();
memoryAlert.setAlertName("HighMemoryUsage");
memoryAlert.setDescription("服务 order-service 的内存使用率持续超过 85%,当前值为 91%。" +
"JVM堆内存使用: 3.8GB/4GB,可能存在内存泄漏风险。" +
"实例: pod-order-service-5c7d8e9f1-m3n2p,命名空间: production");
memoryAlert.setState("firing");
Instant memoryActiveAt = now.minus(15, ChronoUnit.MINUTES);
memoryAlert.setActiveAt(memoryActiveAt.toString());
memoryAlert.setDuration(calculateDuration(memoryActiveAt.toString()));
alerts.add(memoryAlert);
// 告警3: 响应时间过长 - 持续约10分钟
SimplifiedAlert slowAlert = new SimplifiedAlert();
slowAlert.setAlertName("SlowResponse");
slowAlert.setDescription("服务 user-service 的 P99 响应时间持续超过 3 秒,当前值为 4.2 秒。" +
"受影响接口: /api/v1/users/profile, /api/v1/users/orders。" +
"可能原因:数据库慢查询或下游服务延迟");
slowAlert.setState("firing");
Instant slowActiveAt = now.minus(10, ChronoUnit.MINUTES);
slowAlert.setActiveAt(slowActiveAt.toString());
slowAlert.setDuration(calculateDuration(slowActiveAt.toString()));
alerts.add(slowAlert);
return alerts;
}
/**
* 从 Prometheus API 获取告警数据
*/
private PrometheusAlertsResult fetchPrometheusAlerts() throws Exception {
String apiUrl = prometheusBaseUrl + "/api/v1/alerts";
logger.debug("请求 Prometheus API: {}", apiUrl);
Request request = new Request.Builder()
.url(apiUrl)
.get()
.build();
try (Response response = httpClient.newCall(request).execute()) {
if (!response.isSuccessful()) {
throw new RuntimeException("HTTP 请求失败: " + response.code());
}
String responseBody = response.body().string();
return objectMapper.readValue(responseBody, PrometheusAlertsResult.class);
}
}
/**
* 计算从 activeAt 到现在的持续时间
*/
private String calculateDuration(String activeAtStr) {
try {
Instant activeAt = Instant.parse(activeAtStr);
Duration duration = Duration.between(activeAt, Instant.now());
long hours = duration.toHours();
long minutes = duration.toMinutes() % 60;
long seconds = duration.getSeconds() % 60;
if (hours > 0) {
return String.format("%dh%dm%ds", hours, minutes, seconds);
} else if (minutes > 0) {
return String.format("%dm%ds", minutes, seconds);
} else {
return String.format("%ds", seconds);
}
} catch (Exception e) {
logger.warn("解析时间失败: {}", activeAtStr, e);
return "unknown";
}
}
/**
* 构建错误响应
*/
private String buildErrorResponse(String message, String error) {
try {
PrometheusAlertsOutput output = new PrometheusAlertsOutput();
output.setSuccess(false);
output.setMessage(message);
output.setError(error);
return objectMapper.writerWithDefaultPrettyPrinter().writeValueAsString(output);
} catch (Exception e) {
return String.format("{\"success\":false,\"message\":\"%s\",\"error\":\"%s\"}", message, error);
}
}
// ==================== 数据模型 ====================
/**
* Prometheus 告警信息结构
*/
@Data
public static class PrometheusAlert {
private Map<String, String> labels;
private Map<String, String> annotations;
private String state;
private String activeAt;
private String value;
}
/**
* Prometheus 告警查询结果
*/
@Data
public static class PrometheusAlertsResult {
private String status;
private AlertsData data;
private String error;
private String errorType;
}
@Data
public static class AlertsData {
private List<PrometheusAlert> alerts = new ArrayList<>();
}
/**
* 简化的告警信息
*/
@Data
public static class SimplifiedAlert {
@JsonProperty("alert_name")
private String alertName;
@JsonProperty("description")
private String description;
@JsonProperty("state")
private String state;
@JsonProperty("active_at")
private String activeAt;
@JsonProperty("duration")
private String duration;
}
/**
* 告警查询输出
*/
@Data
public static class PrometheusAlertsOutput {
@JsonProperty("success")
private boolean success;
@JsonProperty("alerts")
private List<SimplifiedAlert> alerts;
@JsonProperty("message")
private String message;
@JsonProperty("error")
private String error;
}
}
@@ -0,0 +1,179 @@
package org.example.client;
import io.milvus.client.MilvusServiceClient;
import io.milvus.grpc.DataType;
import io.milvus.param.ConnectParam;
import io.milvus.param.IndexType;
import io.milvus.param.MetricType;
import io.milvus.param.R;
import io.milvus.param.RpcStatus;
import io.milvus.param.collection.*;
import io.milvus.param.index.CreateIndexParam;
import org.example.config.MilvusProperties;
import org.example.constant.MilvusConstants;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.stereotype.Component;
import java.util.concurrent.TimeUnit;
/**
* Milvus 客户端工厂类
* 负责创建和初始化 Milvus 客户端连接
*/
@Component
public class MilvusClientFactory {
private static final Logger logger = LoggerFactory.getLogger(MilvusClientFactory.class);
@Autowired
private MilvusProperties milvusProperties;
/**
* 创建并初始化 Milvus 客户端
*
* 简化版本:直接连接并创建 collection
*
* @return MilvusServiceClient 实例
* @throws RuntimeException 如果连接或初始化失败
*/
public MilvusServiceClient createClient() {
MilvusServiceClient client = null;
try {
// 1. 连接到 Milvus
logger.info("正在连接到 Milvus: {}:{}", milvusProperties.getHost(), milvusProperties.getPort());
client = connectToMilvus();
logger.info("成功连接到 Milvus");
// 2. 检查并创建 biz collection(如果不存在)
if (!collectionExists(client, MilvusConstants.MILVUS_COLLECTION_NAME)) {
logger.info("collection '{}' 不存在,正在创建...", MilvusConstants.MILVUS_COLLECTION_NAME);
createBizCollection(client);
logger.info("成功创建 collection '{}'", MilvusConstants.MILVUS_COLLECTION_NAME);
// 创建索引
createIndexes(client);
logger.info("成功创建索引");
} else {
logger.info("collection '{}' 已存在", MilvusConstants.MILVUS_COLLECTION_NAME);
}
return client;
} catch (Exception e) {
logger.error("创建 Milvus 客户端失败", e);
if (client != null) {
client.close();
}
throw new RuntimeException("创建 Milvus 客户端失败: " + e.getMessage(), e);
}
}
/**
* 连接到 Milvus
*/
private MilvusServiceClient connectToMilvus() {
ConnectParam.Builder builder = ConnectParam.newBuilder()
.withHost(milvusProperties.getHost())
.withPort(milvusProperties.getPort())
.withConnectTimeout(milvusProperties.getTimeout(), TimeUnit.MILLISECONDS);
// 如果配置了用户名和密码
if (milvusProperties.getUsername() != null && !milvusProperties.getUsername().isEmpty()) {
builder.withAuthorization(milvusProperties.getUsername(), milvusProperties.getPassword());
}
return new MilvusServiceClient(builder.build());
}
/**
* 检查 collection 是否存在
*/
private boolean collectionExists(MilvusServiceClient client, String collectionName) {
R<Boolean> response = client.hasCollection(HasCollectionParam.newBuilder()
.withCollectionName(collectionName)
.build());
if (response.getStatus() != 0) {
throw new RuntimeException("检查 collection 失败: " + response.getMessage());
}
return response.getData();
}
/**
* 创建 biz collection
*/
private void createBizCollection(MilvusServiceClient client) {
// 定义字段
FieldType idField = FieldType.newBuilder()
.withName("id")
.withDataType(DataType.VarChar)
.withMaxLength(MilvusConstants.ID_MAX_LENGTH)
.withPrimaryKey(true)
.build();
FieldType vectorField = FieldType.newBuilder()
.withName("vector")
.withDataType(DataType.FloatVector) // 改为 FloatVector
.withDimension(MilvusConstants.VECTOR_DIM)
.build();
FieldType contentField = FieldType.newBuilder()
.withName("content")
.withDataType(DataType.VarChar)
.withMaxLength(MilvusConstants.CONTENT_MAX_LENGTH)
.build();
FieldType metadataField = FieldType.newBuilder()
.withName("metadata")
.withDataType(DataType.JSON)
.build();
// 创建 collection schema
CollectionSchemaParam schema = CollectionSchemaParam.newBuilder()
.withEnableDynamicField(false)
.addFieldType(idField)
.addFieldType(vectorField)
.addFieldType(contentField)
.addFieldType(metadataField)
.build();
// 创建 collection
CreateCollectionParam createParam = CreateCollectionParam.newBuilder()
.withCollectionName(MilvusConstants.MILVUS_COLLECTION_NAME)
.withDescription("Business knowledge collection")
.withSchema(schema)
.withShardsNum(MilvusConstants.DEFAULT_SHARD_NUMBER)
.build();
R<RpcStatus> response = client.createCollection(createParam);
if (response.getStatus() != 0) {
throw new RuntimeException("创建 collection 失败: " + response.getMessage());
}
}
/**
* 为 collection 创建索引
*/
private void createIndexes(MilvusServiceClient client) {
// 为 vector 字段创建索引(FloatVector 使用 IVF_FLAT 和 L2 距离)
CreateIndexParam vectorIndexParam = CreateIndexParam.newBuilder()
.withCollectionName(MilvusConstants.MILVUS_COLLECTION_NAME)
.withFieldName("vector")
.withIndexType(IndexType.IVF_FLAT)
.withMetricType(MetricType.L2) // L2 距离(欧氏距离)
.withExtraParam("{\"nlist\":128}")
.withSyncMode(Boolean.FALSE)
.build();
R<RpcStatus> response = client.createIndex(vectorIndexParam);
if (response.getStatus() != 0) {
throw new RuntimeException("创建 vector 索引失败: " + response.getMessage());
}
logger.info("成功为 vector 字段创建索引");
}
}
@@ -0,0 +1,40 @@
package org.example.config;
import okhttp3.OkHttpClient;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.http.client.OkHttp3ClientHttpRequestFactory;
import org.springframework.web.client.RestClient;
import java.time.Duration;
/**
* DashScope API 配置
* 用于配置超时时间等参数
*/
@Configuration
public class DashScopeConfig {
@Value("${spring.ai.dashscope.chat.options.timeout:180000}")
private long timeout;
/**
* 配置 RestClient.Builder,设置超时时间
* Spring AI 会自动使用这个 Bean
*/
@Bean
public RestClient.Builder restClientBuilder() {
// 创建自定义的 OkHttpClient,设置超时时间
OkHttpClient okHttpClient = new OkHttpClient.Builder()
.connectTimeout(Duration.ofMillis(timeout))
.readTimeout(Duration.ofMillis(timeout))
.writeTimeout(Duration.ofMillis(timeout))
.callTimeout(Duration.ofMillis(timeout))
.build();
// 创建 RestClient.Builder 并配置 OkHttpClient
return RestClient.builder()
.requestFactory(new OkHttp3ClientHttpRequestFactory(okHttpClient));
}
}
@@ -0,0 +1,32 @@
package org.example.config;
import lombok.Getter;
import org.springframework.boot.context.properties.ConfigurationProperties;
import org.springframework.context.annotation.Configuration;
/**
* 文档分片配置
*/
@Getter
@Configuration
@ConfigurationProperties(prefix = "document.chunk")
public class DocumentChunkConfig {
/**
* 每个分片的最大字符数
*/
private int maxSize = 800;
/**
* 分片之间的重叠字符数
*/
private int overlap = 100;
public void setMaxSize(int maxSize) {
this.maxSize = maxSize;
}
public void setOverlap(int overlap) {
this.overlap = overlap;
}
}
@@ -0,0 +1,22 @@
package org.example.config;
import lombok.Getter;
import org.springframework.boot.context.properties.ConfigurationProperties;
import org.springframework.context.annotation.Configuration;
@Getter
@Configuration
@ConfigurationProperties(prefix = "file.upload")
public class FileUploadConfig {
private String path;
private String allowedExtensions;
public void setPath(String path) {
this.path = path;
}
public void setAllowedExtensions(String allowedExtensions) {
this.allowedExtensions = allowedExtensions;
}
}
@@ -0,0 +1,51 @@
package org.example.config;
import io.milvus.client.MilvusServiceClient;
import org.example.client.MilvusClientFactory;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import jakarta.annotation.PreDestroy;
/**
* Milvus 配置类
* 负责创建和管理 MilvusServiceClient Bean
*/
@Configuration
public class MilvusConfig {
private static final Logger logger = LoggerFactory.getLogger(MilvusConfig.class);
@Autowired
private MilvusClientFactory milvusClientFactory;
private MilvusServiceClient milvusClient;
/**
* 创建 MilvusServiceClient Bean
*
* @return MilvusServiceClient 实例
*/
@Bean
public MilvusServiceClient milvusServiceClient() {
logger.info("正在初始化 Milvus 客户端...");
milvusClient = milvusClientFactory.createClient();
logger.info("Milvus 客户端初始化完成");
return milvusClient;
}
/**
* 应用关闭时清理资源
*/
@PreDestroy
public void cleanup() {
if (milvusClient != null) {
logger.info("正在关闭 Milvus 客户端连接...");
milvusClient.close();
logger.info("Milvus 客户端连接已关闭");
}
}
}
@@ -0,0 +1,68 @@
package org.example.config;
import org.springframework.boot.context.properties.ConfigurationProperties;
import org.springframework.context.annotation.Configuration;
@Configuration
@ConfigurationProperties(prefix = "milvus")
public class MilvusProperties {
private String host = "localhost";
private Integer port = 19530;
private String username = "";
private String password = "";
private String database = "default";
private Long timeout = 10000L;
public String getHost() {
return host;
}
public void setHost(String host) {
this.host = host;
}
public Integer getPort() {
return port;
}
public void setPort(Integer port) {
this.port = port;
}
public String getUsername() {
return username;
}
public void setUsername(String username) {
this.username = username;
}
public String getPassword() {
return password;
}
public void setPassword(String password) {
this.password = password;
}
public String getDatabase() {
return database;
}
public void setDatabase(String database) {
this.database = database;
}
public Long getTimeout() {
return timeout;
}
public void setTimeout(Long timeout) {
this.timeout = timeout;
}
public String getAddress() {
return host + ":" + port;
}
}
@@ -0,0 +1,38 @@
package org.example.config;
import com.fasterxml.jackson.databind.ObjectMapper;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.http.converter.HttpMessageConverter;
import org.springframework.http.converter.StringHttpMessageConverter;
import org.springframework.http.converter.json.MappingJackson2HttpMessageConverter;
import org.springframework.web.servlet.config.annotation.WebMvcConfigurer;
import java.nio.charset.StandardCharsets;
import java.util.List;
/**
* Web MVC 配置
* 解决中文乱码问题
*/
@Configuration
public class WebConfig implements WebMvcConfigurer {
@Override
public void configureMessageConverters(List<HttpMessageConverter<?>> converters) {
// 添加 UTF-8 字符串转换器
StringHttpMessageConverter stringConverter = new StringHttpMessageConverter(StandardCharsets.UTF_8);
stringConverter.setWriteAcceptCharset(false); // 不设置 Accept-Charset
converters.add(0, stringConverter);
// 添加 Jackson JSON 转换器,确保 UTF-8 编码
MappingJackson2HttpMessageConverter jsonConverter = new MappingJackson2HttpMessageConverter();
jsonConverter.setDefaultCharset(StandardCharsets.UTF_8);
converters.add(1, jsonConverter);
}
@Bean
public ObjectMapper objectMapper() {
return new ObjectMapper();
}
}
@@ -0,0 +1,30 @@
package org.example.config;
import org.springframework.context.annotation.Configuration;
import org.springframework.web.servlet.config.annotation.CorsRegistry;
import org.springframework.web.servlet.config.annotation.ResourceHandlerRegistry;
import org.springframework.web.servlet.config.annotation.WebMvcConfigurer;
/**
* Web MVC 配置
* 配置跨域和静态资源
*/
@Configuration
public class WebMvcConfig implements WebMvcConfigurer {
@Override
public void addCorsMappings(CorsRegistry registry) {
registry.addMapping("/**")
.allowedOrigins("*")
.allowedMethods("GET", "POST", "PUT", "DELETE", "OPTIONS")
.allowedHeaders("*")
.maxAge(3600);
}
@Override
public void addResourceHandlers(ResourceHandlerRegistry registry) {
// 配置静态资源映射
registry.addResourceHandler("/**")
.addResourceLocations("classpath:/static/");
}
}
@@ -0,0 +1,38 @@
package org.example.constant;
public class MilvusConstants {
/**
* Milvus 数据库名称
*/
public static final String MILVUS_DB_NAME = "default";
/**
* Milvus 集合名称
*/
public static final String MILVUS_COLLECTION_NAME = "biz";
/**
* 向量维度(豆包 embedding 模型的维度)
*/
public static final int VECTOR_DIM = 1024; // 豆包模型返回1024维向量
/**
* ID字段最大长度
*/
public static final int ID_MAX_LENGTH = 256;
/**
* Content字段最大长度
*/
public static final int CONTENT_MAX_LENGTH = 8192;
/**
* 默认分片数
*/
public static final int DEFAULT_SHARD_NUMBER = 2;
private MilvusConstants() {
// 工具类,禁止实例化
}
}
@@ -0,0 +1,629 @@
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;
import com.alibaba.cloud.ai.graph.streaming.OutputType;
import com.alibaba.cloud.ai.graph.streaming.StreamingOutput;
import lombok.Getter;
import lombok.Setter;
import org.example.service.AiOpsService;
import org.example.service.ChatService;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.tool.ToolCallback;
import org.springframework.ai.tool.ToolCallbackProvider;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.http.MediaType;
import org.springframework.http.ResponseEntity;
import org.springframework.web.bind.annotation.*;
import org.springframework.web.servlet.mvc.method.annotation.SseEmitter;
import reactor.core.publisher.Flux;
import java.io.IOException;
import java.util.*;
import java.util.concurrent.ConcurrentHashMap;
import java.util.concurrent.ExecutorService;
import java.util.concurrent.Executors;
import java.util.concurrent.locks.ReentrantLock;
/**
* 统一 API 控制器
* 适配前端接口需求
*/
@RestController
@RequestMapping("/api")
public class ChatController {
private static final Logger logger = LoggerFactory.getLogger(ChatController.class);
@Autowired
private AiOpsService aiOpsService;
@Autowired
private ChatService chatService;
@Autowired
private ToolCallbackProvider tools;
private final ExecutorService executor = Executors.newCachedThreadPool();
// 存储会话信息
private final Map<String, SessionInfo> sessions = new ConcurrentHashMap<>();
// 最大历史消息窗口大小(成对计算:用户消息+AI回复=1对)
private static final int MAX_WINDOW_SIZE = 6;
/**
* 普通对话接口(支持工具调用)
* 与 /chat_react 逻辑一致,但直接返回完整结果而非流式输出
*/
@PostMapping("/chat")
public ResponseEntity<ApiResponse<ChatResponse>> chat(@RequestBody ChatRequest request) {
try {
logger.info("收到对话请求 - SessionId: {}, Question: {}", request.getId(), request.getQuestion());
// 参数校验
if (request.getQuestion() == null || request.getQuestion().trim().isEmpty()) {
logger.warn("问题内容为空");
return ResponseEntity.ok(ApiResponse.success(ChatResponse.error("问题内容不能为空")));
}
// 获取或创建会话
SessionInfo session = getOrCreateSession(request.getId());
// 获取历史消息
List<Map<String, String>> history = session.getHistory();
logger.info("会话历史消息对数: {}", history.size() / 2);
// 创建 DashScope API 和 ChatModel
DashScopeApi dashScopeApi = chatService.createDashScopeApi();
DashScopeChatModel chatModel = chatService.createStandardChatModel(dashScopeApi);
// 记录可用工具
chatService.logAvailableTools();
logger.info("开始 ReactAgent 对话(支持自动工具调用)");
// 构建系统提示词(包含历史消息)
String systemPrompt = chatService.buildSystemPrompt(history);
// 创建 ReactAgent
ReactAgent agent = chatService.createReactAgent(chatModel, systemPrompt);
// 执行对话
String fullAnswer = chatService.executeChat(agent, request.getQuestion());
// 更新会话历史
session.addMessage(request.getQuestion(), fullAnswer);
logger.info("已更新会话历史 - SessionId: {}, 当前消息对数: {}",
request.getId(), session.getMessagePairCount());
return ResponseEntity.ok(ApiResponse.success(ChatResponse.success(fullAnswer)));
} catch (Exception e) {
logger.error("对话失败", e);
return ResponseEntity.ok(ApiResponse.success(ChatResponse.error(e.getMessage())));
}
}
/**
* 清空会话历史
*/
@PostMapping("/chat/clear")
public ResponseEntity<ApiResponse<String>> clearChatHistory(@RequestBody ClearRequest request) {
try {
logger.info("收到清空会话历史请求 - SessionId: {}", request.getId());
if (request.getId() == null || request.getId().isEmpty()) {
return ResponseEntity.ok(ApiResponse.error("会话ID不能为空"));
}
SessionInfo session = sessions.get(request.getId());
if (session != null) {
session.clearHistory();
return ResponseEntity.ok(ApiResponse.success("会话历史已清空"));
} else {
return ResponseEntity.ok(ApiResponse.error("会话不存在"));
}
} catch (Exception e) {
logger.error("清空会话历史失败", e);
return ResponseEntity.ok(ApiResponse.error(e.getMessage()));
}
}
/**
* ReactAgent 对话接口(SSE 流式模式,支持多轮对话,支持自动工具调用,例如获取当前时间,查询日志,告警等)
* 支持 session 管理,保留对话历史
*/
@PostMapping(value = "/chat_stream", produces = "text/event-stream;charset=UTF-8")
public SseEmitter chatStream(@RequestBody ChatRequest request) {
SseEmitter emitter = new SseEmitter(300000L); // 5分钟超时
// 参数校验
if (request.getQuestion() == null || request.getQuestion().trim().isEmpty()) {
logger.warn("问题内容为空");
try {
emitter.send(SseEmitter.event().name("message").data(SseMessage.error("问题内容不能为空"), MediaType.APPLICATION_JSON));
emitter.complete();
} catch (IOException e) {
emitter.completeWithError(e);
}
return emitter;
}
executor.execute(() -> {
try {
logger.info("收到 ReactAgent 对话请求 - SessionId: {}, Question: {}", request.getId(), request.getQuestion());
// 获取或创建会话
SessionInfo session = getOrCreateSession(request.getId());
// 获取历史消息
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);
// 记录可用工具
chatService.logAvailableTools();
logger.info("开始 ReactAgent 流式对话(支持自动工具调用)");
// 构建系统提示词(包含历史消息)
String systemPrompt = chatService.buildSystemPrompt(history);
// 创建 ReactAgent
ReactAgent agent = chatService.createReactAgent(chatModel, systemPrompt);
// 用于累积完整答案
StringBuilder fullAnswerBuilder = new StringBuilder();
// 使用 agent.stream() 进行流式对话
Flux<NodeOutput> stream = agent.stream(request.getQuestion());
stream.subscribe(
output -> {
try {
// 检查是否为 StreamingOutput 类型
if (output instanceof StreamingOutput streamingOutput) {
OutputType type = streamingOutput.getOutputType();
// 处理模型推理的流式输出
if (type == OutputType.AGENT_MODEL_STREAMING) {
// 流式增量内容,逐步显示
String chunk = streamingOutput.message().getText();
if (chunk != null && !chunk.isEmpty()) {
fullAnswerBuilder.append(chunk);
// 实时发送到前端
emitter.send(SseEmitter.event()
.name("message")
.data(SseMessage.content(chunk), MediaType.APPLICATION_JSON));
logger.info("发送流式内容: {}", chunk);
}
} else if (type == OutputType.AGENT_MODEL_FINISHED) {
// 模型推理完成
logger.info("模型输出完成");
} else if (type == OutputType.AGENT_TOOL_FINISHED) {
// 工具调用完成
logger.info("工具调用完成: {}", output.node());
} else if (type == OutputType.AGENT_HOOK_FINISHED) {
// Hook 执行完成
logger.debug("Hook 执行完成: {}", output.node());
}
}
} catch (IOException e) {
logger.error("发送流式消息失败", e);
throw new RuntimeException(e);
}
},
error -> {
// 错误处理
logger.error("ReactAgent 流式对话失败", error);
try {
emitter.send(SseEmitter.event()
.name("message")
.data(SseMessage.error(error.getMessage()), MediaType.APPLICATION_JSON));
} catch (IOException ex) {
logger.error("发送错误消息失败", ex);
}
emitter.completeWithError(error);
},
() -> {
// 完成处理
try {
String fullAnswer = fullAnswerBuilder.toString();
logger.info("ReactAgent 流式对话完成 - SessionId: {}, 答案长度: {}",
request.getId(), fullAnswer.length());
// 更新会话历史
session.addMessage(request.getQuestion(), fullAnswer);
logger.info("已更新会话历史 - SessionId: {}, 当前消息对数: {}",
request.getId(), session.getMessagePairCount());
// 发送完成标记
emitter.send(SseEmitter.event()
.name("message")
.data(SseMessage.done(), MediaType.APPLICATION_JSON));
emitter.complete();
} catch (IOException e) {
logger.error("发送完成消息失败", e);
emitter.completeWithError(e);
}
}
);
} catch (Exception e) {
logger.error("ReactAgent 对话初始化失败", e);
try {
emitter.send(SseEmitter.event()
.name("message")
.data(SseMessage.error(e.getMessage()), MediaType.APPLICATION_JSON));
} catch (IOException ex) {
logger.error("发送错误消息失败", ex);
}
emitter.completeWithError(e);
}
});
return emitter;
}
/**
* AI 智能运维接口(SSE 流式模式)- 自动分析告警并生成运维报告
* 无需用户输入,自动执行告警分析流程
*/
@PostMapping(value = "/ai_ops", produces = "text/event-stream;charset=UTF-8")
public SseEmitter aiOps() {
SseEmitter emitter = new SseEmitter(600000L); // 10分钟超时(告警分析可能较慢)
executor.execute(() -> {
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();
ToolCallback[] toolCallbacks = tools.getToolCallbacks();
emitter.send(SseEmitter.event().name("message").data(SseMessage.content("正在读取告警并拆解任务...\n")));
// 调用 AiOpsService 执行分析流程
Optional<OverAllState> overAllStateOptional = aiOpsService.executeAiOpsAnalysis(chatModel, toolCallbacks);
if (overAllStateOptional.isEmpty()) {
emitter.send(SseEmitter.event().name("message")
.data(SseMessage.error("多 Agent 编排未获取到有效结果"), MediaType.APPLICATION_JSON));
emitter.complete();
return;
}
OverAllState state = overAllStateOptional.get();
logger.info("AI Ops 编排完成,开始提取最终报告...");
// 提取最终报告
Optional<String> finalReportOptional = aiOpsService.extractFinalReport(state);
// 输出最终报告
if (finalReportOptional.isPresent()) {
String finalReportText = finalReportOptional.get();
logger.info("提取到 Planner 最终报告,长度: {}", finalReportText.length());
// 发送分隔线
emitter.send(SseEmitter.event().name("message")
.data(SseMessage.content("\n\n" + "=".repeat(60) + "\n"), MediaType.APPLICATION_JSON));
// 发送完整的告警分析报告
emitter.send(SseEmitter.event().name("message")
.data(SseMessage.content("📋 **告警分析报告**\n\n"), MediaType.APPLICATION_JSON));
int chunkSize = 50;
for (int i = 0; i < finalReportText.length(); i += chunkSize) {
int end = Math.min(i + chunkSize, finalReportText.length());
String chunk = finalReportText.substring(i, end);
emitter.send(SseEmitter.event().name("message")
.data(SseMessage.content(chunk), MediaType.APPLICATION_JSON));
}
// 发送结束分隔线
emitter.send(SseEmitter.event().name("message")
.data(SseMessage.content("\n" + "=".repeat(60) + "\n\n"), MediaType.APPLICATION_JSON));
logger.info("最终报告已完整输出");
} else {
logger.warn("未能提取到 Planner 最终报告");
emitter.send(SseEmitter.event().name("message")
.data(SseMessage.content("⚠️ 多 Agent 流程已完成,但未能生成最终报告。"), MediaType.APPLICATION_JSON));
}
emitter.send(SseEmitter.event().name("message").data(SseMessage.done(), MediaType.APPLICATION_JSON));
emitter.complete();
logger.info("AI Ops 多 Agent 编排完成");
} catch (Exception e) {
logger.error("AI Ops 多 Agent 协作失败", e);
try {
emitter.send(SseEmitter.event().name("message")
.data(SseMessage.error("AI Ops 流程失败: " + e.getMessage()), MediaType.APPLICATION_JSON));
} catch (IOException ex) {
logger.error("发送错误消息失败", ex);
}
emitter.completeWithError(e);
}
});
return emitter;
}
/**
* 获取会话信息
*/
@GetMapping("/chat/session/{sessionId}")
public ResponseEntity<ApiResponse<SessionInfoResponse>> getSessionInfo(@PathVariable String sessionId) {
try {
logger.info("收到获取会话信息请求 - SessionId: {}", sessionId);
SessionInfo session = sessions.get(sessionId);
if (session != null) {
SessionInfoResponse response = new SessionInfoResponse();
response.setSessionId(sessionId);
response.setMessagePairCount(session.getMessagePairCount());
response.setCreateTime(session.createTime);
return ResponseEntity.ok(ApiResponse.success(response));
} else {
return ResponseEntity.ok(ApiResponse.error("会话不存在"));
}
} catch (Exception e) {
logger.error("获取会话信息失败", e);
return ResponseEntity.ok(ApiResponse.error(e.getMessage()));
}
}
// ==================== 辅助方法 ====================
private SessionInfo getOrCreateSession(String sessionId) {
if (sessionId == null || sessionId.isEmpty()) {
sessionId = UUID.randomUUID().toString();
}
return sessions.computeIfAbsent(sessionId, SessionInfo::new);
}
// ==================== 内部类 ====================
/**
* 会话信息
* 管理单个会话的历史消息,支持自动清理和线程安全
*/
private static class SessionInfo {
private final String sessionId;
// 存储历史消息对:[{"role": "user", "content": "..."}, {"role": "assistant", "content": "..."}]
private final List<Map<String, String>> messageHistory;
private final long createTime;
private final ReentrantLock lock;
public SessionInfo(String sessionId) {
this.sessionId = sessionId;
this.messageHistory = new ArrayList<>();
this.createTime = System.currentTimeMillis();
this.lock = new ReentrantLock();
}
/**
* 添加一对消息(用户问题 + AI回复)
* 自动管理历史消息窗口大小
*/
public void addMessage(String userQuestion, String aiAnswer) {
lock.lock();
try {
// 添加用户消息
Map<String, String> userMsg = new HashMap<>();
userMsg.put("role", "user");
userMsg.put("content", userQuestion);
messageHistory.add(userMsg);
// 添加AI回复
Map<String, String> assistantMsg = new HashMap<>();
assistantMsg.put("role", "assistant");
assistantMsg.put("content", aiAnswer);
messageHistory.add(assistantMsg);
// 自动清理:保持最多 MAX_WINDOW_SIZE 对消息
// 每对消息包含2条记录(user + assistant)
int maxMessages = MAX_WINDOW_SIZE * 2;
while (messageHistory.size() > maxMessages) {
// 成对删除最旧的消息(删除前2条)
messageHistory.remove(0); // 删除最旧的用户消息
if (!messageHistory.isEmpty()) {
messageHistory.remove(0); // 删除对应的AI回复
}
}
logger.debug("会话 {} 更新历史消息,当前消息对数: {}",
sessionId, messageHistory.size() / 2);
} finally {
lock.unlock();
}
}
/**
* 获取历史消息(线程安全)
* 返回副本以避免并发修改
*/
public List<Map<String, String>> getHistory() {
lock.lock();
try {
return new ArrayList<>(messageHistory);
} finally {
lock.unlock();
}
}
/**
* 清空历史消息
*/
public void clearHistory() {
lock.lock();
try {
messageHistory.clear();
logger.info("会话 {} 历史消息已清空", sessionId);
} finally {
lock.unlock();
}
}
/**
* 获取当前消息对数
*/
public int getMessagePairCount() {
lock.lock();
try {
return messageHistory.size() / 2;
} finally {
lock.unlock();
}
}
}
/**
* 聊天请求
*/
@Setter
@Getter
public static class ChatRequest {
@com.fasterxml.jackson.annotation.JsonProperty(value = "Id")
@com.fasterxml.jackson.annotation.JsonAlias({"id", "ID"})
private String Id;
@com.fasterxml.jackson.annotation.JsonProperty(value = "Question")
@com.fasterxml.jackson.annotation.JsonAlias({"question", "QUESTION"})
private String Question;
}
/**
* 清空会话请求
*/
@Setter
@Getter
public static class ClearRequest {
@com.fasterxml.jackson.annotation.JsonProperty(value = "Id")
@com.fasterxml.jackson.annotation.JsonAlias({"id", "ID"})
private String Id;
}
// ==================== 内部类 ====================
/**
* 会话信息响应
*/
@Setter
@Getter
public static class SessionInfoResponse {
private String sessionId;
private int messagePairCount;
private long createTime;
}
/**
* 统一聊天响应格式
* 适用于所有普通返回模式的对话接口
*/
@Setter
@Getter
public static class ChatResponse {
private boolean success;
private String answer;
private String errorMessage;
public static ChatResponse success(String answer) {
ChatResponse response = new ChatResponse();
response.setSuccess(true);
response.setAnswer(answer);
return response;
}
public static ChatResponse error(String errorMessage) {
ChatResponse response = new ChatResponse();
response.setSuccess(false);
response.setErrorMessage(errorMessage);
return response;
}
}
/**
* 统一 SSE 流式消息格式
* 适用于所有 SSE 流式返回模式的对话接口
*/
@Setter
@Getter
public static class SseMessage {
private String type; // content: 内容块, error: 错误, done: 完成
private String data;
public static SseMessage content(String data) {
SseMessage message = new SseMessage();
message.setType("content");
message.setData(data);
return message;
}
public static SseMessage error(String errorMessage) {
SseMessage message = new SseMessage();
message.setType("error");
message.setData(errorMessage);
return message;
}
public static SseMessage done() {
SseMessage message = new SseMessage();
message.setType("done");
message.setData(null);
return message;
}
}
@Getter
@Setter
public static class ApiResponse<T> {
private int code;
private String message;
private T data;
public static <T> ApiResponse<T> success(T data) {
ApiResponse<T> response = new ApiResponse<>();
response.setCode(200);
response.setMessage("success");
response.setData(data);
return response;
}
public static <T> ApiResponse<T> error(String message) {
ApiResponse<T> response = new ApiResponse<>();
response.setCode(500);
response.setMessage(message);
return response;
}
}
}
@@ -0,0 +1,154 @@
package org.example.controller;
import org.example.config.FileUploadConfig;
import org.example.dto.FileUploadRes;
import org.example.service.VectorIndexService;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.http.HttpStatus;
import org.springframework.http.ResponseEntity;
import org.springframework.web.bind.annotation.PostMapping;
import org.springframework.web.bind.annotation.RequestParam;
import org.springframework.web.bind.annotation.RestController;
import org.springframework.web.multipart.MultipartFile;
import java.io.IOException;
import java.nio.file.Files;
import java.nio.file.Path;
import java.nio.file.Paths;
import java.util.Arrays;
import java.util.List;
@RestController
public class FileUploadController {
private static final Logger logger = LoggerFactory.getLogger(FileUploadController.class);
@Autowired
private FileUploadConfig fileUploadConfig;
@Autowired
private VectorIndexService vectorIndexService;
@PostMapping(value = "/api/upload", consumes = "multipart/form-data")
public ResponseEntity<?> upload(@RequestParam("file") MultipartFile file) {
if (file.isEmpty()) {
return ResponseEntity.badRequest().body("文件不能为空");
}
String originalFilename = file.getOriginalFilename();
if (originalFilename == null || originalFilename.isEmpty()) {
return ResponseEntity.badRequest().body("文件名不能为空");
}
String fileExtension = getFileExtension(originalFilename);
if (!isAllowedExtension(fileExtension)) {
return ResponseEntity.status(HttpStatus.BAD_REQUEST)
.body("不支持的文件格式,仅支持: " + fileUploadConfig.getAllowedExtensions());
}
try {
String uploadPath = fileUploadConfig.getPath();
Path uploadDir = Paths.get(uploadPath).normalize();
if (!Files.exists(uploadDir)) {
Files.createDirectories(uploadDir);
}
// 使用原始文件名,而不是UUID,以便实现基于文件名的去重
Path filePath = uploadDir.resolve(originalFilename).normalize();
// 如果文件已存在,先删除旧文件(实现覆盖更新)
if (Files.exists(filePath)) {
logger.info("文件已存在,将覆盖: {}", filePath);
Files.delete(filePath);
}
Files.copy(file.getInputStream(), filePath);
logger.info("文件上传成功: {}", filePath);
// 文件上传成功后,自动调用向量索引服务
try {
logger.info("开始为上传文件创建向量索引: {}", filePath);
vectorIndexService.indexSingleFile(filePath.toString());
logger.info("向量索引创建成功: {}", filePath);
} catch (Exception e) {
logger.error("向量索引创建失败: {}, 错误: {}", filePath, e.getMessage(), e);
// 注意:即使索引失败,文件上传仍然成功,只是记录错误日志
// 可以根据业务需求决定是否要删除文件或返回错误
}
FileUploadRes response = new FileUploadRes(
originalFilename,
filePath.toString(),
file.getSize()
);
// 使用统一的API响应格式
ApiResponse<FileUploadRes> apiResponse = new ApiResponse<>();
apiResponse.setCode(200);
apiResponse.setMessage("success");
apiResponse.setData(response);
return ResponseEntity.ok(apiResponse);
} catch (IOException e) {
ApiResponse<String> errorResponse = new ApiResponse<>();
errorResponse.setCode(500);
errorResponse.setMessage("文件上传失败: " + e.getMessage());
return ResponseEntity.status(HttpStatus.INTERNAL_SERVER_ERROR)
.body(errorResponse);
}
}
/**
* 统一 API 响应格式
*/
public static class ApiResponse<T> {
private int code;
private String message;
private T data;
public int getCode() {
return code;
}
public void setCode(int code) {
this.code = code;
}
public String getMessage() {
return message;
}
public void setMessage(String message) {
this.message = message;
}
public T getData() {
return data;
}
public void setData(T data) {
this.data = data;
}
}
private String getFileExtension(String filename) {
int lastIndexOf = filename.lastIndexOf(".");
if (lastIndexOf == -1) {
return "";
}
return filename.substring(lastIndexOf + 1).toLowerCase();
}
private boolean isAllowedExtension(String extension) {
String allowedExtensions = fileUploadConfig.getAllowedExtensions();
if (allowedExtensions == null || allowedExtensions.isEmpty()) {
return false;
}
List<String> allowedList = Arrays.asList(allowedExtensions.split(","));
return allowedList.contains(extension.toLowerCase());
}
}
@@ -0,0 +1,50 @@
package org.example.controller;
import io.milvus.client.MilvusServiceClient;
import io.milvus.grpc.ShowCollectionsResponse;
import io.milvus.param.R;
import io.milvus.param.collection.ShowCollectionsParam;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.http.ResponseEntity;
import org.springframework.web.bind.annotation.*;
import java.util.HashMap;
import java.util.Map;
/**
* Milvus 测试控制器
* 用于测试数据库连接和数据读取
*/
@RestController
@RequestMapping("/milvus")
public class MilvusCheckController {
@Autowired
private MilvusServiceClient milvusClient;
/**
* 简单的健康检查
*/
@GetMapping("/health")
public ResponseEntity<Map<String, Object>> simpleHealth() {
Map<String, Object> result = new HashMap<>();
try {
R<ShowCollectionsResponse> response = milvusClient.showCollections(
ShowCollectionsParam.newBuilder().build()
);
if (response.getStatus() == 0) {
result.put("message", "ok");
result.put("collections", response.getData().getCollectionNamesList());
return ResponseEntity.ok(result);
} else {
result.put("message", response.getMessage());
return ResponseEntity.status(503).body(result);
}
} catch (Exception e) {
result.put("error", e.getMessage());
return ResponseEntity.status(503).body(result);
}
}
}
@@ -0,0 +1,15 @@
package org.example.dto;
import lombok.Data;
/**
* AIOps 请求 DTO
*/
@Data
public class AIOpsRequest {
/**
* 用户请求描述
*/
private String userRequest;
}
@@ -0,0 +1,59 @@
package org.example.dto;
import lombok.Getter;
import lombok.Setter;
/**
* 文档分片
*/
@Setter
@Getter
public class DocumentChunk {
// Getters and Setters
/**
* 分片内容
*/
private String content;
/**
* 分片在原文档中的起始位置
*/
private int startIndex;
/**
* 分片在原文档中的结束位置
*/
private int endIndex;
/**
* 分片序号(从0开始)
*/
private int chunkIndex;
/**
* 分片标题或上下文信息
*/
private String title;
public DocumentChunk() {
}
public DocumentChunk(String content, int startIndex, int endIndex, int chunkIndex) {
this.content = content;
this.startIndex = startIndex;
this.endIndex = endIndex;
this.chunkIndex = chunkIndex;
}
@Override
public String toString() {
return "DocumentChunk{" +
"chunkIndex=" + chunkIndex +
", title='" + title + '\'' +
", contentLength=" + (content != null ? content.length() : 0) +
", startIndex=" + startIndex +
", endIndex=" + endIndex +
'}';
}
}
@@ -0,0 +1,23 @@
package org.example.dto;
import lombok.Getter;
import lombok.Setter;
@Setter
@Getter
public class FileUploadRes {
private String fileName;
private String filePath;
private Long fileSize;
public FileUploadRes() {
}
public FileUploadRes(String fileName, String filePath, Long fileSize) {
this.fileName = fileName;
this.filePath = filePath;
this.fileSize = fileSize;
}
}
@@ -0,0 +1,278 @@
package org.example.service;
import com.alibaba.cloud.ai.dashscope.chat.DashScopeChatModel;
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;
import com.alibaba.cloud.ai.graph.exception.GraphRunnerException;
import org.example.agent.tool.DateTimeTools;
import org.example.agent.tool.InternalDocsTools;
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.messages.AssistantMessage;
import org.springframework.ai.tool.ToolCallback;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.stereotype.Service;
import java.util.List;
import java.util.Optional;
/**
* AI Ops 智能运维服务
* 负责多 Agent 协作的告警分析流程
*/
@Service
public class AiOpsService {
private static final Logger logger = LoggerFactory.getLogger(AiOpsService.class);
@Autowired
private DateTimeTools dateTimeTools;
@Autowired
private InternalDocsTools internalDocsTools;
@Autowired
private QueryMetricsTools queryMetricsTools;
@Autowired(required = false) // Mock 模式下才注册
private QueryLogsTools queryLogsTools;
/**
* 执行 AI Ops 告警分析流程
*
* @param chatModel 大模型实例
* @param toolCallbacks 工具回调数组
* @return 分析结果状态
* @throws GraphRunnerException 如果 Agent 执行失败
*/
public Optional<OverAllState> executeAiOpsAnalysis(DashScopeChatModel chatModel, ToolCallback[] toolCallbacks) throws GraphRunnerException {
logger.info("开始执行 AI Ops 多 Agent 协作流程");
// 构建 Planner 和 Executor Agent
ReactAgent plannerAgent = buildPlannerAgent(chatModel, toolCallbacks);
ReactAgent executorAgent = buildExecutorAgent(chatModel, toolCallbacks);
// 构建 Supervisor Agent
SupervisorAgent supervisorAgent = SupervisorAgent.builder()
.name("ai_ops_supervisor")
.description("负责调度 Planner 与 Executor 的多 Agent 控制器")
.model(chatModel)
.systemPrompt(buildSupervisorSystemPrompt())
.subAgents(List.of(plannerAgent, executorAgent))
.build();
String taskPrompt = "你是企业级 SRE,接到了自动化告警排查任务。请结合工具调用,执行**规划→执行→再规划**的闭环,并最终按照固定模板输出《告警分析报告》。禁止编造虚假数据,如连续多次查询失败需诚实反馈无法完成的原因。";
logger.info("调用 Supervisor Agent 开始编排...");
return supervisorAgent.invoke(taskPrompt);
}
/**
* 从执行结果中提取最终报告文本
*
* @param state 执行状态
* @return 报告文本(如果存在)
*/
public Optional<String> extractFinalReport(OverAllState state) {
logger.info("开始提取最终报告...");
// 提取 Planner 最终输出(包含完整的告警分析报告)
Optional<AssistantMessage> plannerFinalOutput = state.value("planner_plan")
.filter(AssistantMessage.class::isInstance)
.map(AssistantMessage.class::cast);
if (plannerFinalOutput.isPresent()) {
String reportText = plannerFinalOutput.get().getText();
logger.info("成功提取到 Planner 最终报告,长度: {}", reportText.length());
return Optional.of(reportText);
} else {
logger.warn("未能提取到 Planner 最终报告");
return Optional.empty();
}
}
/**
* 构建 Planner Agent
*/
private ReactAgent buildPlannerAgent(DashScopeChatModel chatModel, ToolCallback[] toolCallbacks) {
return ReactAgent.builder()
.name("planner_agent")
.description("负责拆解告警、规划与再规划步骤")
.model(chatModel)
.systemPrompt(buildPlannerPrompt())
.methodTools(buildMethodToolsArray())
.tools(toolCallbacks)
.outputKey("planner_plan")
.build();
}
/**
* 构建 Executor Agent
*/
private ReactAgent buildExecutorAgent(DashScopeChatModel chatModel, ToolCallback[] toolCallbacks) {
return ReactAgent.builder()
.name("executor_agent")
.description("负责执行 Planner 的首个步骤并及时反馈")
.model(chatModel)
.systemPrompt(buildExecutorPrompt())
.methodTools(buildMethodToolsArray())
.tools(toolCallbacks)
.outputKey("executor_feedback")
.build();
}
/**
* 动态构建方法工具数组
* 根据 cls.mock-enabled 决定是否包含 QueryLogsTools
*/
private Object[] buildMethodToolsArray() {
if (queryLogsTools != null) {
// Mock 模式:包含 QueryLogsTools
return new Object[]{dateTimeTools, internalDocsTools, queryMetricsTools, queryLogsTools};
} else {
// 真实模式:不包含 QueryLogsTools(由 MCP 提供日志查询功能)
return new Object[]{dateTimeTools, internalDocsTools, queryMetricsTools};
}
}
/**
* 构建 Planner Agent 系统提示词
*/
private String buildPlannerPrompt() {
return """
你是 Planner Agent,同时承担 Replanner 角色,负责:
1. 读取当前输入任务 {input} 以及 Executor 的最近反馈 {executor_feedback}。
2. 分析 Prometheus 告警、日志、内部文档等信息,制定可执行的下一步步骤。
3. 在执行阶段,输出 JSON,包含 decision (PLAN|EXECUTE|FINISH)、step 描述、预期要调用的工具、以及必要的上下文。
4. 调用任何腾讯云日志/主题相关工具时,region 参数必须使用连字符格式(如 ap-guangzhou),若不确定请省略以使用默认值。
5. 严格禁止编造数据,只能引用工具返回的真实内容;如果连续 3 次调用同一工具仍失败或返回空结果,需停止该方向并在最终报告的结论部分说明"无法完成"的原因。
## 最终报告输出要求(CRITICAL)
当 decision=FINISH 时,你必须:
1. **不要输出 JSON 格式**
2. **直接输出完整的 Markdown 格式报告文本**
3. **报告必须严格遵循以下模板**:
```
# 告警分析报告
---
## 📋 活跃告警清单
| 告警名称 | 级别 | 目标服务 | 首次触发时间 | 最新触发时间 | 状态 |
|---------|------|----------|-------------|-------------|------|
| [告警1名称] | [级别] | [服务名] | [时间] | [时间] | 活跃 |
| [告警2名称] | [级别] | [服务名] | [时间] | [时间] | 活跃 |
---
## 🔍 告警根因分析1 - [告警名称]
### 告警详情
- **告警级别**: [级别]
- **受影响服务**: [服务名]
- **持续时间**: [X分钟]
### 症状描述
[根据监控指标描述症状]
### 日志证据
[引用查询到的关键日志]
### 根因结论
[基于证据得出的根本原因]
---
## 🛠️ 处理方案执行1 - [告警名称]
### 已执行的排查步骤
1. [步骤1]
2. [步骤2]
### 处理建议
[给出具体的处理建议]
### 预期效果
[说明预期的效果]
---
## 🔍 告警根因分析2 - [告警名称]
[如果有第2个告警,重复上述格式]
---
## 📊 结论
### 整体评估
[总结所有告警的整体情况]
### 关键发现
- [发现1]
- [发现2]
### 后续建议
1. [建议1]
2. [建议2]
### 风险评估
[评估当前风险等级和影响范围]
```
**重要提醒**:
- 最终输出必须是纯 Markdown 文本,不要包含 JSON 结构
- 不要使用 "finalReport": "..." 这样的格式
- 直接从 "# 告警分析报告" 开始输出
- 所有内容必须基于工具查询的真实数据,严禁编造
- 如果某个步骤失败,在结论中如实说明,不要跳过
""";
}
/**
* 构建 Executor Agent 系统提示词
*/
private String buildExecutorPrompt() {
return """
你是 Executor Agent,负责读取 Planner 最新输出 {planner_plan},只执行其中的第一步。
- 确认步骤所需的工具与参数,尤其是 region 参数要使用连字符格式(ap-guangzhou);若 Planner 未给出则使用默认区域。
- 调用相应的工具并收集结果,如工具返回错误或空数据,需要将失败原因、请求参数一并记录,并停止进一步调用该工具(同一工具失败达到 3 次时应直接返回 FAILED)。
- 将日志、指标、文档等证据整理成结构化摘要,标注对应的告警名称或资源,方便 Planner 填充"告警根因分析 / 处理方案执行"章节。
- 以 JSON 形式返回执行状态、证据以及给 Planner 的建议,写入 executor_feedback,严禁编造未实际查询到的内容。
输出示例:
{
"status": "SUCCESS",
"summary": "近1小时未见 error 日志,仅有 info",
"evidence": "...",
"nextHint": "建议转向高占用进程"
}
""";
}
/**
* 构建 Supervisor Agent 系统提示词
*/
private String buildSupervisorSystemPrompt() {
return """
你是 AI Ops Supervisor,负责调度 planner_agent 与 executor_agent:
1. 当需要拆解任务或重新制定策略时,调用 planner_agent。
2. 当 planner_agent 输出 decision=EXECUTE 时,调用 executor_agent 执行第一步。
3. 根据 executor_agent 的反馈,评估是否需要再次调用 planner_agent,直到 decision=FINISH。
4. FINISH 后,确保向最终用户输出完整的《告警分析报告》,格式必须严格为:
告警分析报告\n---\n# 告警处理详情\n## 活跃告警清单\n## 告警根因分析N\n## 处理方案执行N\n## 结论。
5. 若步骤涉及腾讯云日志/主题工具,请确保使用连字符区域 ID(ap-guangzhou 等),或省略 region 以采用默认值。
6. 如果发现 Planner/Executor 在同一方向连续 3 次调用工具仍失败或没有数据,必须终止流程,直接输出"任务无法完成"的报告,明确告知失败原因,严禁凭空编造结果。
只允许在 planner_agent、executor_agent 与 FINISH 之间做出选择。
""";
}
}
@@ -0,0 +1,180 @@
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;
import org.example.agent.tool.InternalDocsTools;
import org.example.agent.tool.QueryLogsTools;
import org.example.agent.tool.QueryMetricsTools;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
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;
import java.util.Map;
/**
* 聊天服务
* 封装 ReactAgent 对话的公共逻辑,包括模型创建、系统提示词构建、Agent 配置等
*/
@Service
public class ChatService {
private static final Logger logger = LoggerFactory.getLogger(ChatService.class);
@Autowired
private InternalDocsTools internalDocsTools;
@Autowired
private DateTimeTools dateTimeTools;
@Autowired
private QueryMetricsTools queryMetricsTools;
@Autowired(required = false) // Mock 模式下才注册,所以设置为 optional,真实环境通过mcp配置注入
private QueryLogsTools queryLogsTools;
@Autowired
private ToolCallbackProvider tools;
@Value("${spring.ai.dashscope.api-key}")
private String dashScopeApiKey;
/**
* 创建 DashScope API 实例
*/
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);
}
/**
* 构建系统提示词(包含历史消息)
* @param history 历史消息列表
* @return 完整的系统提示词
*/
public String buildSystemPrompt(List<Map<String, String>> history) {
StringBuilder systemPromptBuilder = new StringBuilder();
// 基础系统提示
systemPromptBuilder.append("你是一个专业的智能助手,可以获取当前时间、查询天气信息、搜索内部文档知识库,以及查询 Prometheus 告警信息。\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");
if ("user".equals(role)) {
systemPromptBuilder.append("用户: ").append(content).append("\n");
} else if ("assistant".equals(role)) {
systemPromptBuilder.append("助手: ").append(content).append("\n");
}
}
systemPromptBuilder.append("--- 对话历史结束 ---\n\n");
}
systemPromptBuilder.append("请基于以上对话历史,回答用户的新问题。");
return systemPromptBuilder.toString();
}
/**
* 动态构建方法工具数组
* 根据 cls.mock-enabled 决定是否包含 QueryLogsTools
*/
public Object[] buildMethodToolsArray() {
if (queryLogsTools != null) {
// Mock 模式:包含 QueryLogsTools
return new Object[]{dateTimeTools, internalDocsTools, queryMetricsTools, queryLogsTools};
} else {
// 真实模式:不包含 QueryLogsTools(由 MCP 提供日志查询功能)
return new Object[]{dateTimeTools, internalDocsTools, queryMetricsTools};
}
}
/**
* 获取工具回调列表,mcp服务提供的工具
*/
public ToolCallback[] getToolCallbacks() {
return tools.getToolCallbacks();
}
/**
* 记录可用工具列表:mcp服务提供的工具
*/
public void logAvailableTools() {
ToolCallback[] toolCallbacks = tools.getToolCallbacks();
logger.info("可用工具列表:");
for (ToolCallback toolCallback : toolCallbacks) {
logger.info(">>> {}", toolCallback.getToolDefinition().name());
}
}
/**
* 创建 ReactAgent
* @param chatModel 聊天模型
* @param systemPrompt 系统提示词
* @return 配置好的 ReactAgent
*/
public ReactAgent createReactAgent(DashScopeChatModel chatModel, String systemPrompt) {
return ReactAgent.builder()
.name("intelligent_assistant")
.model(chatModel)
.systemPrompt(systemPrompt)
.methodTools(buildMethodToolsArray())
.tools(getToolCallbacks())
.build();
}
/**
* 执行 ReactAgent 对话(非流式)
* @param agent ReactAgent 实例
* @param question 用户问题
* @return AI 回复
*/
public String executeChat(ReactAgent agent, String question) throws GraphRunnerException {
logger.info("执行 ReactAgent.call() - 自动处理工具调用");
var response = agent.call(question);
String answer = response.getText();
logger.info("ReactAgent 对话完成,答案长度: {}", answer.length());
return answer;
}
}
@@ -0,0 +1,229 @@
package org.example.service;
import org.example.config.DocumentChunkConfig;
import org.example.dto.DocumentChunk;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.stereotype.Service;
import java.util.ArrayList;
import java.util.List;
import java.util.regex.Matcher;
import java.util.regex.Pattern;
/**
* 文档分片服务
* 负责将长文档切分为多个有语义完整性的小片段
*/
@Service
public class DocumentChunkService {
private static final Logger logger = LoggerFactory.getLogger(DocumentChunkService.class);
@Autowired
private DocumentChunkConfig chunkConfig;
/**
* 智能分片文档
* 优先按照标题、段落边界进行分割,保持语义完整性
*
* @param content 文档内容
* @param filePath 文件路径(用于日志)
* @return 文档分片列表
*/
public List<DocumentChunk> chunkDocument(String content, String filePath) {
List<DocumentChunk> chunks = new ArrayList<>();
if (content == null || content.trim().isEmpty()) {
logger.warn("文档内容为空: {}", filePath);
return chunks;
}
// 1. 首先尝试按标题分割(Markdown格式)
List<Section> sections = splitByHeadings(content);
// 2. 对每个章节进行进一步分片
int globalChunkIndex = 0;
for (Section section : sections) {
List<DocumentChunk> sectionChunks = chunkSection(section, globalChunkIndex);
chunks.addAll(sectionChunks);
globalChunkIndex += sectionChunks.size();
}
logger.info("文档分片完成: {} -> {} 个分片", filePath, chunks.size());
return chunks;
}
/**
* 按照 Markdown 标题分割文档
*/
private List<Section> splitByHeadings(String content) {
List<Section> sections = new ArrayList<>();
// 匹配 Markdown 标题:# 标题, ## 标题, ### 标题等
Pattern headingPattern = Pattern.compile("^(#{1,6})\\s+(.+)$", Pattern.MULTILINE);
Matcher matcher = headingPattern.matcher(content);
int lastEnd = 0;
String currentTitle = null;
while (matcher.find()) {
// 保存上一个章节
if (lastEnd < matcher.start()) {
String sectionContent = content.substring(lastEnd, matcher.start()).trim();
if (!sectionContent.isEmpty()) {
sections.add(new Section(currentTitle, sectionContent, lastEnd));
}
}
// 更新当前标题
currentTitle = matcher.group(2).trim();
lastEnd = matcher.start();
}
// 添加最后一个章节
if (lastEnd < content.length()) {
String sectionContent = content.substring(lastEnd).trim();
if (!sectionContent.isEmpty()) {
sections.add(new Section(currentTitle, sectionContent, lastEnd));
}
}
// 如果没有找到任何标题,将整个文档作为一个章节
if (sections.isEmpty()) {
sections.add(new Section(null, content, 0));
}
return sections;
}
/**
* 对单个章节进行分片
*/
private List<DocumentChunk> chunkSection(Section section, int startChunkIndex) {
List<DocumentChunk> chunks = new ArrayList<>();
String content = section.content;
String title = section.title;
// 如果章节内容小于最大尺寸,直接作为一个分片
if (content.length() <= chunkConfig.getMaxSize()) {
DocumentChunk chunk = new DocumentChunk(
content,
section.startIndex,
section.startIndex + content.length(),
startChunkIndex
);
chunk.setTitle(title);
chunks.add(chunk);
return chunks;
}
// 章节内容较长,需要进一步分片
// 优先在段落边界分割
List<String> paragraphs = splitByParagraphs(content);
StringBuilder currentChunk = new StringBuilder();
int currentStartIndex = section.startIndex;
int chunkIndex = startChunkIndex;
for (String paragraph : paragraphs) {
// 如果当前分片加上新段落超过最大尺寸
if (currentChunk.length() > 0 &&
currentChunk.length() + paragraph.length() > chunkConfig.getMaxSize()) {
// 保存当前分片
String chunkContent = currentChunk.toString().trim();
DocumentChunk chunk = new DocumentChunk(
chunkContent,
currentStartIndex,
currentStartIndex + chunkContent.length(),
chunkIndex++
);
chunk.setTitle(title);
chunks.add(chunk);
// 开始新分片,包含重叠部分
String overlap = getOverlapText(chunkContent);
currentChunk = new StringBuilder(overlap);
currentStartIndex = currentStartIndex + chunkContent.length() - overlap.length();
}
currentChunk.append(paragraph).append("\n\n");
}
// 保存最后一个分片
if (currentChunk.length() > 0) {
String chunkContent = currentChunk.toString().trim();
DocumentChunk chunk = new DocumentChunk(
chunkContent,
currentStartIndex,
currentStartIndex + chunkContent.length(),
chunkIndex
);
chunk.setTitle(title);
chunks.add(chunk);
}
return chunks;
}
/**
* 按段落分割文本
*/
private List<String> splitByParagraphs(String content) {
List<String> paragraphs = new ArrayList<>();
// 按双换行符分割段落
String[] parts = content.split("\n\n+");
for (String part : parts) {
String trimmed = part.trim();
if (!trimmed.isEmpty()) {
paragraphs.add(trimmed);
}
}
return paragraphs;
}
/**
* 获取重叠文本
* 从文本末尾提取指定长度的内容作为下一个分片的开头
*/
private String getOverlapText(String text) {
int overlapSize = Math.min(chunkConfig.getOverlap(), text.length());
if (overlapSize <= 0) {
return "";
}
// 从末尾提取重叠内容
String overlap = text.substring(text.length() - overlapSize);
// 尝试在句子边界截断(查找最后一个句号、问号、感叹号)
int lastSentenceEnd = Math.max(
overlap.lastIndexOf('。'),
Math.max(overlap.lastIndexOf('?'), overlap.lastIndexOf('!'))
);
if (lastSentenceEnd > overlapSize / 2) {
return overlap.substring(lastSentenceEnd + 1).trim();
}
return overlap.trim();
}
/**
* 章节数据类
*/
private static class Section {
String title;
String content;
int startIndex;
Section(String title, String content, int startIndex) {
this.title = title;
this.content = content;
this.startIndex = startIndex;
}
}
}
@@ -0,0 +1,233 @@
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.beans.factory.annotation.Autowired;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.stereotype.Service;
import jakarta.annotation.PostConstruct;
import java.util.ArrayList;
import java.util.List;
import java.util.Map;
/**
* RAG (Retrieval-Augmented Generation) 服务
* 结合向量检索和大语言模型生成答案
*/
@Service
public class RagService {
private static final Logger logger = LoggerFactory.getLogger(RagService.class);
@Autowired
private VectorSearchService vectorSearchService;
@Value("${dashscope.api.key}")
private String apiKey;
@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);
}
/**
* 流式处理用户问题(不带历史消息)
*
* @param question 用户问题
* @param callback 流式回调接口
*/
public void queryStream(String question, StreamCallback callback) {
queryStream(question, new ArrayList<>(), callback);
}
/**
* 流式处理用户问题(带历史消息)
*
* @param question 用户问题
* @param history 历史消息列表,格式:[{"role": "user", "content": "..."}, {"role": "assistant", "content": "..."}]
* @param callback 流式回调接口
*/
public void queryStream(String question, List<Map<String, String>> history, StreamCallback callback) {
try {
logger.info("收到 RAG 流式查询: {}", question);
// 1. 从向量数据库检索相关文档
List<VectorSearchService.SearchResult> searchResults =
vectorSearchService.searchSimilarDocuments(question, topK);
// 发送检索结果
callback.onSearchResults(searchResults);
if (searchResults.isEmpty()) {
logger.warn("未找到相关文档");
callback.onComplete("抱歉,我在知识库中没有找到相关信息来回答您的问题。", "");
return;
}
// 2. 构建上下文和提示词
String context = buildContext(searchResults);
String prompt = buildPrompt(question, context);
// 3. 流式调用大语言模型(传入历史消息)
generateAnswerStream(prompt, history, callback);
} catch (Exception e) {
logger.error("RAG 流式查询失败", e);
callback.onError(e);
}
}
/**
* 构建上下文
*/
private String buildContext(List<VectorSearchService.SearchResult> searchResults) {
StringBuilder context = new StringBuilder();
for (int i = 0; i < searchResults.size(); i++) {
VectorSearchService.SearchResult result = searchResults.get(i);
context.append("【参考资料 ").append(i + 1).append("】\n");
context.append(result.getContent()).append("\n\n");
}
return context.toString();
}
/**
* 构建提示词
*/
private String buildPrompt(String question, String context) {
return String.format(
"你是一个专业的AI助手。请根据以下参考资料回答用户的问题。\n\n" +
"参考资料:\n%s\n" +
"用户问题:%s\n\n" +
"请基于上述参考资料给出准确、详细的回答。如果参考资料中没有相关信息,请明确说明。",
context, question
);
}
/**
* 生成答案(流式)
*
* @param prompt 当前问题的提示词
* @param history 历史消息列表
* @param callback 流式回调接口
*/
private void generateAnswerStream(String prompt, List<Map<String, String>> history, StreamCallback callback)
throws NoApiKeyException, ApiException, InputRequiredException {
// 构建消息列表:历史消息 + 当前问题
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());
} else if ("assistant".equals(role)) {
messages.add(Message.builder()
.role(Role.ASSISTANT.getValue())
.content(content)
.build());
}
}
// 添加当前用户问题
Message userMsg = Message.builder()
.role(Role.USER.getValue())
.content(prompt)
.build();
messages.add(userMsg);
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();
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();
if (content != null && !content.isEmpty()) {
logger.debug("收到AI模型内容块: {}", content);
// 对于 thinking 模型,content 可能包含思考过程和最终答案
// 这里我们将所有内容都作为答案返回
finalContent.append(content);
callback.onContentChunk(content);
logger.debug("已调用 onContentChunk 回调");
} else {
logger.debug("收到空内容块,跳过");
}
}
});
logger.info("AI模型流式响应完成,总内容长度: {}", finalContent.length());
callback.onComplete(finalContent.toString(), reasoningContent.toString());
logger.info("已调用 onComplete 回调");
}
/**
* 流式回调接口
*/
public interface StreamCallback {
void onSearchResults(List<VectorSearchService.SearchResult> results);
void onReasoningChunk(String chunk);
void onContentChunk(String chunk);
void onComplete(String fullContent, String fullReasoning);
void onError(Exception e);
}
}
@@ -0,0 +1,247 @@
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.stereotype.Service;
import jakarta.annotation.PostConstruct;
import java.util.ArrayList;
import java.util.Collections;
import java.util.List;
/**
* 向量嵌入服务
* 使用阿里云 DashScope Text Embedding API
*/
@Service
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);
}
/**
* 生成向量嵌入
* 调用阿里云 DashScope Text Embedding API
*
* @param content 文本内容
* @return 向量嵌入(浮点数列表)
*/
public List<Float> generateEmbedding(String content) {
try {
if (content == null || content.trim().isEmpty()) {
logger.warn("内容为空,无法生成向量");
throw new IllegalArgumentException("内容不能为空");
}
logger.debug("开始生成向量嵌入, 内容长度: {} 字符", content.length());
// 确保 API Key 已设置(防止被其他地方覆盖)
if (Constants.apiKey == null || Constants.apiKey.isEmpty()) {
logger.warn("检测到 Constants.apiKey 为空,重新设置");
Constants.apiKey = apiKey;
}
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("成功生成向量嵌入, 内容长度: {} 字符, 向量维度: {}",
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()) {
logger.warn("内容列表为空,无法生成向量");
return Collections.emptyList();
}
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();
// 调用 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());
}
embeddings.add(embedding);
}
logger.info("成功批量生成向量嵌入, 数量: {}, 维度: {}",
embeddings.size(),
embeddings.isEmpty() ? 0 : embeddings.get(0).size());
return embeddings;
} 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);
}
}
/**
* 生成查询向量
*
* @param query 查询文本
* @return 向量嵌入
*/
public List<Float> generateQueryVector(String query) {
return generateEmbedding(query);
}
/**
* 计算两个向量的余弦相似度
*
* @param vector1 向量1
* @param vector2 向量2
* @return 余弦相似度 [-1, 1]
*/
public float calculateCosineSimilarity(List<Float> vector1, List<Float> vector2) {
if (vector1.size() != vector2.size()) {
throw new IllegalArgumentException("向量维度不匹配");
}
float dotProduct = 0.0f;
float norm1 = 0.0f;
float norm2 = 0.0f;
for (int i = 0; i < vector1.size(); i++) {
dotProduct += vector1.get(i) * vector2.get(i);
norm1 += vector1.get(i) * vector1.get(i);
norm2 += vector2.get(i) * vector2.get(i);
}
return dotProduct / (float) (Math.sqrt(norm1) * Math.sqrt(norm2));
}
}
@@ -0,0 +1,351 @@
package org.example.service;
import io.milvus.client.MilvusServiceClient;
import io.milvus.grpc.MutationResult;
import io.milvus.param.R;
import io.milvus.param.RpcStatus;
import io.milvus.param.collection.LoadCollectionParam;
import io.milvus.param.dml.DeleteParam;
import io.milvus.param.dml.InsertParam;
import lombok.Getter;
import lombok.Setter;
import org.example.constant.MilvusConstants;
import org.example.dto.DocumentChunk;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.stereotype.Service;
import java.io.File;
import java.nio.file.Files;
import java.nio.file.Path;
import java.nio.file.Paths;
import java.time.LocalDateTime;
import java.util.*;
/**
* 向量索引服务
* 负责读取文件、生成向量、存储到 Milvus
*/
@Service
public class VectorIndexService {
private static final Logger logger = LoggerFactory.getLogger(VectorIndexService.class);
@Autowired
private MilvusServiceClient milvusClient;
@Autowired
private VectorEmbeddingService embeddingService;
@Autowired
private DocumentChunkService chunkService;
@Value("${file.upload.path}")
private String uploadPath;
/**
* 索引指定目录下的所有文件
*
* @param directoryPath 目录路径(可选,默认使用配置的上传目录)
* @return 索引结果 这里可以优化:定时重建目录下所有文件的索引
*/
public IndexingResult indexDirectory(String directoryPath) {
IndexingResult result = new IndexingResult();
result.setStartTime(LocalDateTime.now());
try {
// 使用指定目录或默认上传目录
String targetPath = (directoryPath != null && !directoryPath.trim().isEmpty())
? directoryPath : uploadPath;
Path dirPath = Paths.get(targetPath).normalize();
File directory = dirPath.toFile();
if (!directory.exists() || !directory.isDirectory()) {
throw new IllegalArgumentException("目录不存在或不是有效目录: " + targetPath);
}
result.setDirectoryPath(directory.getAbsolutePath());
// 获取所有支持的文件
File[] files = directory.listFiles((dir, name) ->
name.endsWith(".txt") || name.endsWith(".md")
);
if (files == null || files.length == 0) {
logger.warn("目录中没有找到支持的文件: {}", targetPath);
result.setTotalFiles(0);
result.setSuccess(true);
result.setEndTime(LocalDateTime.now());
return result;
}
result.setTotalFiles(files.length);
logger.info("开始索引目录: {}, 找到 {} 个文件", targetPath, files.length);
// 遍历并索引每个文件
for (File file : files) {
try {
indexSingleFile(file.getAbsolutePath());
result.incrementSuccessCount();
logger.info("✓ 文件索引成功: {}", file.getName());
} catch (Exception e) {
result.incrementFailCount();
result.addFailedFile(file.getAbsolutePath(), e.getMessage());
logger.error("✗ 文件索引失败: {}", file.getName(), e);
}
}
result.setSuccess(result.getFailCount() == 0);
result.setEndTime(LocalDateTime.now());
logger.info("目录索引完成: 总数={}, 成功={}, 失败={}",
result.getTotalFiles(), result.getSuccessCount(), result.getFailCount());
return result;
} catch (Exception e) {
logger.error("索引目录失败", e);
result.setSuccess(false);
result.setErrorMessage(e.getMessage());
result.setEndTime(LocalDateTime.now());
return result;
}
}
/**
* 索引单个文件
*
* @param filePath 文件路径
* @throws Exception 索引失败时抛出异常
*/
public void indexSingleFile(String filePath) throws Exception {
Path path = Paths.get(filePath).normalize();
File file = path.toFile();
if (!file.exists() || !file.isFile()) {
throw new IllegalArgumentException("文件不存在: " + filePath);
}
logger.info("开始索引文件: {}", path);
// 1. 读取文件内容
String content = Files.readString(path);
logger.info("读取文件: {}, 内容长度: {} 字符", path, content.length());
// 2. 删除该文件的旧数据(如果存在)
deleteExistingData(path.toString());
// 3. 文档分片
List<DocumentChunk> chunks = chunkService.chunkDocument(content, path.toString());
logger.info("文档分片完成: {} -> {} 个分片", filePath, chunks.size());
// 4. 为每个分片生成向量并插入 Milvus
for (int i = 0; i < chunks.size(); i++) {
DocumentChunk chunk = chunks.get(i);
try {
// 生成向量
List<Float> vector = embeddingService.generateEmbedding(chunk.getContent());
// 构建元数据(包含文件信息)
Map<String, Object> metadata = buildMetadata(path.toString(), chunk, chunks.size());
// 插入到 Milvus
insertToMilvus(chunk.getContent(), vector, metadata, chunk.getChunkIndex());
logger.info("✓ 分片 {}/{} 索引成功", i + 1, chunks.size());
} catch (Exception e) {
logger.error("✗ 分片 {}/{} 索引失败", i + 1, chunks.size(), e);
throw new RuntimeException("分片索引失败: " + e.getMessage(), e);
}
}
logger.info("文件索引完成: {}, 共 {} 个分片", filePath, chunks.size());
}
/**
* 删除文件的旧数据(根据 metadata._source)
*/
private void deleteExistingData(String filePath) {
try {
// 使用统一的路径分隔符(正斜杠)用于Milvus存储,避免表达式解析错误
// 将系统路径转换为统一格式
Path path = Paths.get(filePath).normalize();
String normalizedPath = path.toString().replace(File.separator, "/");
// 构建删除表达式:metadata["_source"] == "xxx"
String expr = String.format("metadata[\"_source\"] == \"%s\"", normalizedPath);
logger.info("准备删除旧数据,路径: {}, 表达式: {}", normalizedPath, expr);
// 确保 collection 已加载(删除操作需要集合已加载)
R<RpcStatus> loadResponse = milvusClient.loadCollection(
LoadCollectionParam.newBuilder()
.withCollectionName(MilvusConstants.MILVUS_COLLECTION_NAME)
.build()
);
// 状态码 65535 表示集合已经加载,这不是错误
if (loadResponse.getStatus() != 0 && loadResponse.getStatus() != 65535) {
logger.warn("加载 collection 失败: {}", loadResponse.getMessage());
return;
}
DeleteParam deleteParam = DeleteParam.newBuilder()
.withCollectionName(MilvusConstants.MILVUS_COLLECTION_NAME)
.withExpr(expr)
.build();
R<MutationResult> response = milvusClient.delete(deleteParam);
if (response.getStatus() != 0) {
logger.warn("删除旧数据时出现警告: {}", response.getMessage());
} else {
long deletedCount = response.getData().getDeleteCnt();
logger.info("✓ 已删除文件的旧数据: {}, 删除记录数: {}", normalizedPath, deletedCount);
}
} catch (Exception e) {
logger.warn("删除旧数据失败(可能是首次索引): {}", e.getMessage());
}
}
/**
* 构建元数据(包含文件信息)
*/
private Map<String, Object> buildMetadata(String filePath, DocumentChunk chunk, int totalChunks) {
Map<String, Object> metadata = new HashMap<>();
// 标准化路径:使用统一的路径分隔符(正斜杠)用于存储,确保跨平台一致性
Path path = Paths.get(filePath).normalize();
String normalizedPath = path.toString().replace(File.separator, "/");
// 文件信息
Path fileName = path.getFileName();
String fileNameStr = fileName != null ? fileName.toString() : "";
String extension = "";
int dotIndex = fileNameStr.lastIndexOf('.');
if (dotIndex > 0) {
extension = fileNameStr.substring(dotIndex);
}
metadata.put("_source", normalizedPath);
metadata.put("_extension", extension);
metadata.put("_file_name", fileNameStr);
// 分片信息
metadata.put("chunkIndex", chunk.getChunkIndex());
metadata.put("totalChunks", totalChunks);
// 标题信息
if (chunk.getTitle() != null && !chunk.getTitle().isEmpty()) {
metadata.put("title", chunk.getTitle());
}
return metadata;
}
/**
* 插入向量到 Milvus
*/
private void insertToMilvus(String content, List<Float> vector,
Map<String, Object> metadata, int chunkIndex) throws Exception {
try {
// 确保 collection 已加载
R<RpcStatus> loadResponse = milvusClient.loadCollection(
LoadCollectionParam.newBuilder()
.withCollectionName(MilvusConstants.MILVUS_COLLECTION_NAME)
.build()
);
if (loadResponse.getStatus() != 0 && loadResponse.getStatus() != 65535) {
throw new RuntimeException("加载 collection 失败: " + loadResponse.getMessage());
}
// 生成唯一 ID(使用 _source + 分片索引)
String source = (String) metadata.get("_source");
String id = UUID.nameUUIDFromBytes((source + "_" + chunkIndex).getBytes()).toString();
// 构建字段数据
List<InsertParam.Field> fields = new ArrayList<>();
// ID 字段
fields.add(new InsertParam.Field("id", Collections.singletonList(id)));
// content 字段
fields.add(new InsertParam.Field("content", Collections.singletonList(content)));
// vector 字段
fields.add(new InsertParam.Field("vector", Collections.singletonList(vector)));
// metadata 字段(JSON 对象)
com.google.gson.Gson gson = new com.google.gson.Gson();
com.google.gson.JsonObject metadataJson = gson.toJsonTree(metadata).getAsJsonObject();
fields.add(new InsertParam.Field("metadata", Collections.singletonList(metadataJson)));
// 构建插入参数
InsertParam insertParam = InsertParam.newBuilder()
.withCollectionName(MilvusConstants.MILVUS_COLLECTION_NAME)
.withFields(fields)
.build();
// 执行插入
R<MutationResult> insertResponse = milvusClient.insert(insertParam);
if (insertResponse.getStatus() != 0) {
throw new RuntimeException("插入向量失败: " + insertResponse.getMessage());
}
logger.debug("向量插入成功: id={}, source={}, chunk={}", id, source, chunkIndex);
} catch (Exception e) {
logger.error("插入向量到 Milvus 失败", e);
throw e;
}
}
/**
* 索引结果类
*/
@Getter
public static class IndexingResult {
@Setter
private boolean success;
@Setter
private String directoryPath;
@Setter
private int totalFiles;
private int successCount;
private int failCount;
@Setter
private LocalDateTime startTime;
@Setter
private LocalDateTime endTime;
@Setter
private String errorMessage;
private Map<String, String> failedFiles = new HashMap<>();
public void incrementSuccessCount() {
this.successCount++;
}
public void incrementFailCount() {
this.failCount++;
}
public long getDurationMs() {
if (startTime != null && endTime != null) {
return java.time.Duration.between(startTime, endTime).toMillis();
}
return 0;
}
public void addFailedFile(String filePath, String error) {
this.failedFiles.put(filePath, error);
}
}
}
@@ -0,0 +1,108 @@
package org.example.service;
import io.milvus.client.MilvusServiceClient;
import io.milvus.grpc.SearchResults;
import io.milvus.param.R;
import io.milvus.param.dml.SearchParam;
import io.milvus.response.SearchResultsWrapper;
import lombok.Getter;
import lombok.Setter;
import org.example.constant.MilvusConstants;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.stereotype.Service;
import java.util.ArrayList;
import java.util.Collections;
import java.util.List;
/**
* 向量搜索服务
* 负责从 Milvus 中搜索相似向量
*/
@Service
public class VectorSearchService {
private static final Logger logger = LoggerFactory.getLogger(VectorSearchService.class);
@Autowired
private MilvusServiceClient milvusClient;
@Autowired
private VectorEmbeddingService embeddingService;
/**
* 搜索相似文档
*
* @param query 查询文本
* @param topK 返回最相似的K个结果
* @return 搜索结果列表
*/
public List<SearchResult> searchSimilarDocuments(String query, int topK) {
try {
logger.info("开始搜索相似文档, 查询: {}, topK: {}", query, topK);
// 1. 将查询文本向量化
List<Float> queryVector = embeddingService.generateQueryVector(query);
logger.debug("查询向量生成成功, 维度: {}", queryVector.size());
// 2. 构建搜索参数
SearchParam searchParam = SearchParam.newBuilder()
.withCollectionName(MilvusConstants.MILVUS_COLLECTION_NAME)
.withVectorFieldName("vector")
.withVectors(Collections.singletonList(queryVector))
.withTopK(topK)
.withMetricType(io.milvus.param.MetricType.L2)
.withOutFields(List.of("id", "content", "metadata"))
.withParams("{\"nprobe\":10}")
.build();
// 3. 执行搜索
R<SearchResults> searchResponse = milvusClient.search(searchParam);
if (searchResponse.getStatus() != 0) {
throw new RuntimeException("向量搜索失败: " + searchResponse.getMessage());
}
// 4. 解析搜索结果
SearchResultsWrapper wrapper = new SearchResultsWrapper(searchResponse.getData().getResults());
List<SearchResult> results = new ArrayList<>();
for (int i = 0; i < wrapper.getRowRecords(0).size(); i++) {
SearchResult result = new SearchResult();
result.setId((String) wrapper.getIDScore(0).get(i).get("id"));
result.setContent((String) wrapper.getFieldData("content", 0).get(i));
result.setScore(wrapper.getIDScore(0).get(i).getScore());
// 解析 metadata
Object metadataObj = wrapper.getFieldData("metadata", 0).get(i);
if (metadataObj != null) {
result.setMetadata(metadataObj.toString());
}
results.add(result);
}
logger.info("搜索完成, 找到 {} 个相似文档", results.size());
return results;
} catch (Exception e) {
logger.error("搜索相似文档失败", e);
throw new RuntimeException("搜索失败: " + e.getMessage(), e);
}
}
/**
* 搜索结果类
*/
@Setter
@Getter
public static class SearchResult {
private String id;
private String content;
private float score;
private String metadata;
}
}
@@ -0,0 +1,69 @@
package org.example.tool;
import io.milvus.client.MilvusServiceClient;
import io.milvus.param.ConnectParam;
import io.milvus.param.R;
import io.milvus.param.RpcStatus;
import io.milvus.param.collection.DropCollectionParam;
import io.milvus.param.collection.HasCollectionParam;
/**
* 删除 Milvus Collection 的工具类
* 用于重建 Collection 时清理旧数据
*/
public class DropCollection {
public static void main(String[] args) {
MilvusServiceClient client = null;
try {
// 连接到 Milvus
System.out.println("正在连接到 Milvus localhost:19530...");
client = new MilvusServiceClient(
ConnectParam.newBuilder()
.withHost("localhost")
.withPort(19530)
.build()
);
System.out.println("✓ 连接成功");
String collectionName = "biz";
// 检查 Collection 是否存在
R<Boolean> hasResponse = client.hasCollection(
HasCollectionParam.newBuilder()
.withCollectionName(collectionName)
.build()
);
if (hasResponse.getData()) {
System.out.println("发现 Collection: " + collectionName);
System.out.println("正在删除...");
// 删除 Collection
R<RpcStatus> dropResponse = client.dropCollection(
DropCollectionParam.newBuilder()
.withCollectionName(collectionName)
.build()
);
if (dropResponse.getStatus() == 0) {
System.out.println("✓ Collection 已成功删除");
System.out.println("\n请重启 Spring Boot 应用,它会自动创建新的 FloatVector Collection");
} else {
System.err.println("✗ 删除失败: " + dropResponse.getMessage());
}
} else {
System.out.println("Collection '" + collectionName + "' 不存在");
}
} catch (Exception e) {
System.err.println("错误: " + e.getMessage());
e.printStackTrace();
} finally {
if (client != null) {
client.close();
}
}
}
}