Compare commits
9 Commits
93296eb810
...
develop
| Author | SHA1 | Date | |
|---|---|---|---|
| 130423edb4 | |||
| 7900552ede | |||
| 368b90b211 | |||
| 1ea7b7527f | |||
| 68fd303656 | |||
| c1fe64cefa | |||
| 1b36067e6c | |||
| 876517f821 | |||
| 2d50f7de15 |
@@ -43,6 +43,11 @@ public class StoreOptions extends Metadata {
|
||||
public void setEmbeddingOptions(EmbeddingOptions embeddingOptions) {
|
||||
throw new IllegalStateException("Can not set embeddingOptions to the default instance.");
|
||||
}
|
||||
|
||||
@Override
|
||||
public void setTimeoutMillis(Long timeoutMillis) {
|
||||
throw new IllegalStateException("Can not set timeoutMillis to the default instance.");
|
||||
}
|
||||
};
|
||||
|
||||
/**
|
||||
@@ -65,6 +70,11 @@ public class StoreOptions extends Metadata {
|
||||
*/
|
||||
private EmbeddingOptions embeddingOptions = EmbeddingOptions.DEFAULT;
|
||||
|
||||
/**
|
||||
* Optional upper bound for one store operation.
|
||||
*/
|
||||
private Long timeoutMillis;
|
||||
|
||||
|
||||
public String getCollectionName() {
|
||||
return collectionName;
|
||||
@@ -111,6 +121,17 @@ public class StoreOptions extends Metadata {
|
||||
this.embeddingOptions = embeddingOptions;
|
||||
}
|
||||
|
||||
public Long getTimeoutMillis() {
|
||||
return timeoutMillis;
|
||||
}
|
||||
|
||||
public void setTimeoutMillis(Long timeoutMillis) {
|
||||
if (timeoutMillis != null && timeoutMillis <= 0L) {
|
||||
throw new IllegalArgumentException("timeoutMillis must be greater than zero");
|
||||
}
|
||||
this.timeoutMillis = timeoutMillis;
|
||||
}
|
||||
|
||||
|
||||
public static StoreOptions ofCollectionName(String collectionName) {
|
||||
StoreOptions storeOptions = new StoreOptions();
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
/*
|
||||
* Copyright (c) 2023-2026, Easy-Agents (fuhai999@gmail.com).
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
*/
|
||||
package com.easyagents.core.store;
|
||||
|
||||
/**
|
||||
* Indicates that a store operation exhausted its caller-provided time budget.
|
||||
*/
|
||||
public class StoreTimeoutException extends RuntimeException {
|
||||
|
||||
public StoreTimeoutException(String message) {
|
||||
super(message);
|
||||
}
|
||||
|
||||
public StoreTimeoutException(String message, Throwable cause) {
|
||||
super(message, cause);
|
||||
}
|
||||
}
|
||||
@@ -2,6 +2,7 @@ package com.easyagents.document.core.async;
|
||||
|
||||
import com.easyagents.core.util.StringUtil;
|
||||
import com.easyagents.document.core.exception.DocumentParseException;
|
||||
import com.easyagents.document.core.exception.DocumentAsyncTaskNotFoundException;
|
||||
import com.easyagents.document.core.entity.ParseResponse;
|
||||
import com.easyagents.document.core.entity.ParseTaskInfo;
|
||||
import com.easyagents.document.core.entity.ParseTaskStatus;
|
||||
@@ -135,7 +136,7 @@ public class DocumentAsyncTaskManager {
|
||||
}
|
||||
DocumentAsyncTaskRecord record = repository.find(taskId);
|
||||
if (record == null) {
|
||||
throw new DocumentParseException("Document async task not found: " + taskId);
|
||||
throw new DocumentAsyncTaskNotFoundException(taskId);
|
||||
}
|
||||
return record;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
package com.easyagents.document.core.exception;
|
||||
|
||||
/**
|
||||
* 进程内异步文档任务不存在。
|
||||
*
|
||||
* <p>本地 Office 解析任务允许使用内存仓库;进程重启后,调用方可以
|
||||
* 通过该异常识别执行实例已经丢失,并从持久化业务任务重新提交。</p>
|
||||
*
|
||||
* @author Codex
|
||||
* @since 2026-09-02
|
||||
*/
|
||||
public class DocumentAsyncTaskNotFoundException extends DocumentParseException {
|
||||
|
||||
private final String taskId;
|
||||
|
||||
public DocumentAsyncTaskNotFoundException(String taskId) {
|
||||
super("Document async task not found: " + taskId);
|
||||
this.taskId = taskId;
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取已丢失的任务 ID。
|
||||
*
|
||||
* @return 任务 ID
|
||||
*/
|
||||
public String getTaskId() {
|
||||
return taskId;
|
||||
}
|
||||
}
|
||||
@@ -4,6 +4,7 @@ import com.easyagents.document.core.entity.ParseResponse;
|
||||
import com.easyagents.document.core.entity.ParseResult;
|
||||
import com.easyagents.document.core.entity.ParseTaskInfo;
|
||||
import com.easyagents.document.core.entity.ParseTaskStatus;
|
||||
import com.easyagents.document.core.exception.DocumentAsyncTaskNotFoundException;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
|
||||
@@ -18,6 +19,21 @@ import java.util.concurrent.Executor;
|
||||
*/
|
||||
public class DocumentAsyncTaskManagerTest {
|
||||
|
||||
@Test
|
||||
public void shouldExposeMissingInMemoryTaskAsRecoverableSignal() {
|
||||
DocumentAsyncTaskManager manager = new DocumentAsyncTaskManager(
|
||||
new InMemoryDocumentAsyncTaskRepository(),
|
||||
Runnable::run
|
||||
);
|
||||
|
||||
try {
|
||||
manager.queryTaskInfo("lost-task");
|
||||
Assert.fail("expected DocumentAsyncTaskNotFoundException");
|
||||
} catch (DocumentAsyncTaskNotFoundException error) {
|
||||
Assert.assertEquals("lost-task", error.getTaskId());
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldTrackTaskLifecycleAndResult() {
|
||||
Executor directExecutor = new Executor() {
|
||||
|
||||
@@ -14,6 +14,7 @@ import com.easyagents.document.core.entity.ParseRequest;
|
||||
import com.easyagents.document.core.entity.ParseResponse;
|
||||
import com.easyagents.document.core.entity.ParseResult;
|
||||
import com.easyagents.document.core.entity.XlsxParseRequest;
|
||||
import com.easyagents.document.core.exception.DocumentParseException;
|
||||
import com.easyagents.document.core.support.AbstractAsyncDocumentParseService;
|
||||
import com.easyagents.document.xlsx.XlsxDocumentProvider;
|
||||
import com.easyagents.document.xlsx.model.XlsxCellArtifact;
|
||||
@@ -52,6 +53,10 @@ import java.util.concurrent.Executors;
|
||||
public class MineruXlsxDocumentParseService extends AbstractAsyncDocumentParseService<XlsxParseRequest> implements XlsxDocumentProvider {
|
||||
|
||||
public static final String PROVIDER_NAME = "mineru";
|
||||
private static final byte[] OLE2_SIGNATURE = new byte[] {
|
||||
(byte) 0xd0, (byte) 0xcf, 0x11, (byte) 0xe0,
|
||||
(byte) 0xa1, (byte) 0xb1, 0x1a, (byte) 0xe1
|
||||
};
|
||||
|
||||
private final MineruProperties properties;
|
||||
private final MineruClient client;
|
||||
@@ -145,6 +150,9 @@ public class MineruXlsxDocumentParseService extends AbstractAsyncDocumentParseSe
|
||||
|
||||
@Override
|
||||
protected ParseResponse doParse(XlsxParseRequest request, DocumentAsyncTaskUpdater updater) {
|
||||
for (ParseFile file : request.getFiles()) {
|
||||
validateXlsxContent(file);
|
||||
}
|
||||
ParseResponse response = new ParseResponse();
|
||||
List<ParseResult> results = new ArrayList<ParseResult>();
|
||||
String backend = null;
|
||||
@@ -211,6 +219,50 @@ public class MineruXlsxDocumentParseService extends AbstractAsyncDocumentParseSe
|
||||
return aggregate;
|
||||
}
|
||||
|
||||
/**
|
||||
* 在 POI 打开工作簿前校验 XLSX 容器签名,避免向用户暴露底层格式异常。
|
||||
*
|
||||
* @param file 待解析文件
|
||||
*/
|
||||
private void validateXlsxContent(ParseFile file) {
|
||||
byte[] content = file == null ? null : file.getContent();
|
||||
if (hasZipSignature(content)) {
|
||||
return;
|
||||
}
|
||||
String fileName = file == null || !StringUtil.hasText(file.getFileName())
|
||||
? "当前文件"
|
||||
: "文件“" + file.getFileName() + "”";
|
||||
String reason = startsWith(content, OLE2_SIGNATURE)
|
||||
? "可能是旧版 XLS 或已加密文件"
|
||||
: "文件内容与 .xlsx 扩展名不一致或文件已损坏";
|
||||
throw new DocumentParseException(
|
||||
fileName + "不是标准 XLSX," + reason
|
||||
+ "。请解除保护后用 Excel/WPS 另存为 XLSX(修改文件后缀无效)"
|
||||
);
|
||||
}
|
||||
|
||||
private boolean hasZipSignature(byte[] content) {
|
||||
return content != null
|
||||
&& content.length >= 4
|
||||
&& content[0] == 'P'
|
||||
&& content[1] == 'K'
|
||||
&& ((content[2] == 3 && content[3] == 4)
|
||||
|| (content[2] == 5 && content[3] == 6)
|
||||
|| (content[2] == 7 && content[3] == 8));
|
||||
}
|
||||
|
||||
private boolean startsWith(byte[] content, byte[] signature) {
|
||||
if (content == null || content.length < signature.length) {
|
||||
return false;
|
||||
}
|
||||
for (int index = 0; index < signature.length; index++) {
|
||||
if (content[index] != signature[index]) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
private SheetExtraction extractSheet(XSSFSheet sheet,
|
||||
int sheetIndex,
|
||||
DataFormatter formatter,
|
||||
|
||||
@@ -16,6 +16,7 @@ import com.easyagents.document.core.entity.ParseTaskStatus;
|
||||
import com.easyagents.document.core.entity.XlsxParseRequest;
|
||||
import com.easyagents.document.core.exception.DocumentParseException;
|
||||
import com.easyagents.document.xlsx.model.XlsxParseArtifact;
|
||||
import org.apache.poi.hssf.usermodel.HSSFWorkbook;
|
||||
import org.apache.poi.ss.usermodel.ClientAnchor;
|
||||
import org.apache.poi.xssf.usermodel.XSSFDrawing;
|
||||
import org.apache.poi.xssf.usermodel.XSSFSheet;
|
||||
@@ -119,6 +120,31 @@ public class MineruXlsxDocumentParseServiceTest {
|
||||
Assert.assertEquals("image/jpeg", result.getImages().get(0).getMimeType());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRejectLegacyXlsContentWithActionableMessage() throws Exception {
|
||||
RecordingClient client = new RecordingClient(defaultProperties());
|
||||
MineruMapper mapper = new MineruMapper(defaultProperties());
|
||||
MineruXlsxDocumentParseService service = new MineruXlsxDocumentParseService(
|
||||
defaultProperties(),
|
||||
client,
|
||||
mapper,
|
||||
new DocumentAsyncTaskManager(new InMemoryDocumentAsyncTaskRepository(), directExecutor())
|
||||
);
|
||||
XlsxParseRequest request = new XlsxParseRequest();
|
||||
request.addFile(ParseFile.of("legacy.xlsx", buildLegacyWorkbookBytes()));
|
||||
|
||||
DocumentParseException error = Assert.assertThrows(
|
||||
DocumentParseException.class,
|
||||
() -> service.parse(request)
|
||||
);
|
||||
|
||||
Assert.assertEquals(
|
||||
"文件“legacy.xlsx”不是标准 XLSX,可能是旧版 XLS 或已加密文件。"
|
||||
+ "请解除保护后用 Excel/WPS 另存为 XLSX(修改文件后缀无效)",
|
||||
error.getMessage()
|
||||
);
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldAppendImageReferenceForImageOnlySheet() throws Exception {
|
||||
RecordingClient client = new RecordingClient(defaultProperties());
|
||||
@@ -255,6 +281,15 @@ public class MineruXlsxDocumentParseServiceTest {
|
||||
return writeWorkbook(workbook);
|
||||
}
|
||||
|
||||
private byte[] buildLegacyWorkbookBytes() throws Exception {
|
||||
try (HSSFWorkbook workbook = new HSSFWorkbook();
|
||||
ByteArrayOutputStream outputStream = new ByteArrayOutputStream()) {
|
||||
workbook.createSheet("Sheet1").createRow(0).createCell(0).setCellValue("旧版表格");
|
||||
workbook.write(outputStream);
|
||||
return outputStream.toByteArray();
|
||||
}
|
||||
}
|
||||
|
||||
private void addPicture(XSSFWorkbook workbook,
|
||||
XSSFSheet sheet,
|
||||
int rowIndex,
|
||||
|
||||
@@ -0,0 +1,129 @@
|
||||
package com.easyagents.federation.sql.adapter.jdbc;
|
||||
|
||||
import com.easyagents.federation.sql.api.FederationSqlEngine;
|
||||
import com.easyagents.federation.sql.api.FederationSqlEngines;
|
||||
import com.easyagents.federation.sql.api.SqlQueryCommand;
|
||||
import com.easyagents.federation.sql.execute.FederationResultCursor;
|
||||
import com.easyagents.federation.sql.federation.FederationExecutionPolicy;
|
||||
import com.easyagents.federation.sql.federation.FederationQueryScopeDefinition;
|
||||
import com.easyagents.federation.sql.federation.FederationSourceBindingDefinition;
|
||||
import com.easyagents.federation.sql.source.FederationDataSourceHandles;
|
||||
import com.easyagents.federation.sql.source.FederationSourceDefinition;
|
||||
import com.easyagents.federation.sql.source.RuntimeFingerprint;
|
||||
import com.easyagents.federation.sql.source.SourceApplyOptions;
|
||||
import com.easyagents.federation.sql.source.SourceId;
|
||||
import com.mysql.cj.jdbc.MysqlDataSource;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Assume;
|
||||
import org.junit.Test;
|
||||
|
||||
import java.math.BigDecimal;
|
||||
import java.sql.Connection;
|
||||
import java.sql.PreparedStatement;
|
||||
import java.sql.ResultSet;
|
||||
import java.sql.SQLException;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
/**
|
||||
* 亿级数据集 100,000 行 Cursor 手工基准;默认跳过,必须显式启用。
|
||||
*/
|
||||
public class Jdbc100mCursorManualBenchmarkTest {
|
||||
|
||||
/**
|
||||
* 从支付库顺序消费恰好 100,000 行,验证 Cursor 流式消费和资源上限。
|
||||
*/
|
||||
@Test
|
||||
public void shouldConsumeOneHundredThousandRows() {
|
||||
Assume.assumeTrue(Boolean.getBoolean("federation.100m.enabled"));
|
||||
String host = System.getProperty("federation.100m.mysql.host", "127.0.0.1");
|
||||
int port = Integer.getInteger("federation.100m.mysql.port", 33307);
|
||||
String database = System.getProperty(
|
||||
"federation.100m.mysql.database", "ef_bank_payment_perf");
|
||||
String username = System.getProperty("federation.100m.mysql.username", "ef_bench_ro");
|
||||
String password = System.getProperty("federation.100m.mysql.password");
|
||||
Assert.assertNotNull("必须通过系统属性提供只读密码", password);
|
||||
|
||||
MysqlDataSource dataSource = new MysqlDataSource();
|
||||
dataSource.setUrl("jdbc:mysql://" + host + ':' + port + '/' + database
|
||||
+ "?useSSL=false&allowPublicKeyRetrieval=true&serverTimezone=Asia/Shanghai"
|
||||
+ "&useCursorFetch=true&defaultFetchSize=1000");
|
||||
dataSource.setUser(username);
|
||||
dataSource.setPassword(password);
|
||||
|
||||
SourceId sourceId = new SourceId("payment-100m");
|
||||
FederationSourceDefinition definition = new FederationSourceDefinition(
|
||||
sourceId, 1, JdbcFederationSqlAdapterProvider.ADAPTER_ID,
|
||||
List.of(new JdbcSchemaDefinition("MAIN", database, null)), Map.of());
|
||||
FederationExecutionPolicy policy = new FederationExecutionPolicy(
|
||||
6, 12, 2, 100_000, 64L * 1024L * 1024L, 60_000);
|
||||
FederationQueryScopeDefinition scope = FederationQueryScopeDefinition.virtual(
|
||||
"payment-cursor-100m", 1,
|
||||
Map.of("PAY", FederationSourceBindingDefinition.of(sourceId, 1)),
|
||||
"PAY", policy);
|
||||
|
||||
CursorResult nativeResult = consumeNative(dataSource);
|
||||
CursorResult federationResult;
|
||||
long federationStarted = System.nanoTime();
|
||||
try (FederationSqlEngine engine = FederationSqlEngines.builder()
|
||||
.dataSourceResolver(ignored -> FederationDataSourceHandles.shared(
|
||||
dataSource, new RuntimeFingerprint("MySQL", "8.0", "JDBC", "8.4", "1")))
|
||||
.adapter(new JdbcFederationSqlAdapterProvider())
|
||||
.federationExecutionPolicy(policy)
|
||||
.build()) {
|
||||
engine.sources().apply(definition, SourceApplyOptions.prewarmNow());
|
||||
String sql = "SELECT payment_id, amount FROM PAY.MAIN.payment_order "
|
||||
+ "WHERE payment_id BETWEEN 1 AND 100000 ORDER BY payment_id";
|
||||
int rows = 0;
|
||||
long amountCents = 0;
|
||||
try (FederationResultCursor cursor = engine.query(
|
||||
SqlQueryCommand.of(sql, scope, List.of()))) {
|
||||
while (cursor.next()) {
|
||||
rows++;
|
||||
amountCents += ((BigDecimal) cursor.getObject(2))
|
||||
.movePointRight(2).longValue();
|
||||
}
|
||||
}
|
||||
federationResult = new CursorResult(rows, amountCents,
|
||||
elapsedMillis(federationStarted));
|
||||
}
|
||||
Assert.assertEquals(100_000, nativeResult.rows());
|
||||
Assert.assertEquals(100_000, federationResult.rows());
|
||||
Assert.assertEquals(nativeResult.amountCents(), federationResult.amountCents());
|
||||
Assert.assertTrue(federationResult.amountCents() > 0L);
|
||||
System.out.println("BENCH_CURSOR_100K rows=" + federationResult.rows()
|
||||
+ " nativeMillis=" + nativeResult.elapsedMillis()
|
||||
+ " federationMillis=" + federationResult.elapsedMillis()
|
||||
+ " amountCents=" + federationResult.amountCents());
|
||||
}
|
||||
|
||||
private CursorResult consumeNative(MysqlDataSource dataSource) {
|
||||
String sql = "SELECT payment_id, amount FROM payment_order "
|
||||
+ "WHERE payment_id BETWEEN 1 AND 100000 ORDER BY payment_id";
|
||||
long started = System.nanoTime();
|
||||
int rows = 0;
|
||||
long amountCents = 0;
|
||||
try (Connection connection = dataSource.getConnection();
|
||||
PreparedStatement statement = connection.prepareStatement(sql,
|
||||
ResultSet.TYPE_FORWARD_ONLY, ResultSet.CONCUR_READ_ONLY)) {
|
||||
statement.setFetchSize(1000);
|
||||
try (ResultSet resultSet = statement.executeQuery()) {
|
||||
while (resultSet.next()) {
|
||||
rows++;
|
||||
amountCents += resultSet.getBigDecimal(2)
|
||||
.movePointRight(2).longValue();
|
||||
}
|
||||
}
|
||||
} catch (SQLException exception) {
|
||||
throw new IllegalStateException("原生 JDBC Cursor 基准执行失败", exception);
|
||||
}
|
||||
return new CursorResult(rows, amountCents, elapsedMillis(started));
|
||||
}
|
||||
|
||||
private long elapsedMillis(long started) {
|
||||
return (System.nanoTime() - started) / 1_000_000L;
|
||||
}
|
||||
|
||||
private record CursorResult(int rows, long amountCents, long elapsedMillis) {
|
||||
}
|
||||
}
|
||||
@@ -27,6 +27,7 @@ import com.easyagents.federation.sql.source.RuntimeFingerprint;
|
||||
import com.easyagents.federation.sql.source.SourceApplyOptions;
|
||||
import com.easyagents.federation.sql.source.SourceId;
|
||||
import java.math.BigDecimal;
|
||||
import java.sql.SQLException;
|
||||
import java.time.Instant;
|
||||
import java.sql.Connection;
|
||||
import java.sql.Statement;
|
||||
@@ -34,6 +35,9 @@ import java.sql.Types;
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.concurrent.CountDownLatch;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
import java.util.concurrent.atomic.AtomicLong;
|
||||
import java.util.concurrent.atomic.AtomicReference;
|
||||
import org.h2.jdbcx.JdbcDataSource;
|
||||
@@ -50,6 +54,13 @@ public class JdbcFederatedQueryEngineTest {
|
||||
private static final SourceId SALES_SOURCE = new SourceId("sales-source");
|
||||
private static final SourceId BILLING_SOURCE = new SourceId("billing-source");
|
||||
private static final SourceId REGION_SOURCE = new SourceId("region-source");
|
||||
private static final AtomicInteger ACTIVE_FRAGMENTS = new AtomicInteger();
|
||||
private static final AtomicInteger MAXIMUM_ACTIVE_FRAGMENTS = new AtomicInteger();
|
||||
private static final AtomicInteger FAILING_TRACKED_VALUE = new AtomicInteger(
|
||||
Integer.MIN_VALUE
|
||||
);
|
||||
private static final AtomicReference<CountDownLatch> FRAGMENT_START_BARRIER =
|
||||
new AtomicReference<>(new CountDownLatch(0));
|
||||
|
||||
private JdbcDataSource sales;
|
||||
private JdbcDataSource billing;
|
||||
@@ -80,7 +91,10 @@ public class JdbcFederatedQueryEngineTest {
|
||||
+ "TIMESTAMP WITH TIME ZONE '2026-08-21 12:00:00.123456+08:00')",
|
||||
"CREATE TABLE PRECISE_EVENT (ID INT PRIMARY KEY, EVENT_AT TIMESTAMP(6))",
|
||||
"INSERT INTO PRECISE_EVENT VALUES "
|
||||
+ "(1, TIMESTAMP '2026-08-21 12:00:00.123456')");
|
||||
+ "(1, TIMESTAMP '2026-08-21 12:00:00.123456')",
|
||||
fragmentTrackingAlias(),
|
||||
"CREATE VIEW TRACKED_VALUE AS SELECT TRACK_FRAGMENT(ID) AS TRACKED_KEY "
|
||||
+ "FROM CUSTOMER WHERE ID <= 2");
|
||||
execute(billing,
|
||||
"CREATE TABLE ORDER_ITEM (ID INT PRIMARY KEY, CUSTOMER_ID INT NOT NULL, AMOUNT DECIMAL(12,2), CODE VARCHAR(64))",
|
||||
"INSERT INTO ORDER_ITEM VALUES (10, 1, 30.00, 'Alice'), (11, 1, 20.00, 'Alice'), (12, 2, 80.00, 'Bob')",
|
||||
@@ -90,10 +104,16 @@ public class JdbcFederatedQueryEngineTest {
|
||||
+ "TIMESTAMP WITH TIME ZONE '2026-08-21 06:00:00.123456+02:00')",
|
||||
"CREATE TABLE PRECISE_EVENT (ID INT PRIMARY KEY, EVENT_AT TIMESTAMP(6))",
|
||||
"INSERT INTO PRECISE_EVENT VALUES "
|
||||
+ "(2, TIMESTAMP '2026-08-21 08:30:00.654321')");
|
||||
+ "(2, TIMESTAMP '2026-08-21 08:30:00.654321')",
|
||||
fragmentTrackingAlias(),
|
||||
"CREATE VIEW TRACKED_VALUE AS SELECT TRACK_FRAGMENT(ID) AS TRACKED_KEY "
|
||||
+ "FROM ORDER_ITEM WHERE ID <= 11");
|
||||
execute(region,
|
||||
"CREATE TABLE CUSTOMER_REGION (CUSTOMER_ID INT PRIMARY KEY, REGION VARCHAR(64))",
|
||||
"INSERT INTO CUSTOMER_REGION VALUES (1, 'North'), (2, 'South'), (3, 'West')");
|
||||
"INSERT INTO CUSTOMER_REGION VALUES (1, 'North'), (2, 'South'), (3, 'West')",
|
||||
fragmentTrackingAlias(),
|
||||
"CREATE VIEW TRACKED_VALUE AS SELECT TRACK_FRAGMENT(CUSTOMER_ID) AS TRACKED_KEY "
|
||||
+ "FROM CUSTOMER_REGION WHERE CUSTOMER_ID <= 2");
|
||||
|
||||
RuntimeFingerprint fingerprint = new RuntimeFingerprint(
|
||||
"H2", "2", "H2 JDBC Driver", "2", "1"
|
||||
@@ -356,6 +376,91 @@ public class JdbcFederatedQueryEngineTest {
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证真实 JDBC Fragment 执行会重叠,且始终受查询级并发上限约束。
|
||||
*/
|
||||
@Test
|
||||
public void shouldBoundConcurrentJdbcFragments() {
|
||||
ACTIVE_FRAGMENTS.set(0);
|
||||
MAXIMUM_ACTIVE_FRAGMENTS.set(0);
|
||||
FRAGMENT_START_BARRIER.set(new CountDownLatch(2));
|
||||
FederationQueryScopeDefinition threeSourceScope =
|
||||
new FederationQueryScopeDefinition(
|
||||
"tracked-fragments",
|
||||
1,
|
||||
Map.of(
|
||||
"SALES", FederationSourceBindingDefinition.of(SALES_SOURCE, 1),
|
||||
"BILLING", FederationSourceBindingDefinition.of(BILLING_SOURCE, 1),
|
||||
"REGION", FederationSourceBindingDefinition.of(REGION_SOURCE, 1)
|
||||
),
|
||||
"SALES",
|
||||
threeSourcePolicy()
|
||||
);
|
||||
|
||||
try (FederationResultCursor cursor = engine.query(SqlQueryCommand.of(
|
||||
"SELECT COUNT(*) FROM SALES.APP.TRACKED_VALUE s "
|
||||
+ "JOIN BILLING.APP.TRACKED_VALUE b ON s.TRACKED_KEY = b.TRACKED_KEY "
|
||||
+ "JOIN REGION.APP.TRACKED_VALUE r ON s.TRACKED_KEY = r.TRACKED_KEY",
|
||||
threeSourceScope,
|
||||
List.of()
|
||||
))) {
|
||||
Assert.assertTrue(cursor.next());
|
||||
Assert.assertEquals(8L, cursor.getObject(1));
|
||||
Assert.assertFalse(cursor.next());
|
||||
}
|
||||
|
||||
Assert.assertEquals(2, MAXIMUM_ACTIVE_FRAGMENTS.get());
|
||||
Assert.assertEquals(0, ACTIVE_FRAGMENTS.get());
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证任一并行 Fragment 失败后立即传播原始执行错误,不等待未启动任务超时。
|
||||
*/
|
||||
@Test
|
||||
public void shouldFailFastWhenConcurrentFragmentFails() {
|
||||
ACTIVE_FRAGMENTS.set(0);
|
||||
MAXIMUM_ACTIVE_FRAGMENTS.set(0);
|
||||
FRAGMENT_START_BARRIER.set(new CountDownLatch(0));
|
||||
FAILING_TRACKED_VALUE.set(10);
|
||||
FederationQueryScopeDefinition threeSourceScope =
|
||||
new FederationQueryScopeDefinition(
|
||||
"tracked-fragment-failure",
|
||||
1,
|
||||
Map.of(
|
||||
"SALES", FederationSourceBindingDefinition.of(SALES_SOURCE, 1),
|
||||
"BILLING", FederationSourceBindingDefinition.of(BILLING_SOURCE, 1),
|
||||
"REGION", FederationSourceBindingDefinition.of(REGION_SOURCE, 1)
|
||||
),
|
||||
"SALES",
|
||||
new FederationExecutionPolicy(
|
||||
3,
|
||||
8,
|
||||
2,
|
||||
100_000,
|
||||
64L * 1024L * 1024L,
|
||||
1_000
|
||||
)
|
||||
);
|
||||
|
||||
try (FederationResultCursor cursor = engine.query(SqlQueryCommand.of(
|
||||
"SELECT COUNT(*) FROM SALES.APP.TRACKED_VALUE s "
|
||||
+ "JOIN BILLING.APP.TRACKED_VALUE b ON s.TRACKED_KEY = b.TRACKED_KEY "
|
||||
+ "JOIN REGION.APP.TRACKED_VALUE r ON s.TRACKED_KEY = r.TRACKED_KEY",
|
||||
threeSourceScope,
|
||||
List.of()
|
||||
))) {
|
||||
cursor.next();
|
||||
Assert.fail("failing fragment should terminate the federated query");
|
||||
} catch (FederationSqlException exception) {
|
||||
Assert.assertEquals(
|
||||
FederationSqlErrorCode.EXECUTION_FAILED,
|
||||
exception.errorCode()
|
||||
);
|
||||
} finally {
|
||||
FAILING_TRACKED_VALUE.set(Integer.MIN_VALUE);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证两个分片分别绑定原始查询中的动态参数。
|
||||
*/
|
||||
@@ -717,6 +822,41 @@ public class JdbcFederatedQueryEngineTest {
|
||||
return dataSource;
|
||||
}
|
||||
|
||||
/**
|
||||
* H2 视图调用的并发跟踪函数。
|
||||
*
|
||||
* @param value 原始列值
|
||||
* @return 固定键值
|
||||
* @throws SQLException 分片未能在时限内重叠启动
|
||||
*/
|
||||
public static int trackFragment(int value) throws SQLException {
|
||||
int active = ACTIVE_FRAGMENTS.incrementAndGet();
|
||||
MAXIMUM_ACTIVE_FRAGMENTS.accumulateAndGet(active, Math::max);
|
||||
CountDownLatch barrier = FRAGMENT_START_BARRIER.get();
|
||||
barrier.countDown();
|
||||
try {
|
||||
if (value == FAILING_TRACKED_VALUE.get()) {
|
||||
throw new SQLException("simulated federation fragment failure");
|
||||
}
|
||||
if (!barrier.await(2, TimeUnit.SECONDS)) {
|
||||
throw new SQLException("federation fragments did not overlap");
|
||||
}
|
||||
Thread.sleep(100L);
|
||||
return 1;
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
throw new SQLException("fragment tracking was interrupted", exception);
|
||||
} finally {
|
||||
ACTIVE_FRAGMENTS.decrementAndGet();
|
||||
}
|
||||
}
|
||||
|
||||
private static String fragmentTrackingAlias() {
|
||||
return "CREATE ALIAS TRACK_FRAGMENT FOR \""
|
||||
+ JdbcFederatedQueryEngineTest.class.getName()
|
||||
+ ".trackFragment\"";
|
||||
}
|
||||
|
||||
/**
|
||||
* 根据物理源选择测试数据库。
|
||||
*
|
||||
|
||||
@@ -0,0 +1,185 @@
|
||||
package com.easyagents.federation.sql.adapter.jdbc;
|
||||
|
||||
import com.easyagents.federation.sql.api.FederationSqlEngine;
|
||||
import com.easyagents.federation.sql.api.FederationSqlEngines;
|
||||
import com.easyagents.federation.sql.api.FederationSqlErrorCode;
|
||||
import com.easyagents.federation.sql.api.FederationSqlException;
|
||||
import com.easyagents.federation.sql.api.SqlQueryCommand;
|
||||
import com.easyagents.federation.sql.compile.FederationSqlPlan;
|
||||
import com.easyagents.federation.sql.compile.SqlCompileRequest;
|
||||
import com.easyagents.federation.sql.execute.FederationResultCursor;
|
||||
import com.easyagents.federation.sql.federation.FederationExecutionPolicy;
|
||||
import com.easyagents.federation.sql.federation.FederationQueryMode;
|
||||
import com.easyagents.federation.sql.federation.FederationQueryScopeDefinition;
|
||||
import com.easyagents.federation.sql.federation.FederationSourceBindingDefinition;
|
||||
import com.easyagents.federation.sql.source.FederationDataSourceHandles;
|
||||
import com.easyagents.federation.sql.source.FederationSourceDefinition;
|
||||
import com.easyagents.federation.sql.source.RuntimeFingerprint;
|
||||
import com.easyagents.federation.sql.source.SourceApplyOptions;
|
||||
import com.easyagents.federation.sql.source.SourceId;
|
||||
import org.h2.jdbcx.JdbcDataSource;
|
||||
import org.junit.After;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Before;
|
||||
import org.junit.Test;
|
||||
|
||||
import java.sql.Connection;
|
||||
import java.sql.Statement;
|
||||
import java.util.LinkedHashMap;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
/**
|
||||
* 四源和六源 JDBC 多阶段 Join 集成测试。
|
||||
*/
|
||||
public class JdbcSixSourceQueryEngineTest {
|
||||
|
||||
private static final FederationExecutionPolicy SIX_SOURCE_POLICY =
|
||||
new FederationExecutionPolicy(6, 12, 2, 100_000,
|
||||
64L * 1024L * 1024L, 60_000);
|
||||
|
||||
private final Map<SourceId, JdbcDataSource> dataSources = new LinkedHashMap<>();
|
||||
private final Map<String, FederationSourceBindingDefinition> bindings = new LinkedHashMap<>();
|
||||
private FederationSqlEngine engine;
|
||||
|
||||
/**
|
||||
* 创建六个独立 H2 JDBC 数据源。
|
||||
*/
|
||||
@Before
|
||||
public void setUp() throws Exception {
|
||||
for (int index = 1; index <= 6; index++) {
|
||||
SourceId sourceId = new SourceId("source-" + index);
|
||||
JdbcDataSource dataSource = dataSource("six-source-" + index);
|
||||
execute(dataSource,
|
||||
"DROP TABLE IF EXISTS SUBJECT_METRIC",
|
||||
"CREATE TABLE SUBJECT_METRIC (SUBJECT_ID INT PRIMARY KEY, METRIC_VALUE INT NOT NULL)",
|
||||
"INSERT INTO SUBJECT_METRIC VALUES (1, " + (index * 10)
|
||||
+ "), (2, " + (index * 100) + ")");
|
||||
dataSources.put(sourceId, dataSource);
|
||||
bindings.put("S" + index, FederationSourceBindingDefinition.of(sourceId, 1));
|
||||
}
|
||||
RuntimeFingerprint fingerprint = new RuntimeFingerprint(
|
||||
"H2", "2", "H2 JDBC Driver", "2", "1");
|
||||
engine = FederationSqlEngines.builder()
|
||||
.dataSourceResolver(definition -> FederationDataSourceHandles.shared(
|
||||
dataSources.get(definition.sourceId()), fingerprint))
|
||||
.adapter(new JdbcFederationSqlAdapterProvider())
|
||||
.federationExecutionPolicy(SIX_SOURCE_POLICY)
|
||||
.maximumPlanCacheEntries(32)
|
||||
.build();
|
||||
for (SourceId sourceId : dataSources.keySet()) {
|
||||
engine.sources().apply(definition(sourceId), SourceApplyOptions.prewarmNow());
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 关闭 Engine。
|
||||
*/
|
||||
@After
|
||||
public void tearDown() {
|
||||
if (engine != null) {
|
||||
engine.close();
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证同一六源 Scope 中只引用四源时生成四个 Fragment。
|
||||
*/
|
||||
@Test
|
||||
public void shouldCompileFourReferencedSources() {
|
||||
FederationSqlPlan plan = engine.compile(SqlCompileRequest.of(
|
||||
selectSql(4), scope(SIX_SOURCE_POLICY)));
|
||||
|
||||
Assert.assertEquals(FederationQueryMode.FEDERATED, plan.queryMode());
|
||||
Assert.assertEquals(4, plan.referencedSources().size());
|
||||
Assert.assertEquals(4, plan.fragments().size());
|
||||
Assert.assertEquals(3, plan.joinOptimizations().size());
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证六个独立 JDBC 数据源可完成五阶段 Join 并返回稳定结果。
|
||||
*/
|
||||
@Test
|
||||
public void shouldExecuteJoinAcrossSixSources() {
|
||||
String sql = selectSql(6);
|
||||
FederationSqlPlan plan = engine.compile(SqlCompileRequest.of(
|
||||
sql, scope(SIX_SOURCE_POLICY)));
|
||||
|
||||
Assert.assertEquals(6, plan.referencedSources().size());
|
||||
Assert.assertEquals(6, plan.fragments().size());
|
||||
Assert.assertEquals(5, plan.joinOptimizations().size());
|
||||
|
||||
try (FederationResultCursor cursor = engine.query(
|
||||
SqlQueryCommand.of(sql, scope(SIX_SOURCE_POLICY), List.of()))) {
|
||||
Assert.assertTrue(cursor.next());
|
||||
Assert.assertEquals(1, cursor.getObject(1));
|
||||
Assert.assertEquals(10, cursor.getObject(2));
|
||||
Assert.assertEquals(60, cursor.getObject(7));
|
||||
Assert.assertTrue(cursor.next());
|
||||
Assert.assertEquals(2, cursor.getObject(1));
|
||||
Assert.assertEquals(600, cursor.getObject(7));
|
||||
Assert.assertFalse(cursor.next());
|
||||
Assert.assertEquals(6, cursor.metrics().fragments().size());
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证 Scope 收紧为四源时,六源 SQL 在编译期稳定拒绝。
|
||||
*/
|
||||
@Test
|
||||
public void shouldRejectSixSourcesWhenScopeLimitIsFour() {
|
||||
FederationExecutionPolicy fourSourcePolicy = new FederationExecutionPolicy(
|
||||
4, 12, 2, 100_000, 64L * 1024L * 1024L, 60_000);
|
||||
|
||||
try {
|
||||
engine.compile(SqlCompileRequest.of(selectSql(6), scope(fourSourcePolicy)));
|
||||
Assert.fail("expected source limit rejection");
|
||||
} catch (FederationSqlException exception) {
|
||||
Assert.assertEquals(FederationSqlErrorCode.FEDERATION_RESOURCE_LIMIT_EXCEEDED,
|
||||
exception.errorCode());
|
||||
}
|
||||
}
|
||||
|
||||
private FederationQueryScopeDefinition scope(FederationExecutionPolicy policy) {
|
||||
return new FederationQueryScopeDefinition(
|
||||
"six-source-scope", 1, bindings, "S1", policy);
|
||||
}
|
||||
|
||||
private static String selectSql(int sourceCount) {
|
||||
StringBuilder sql = new StringBuilder("SELECT s1.SUBJECT_ID");
|
||||
for (int index = 1; index <= sourceCount; index++) {
|
||||
sql.append(", s").append(index).append(".METRIC_VALUE");
|
||||
}
|
||||
sql.append(" FROM S1.APP.SUBJECT_METRIC s1");
|
||||
for (int index = 2; index <= sourceCount; index++) {
|
||||
sql.append(" JOIN S").append(index).append(".APP.SUBJECT_METRIC s")
|
||||
.append(index).append(" ON s").append(index)
|
||||
.append(".SUBJECT_ID = s1.SUBJECT_ID");
|
||||
}
|
||||
return sql.append(" ORDER BY s1.SUBJECT_ID").toString();
|
||||
}
|
||||
|
||||
private static JdbcDataSource dataSource(String name) {
|
||||
JdbcDataSource dataSource = new JdbcDataSource();
|
||||
dataSource.setURL("jdbc:h2:mem:" + name + ";DB_CLOSE_DELAY=-1");
|
||||
dataSource.setUser("sa");
|
||||
dataSource.setPassword("");
|
||||
return dataSource;
|
||||
}
|
||||
|
||||
private static void execute(JdbcDataSource dataSource, String... statements)
|
||||
throws Exception {
|
||||
try (Connection connection = dataSource.getConnection();
|
||||
Statement statement = connection.createStatement()) {
|
||||
for (String sql : statements) {
|
||||
statement.execute(sql);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private static FederationSourceDefinition definition(SourceId sourceId) {
|
||||
return new FederationSourceDefinition(
|
||||
sourceId, 1, JdbcFederationSqlAdapterProvider.ADAPTER_ID,
|
||||
List.of(new JdbcSchemaDefinition("APP", null, "PUBLIC")), Map.of());
|
||||
}
|
||||
}
|
||||
@@ -70,6 +70,7 @@ public final class DefaultFederationSqlEngine implements FederationSqlEngine {
|
||||
private final CalciteFederationSqlCompiler compiler;
|
||||
private final AdapterFederationStatisticsProvider automaticStatisticsProvider;
|
||||
private final NodeMemoryAdmissionController nodeMemoryAdmission;
|
||||
private final FederationFragmentScheduler fragmentScheduler;
|
||||
private final CalciteSqlCompleter completer = new CalciteSqlCompleter();
|
||||
private final QueryCancellationRegistry cancellations = new QueryCancellationRegistry();
|
||||
private final ScheduledThreadPoolExecutor deadlineScheduler = deadlineScheduler();
|
||||
@@ -297,6 +298,7 @@ public final class DefaultFederationSqlEngine implements FederationSqlEngine {
|
||||
this.nodeMemoryAdmission = new NodeMemoryAdmissionController(
|
||||
maximumNodeIntermediateBytes
|
||||
);
|
||||
this.fragmentScheduler = new FederationFragmentScheduler(executionPolicy);
|
||||
this.planCache = new BoundedPlanCache(
|
||||
maximumPlanCacheEntries,
|
||||
maximumConcurrentCompilations,
|
||||
@@ -650,7 +652,8 @@ public final class DefaultFederationSqlEngine implements FederationSqlEngine {
|
||||
registration,
|
||||
compiler.effectivePolicy(plan.queryScope()),
|
||||
metrics,
|
||||
queryDeadline
|
||||
queryDeadline,
|
||||
fragmentScheduler
|
||||
);
|
||||
queryCursor = new FederatedResultCursor(
|
||||
plan,
|
||||
@@ -1080,6 +1083,11 @@ public final class DefaultFederationSqlEngine implements FederationSqlEngine {
|
||||
} catch (RuntimeException exception) {
|
||||
failure = exception;
|
||||
}
|
||||
try {
|
||||
fragmentScheduler.close();
|
||||
} catch (RuntimeException exception) {
|
||||
failure = append(failure, exception);
|
||||
}
|
||||
try {
|
||||
sourceManager.close();
|
||||
} catch (RuntimeException exception) {
|
||||
|
||||
@@ -62,6 +62,7 @@ final class FederatedResultCursor
|
||||
this.session = session;
|
||||
this.metrics = metrics;
|
||||
this.maximumRows = maximumRows;
|
||||
session.prepareFragments();
|
||||
this.enumerator = plan.localBindable()
|
||||
.bind(new FederationDataContext(
|
||||
plan.localRootSchema(),
|
||||
|
||||
@@ -20,14 +20,21 @@ import java.time.LocalDateTime;
|
||||
import java.time.LocalTime;
|
||||
import java.time.OffsetDateTime;
|
||||
import java.time.ZoneOffset;
|
||||
import java.util.ArrayDeque;
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Set;
|
||||
import java.util.concurrent.CompletableFuture;
|
||||
import java.util.concurrent.ConcurrentHashMap;
|
||||
import java.util.concurrent.ExecutionException;
|
||||
import java.util.concurrent.Future;
|
||||
import java.util.concurrent.FutureTask;
|
||||
import java.util.concurrent.Semaphore;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.concurrent.TimeoutException;
|
||||
import java.util.concurrent.atomic.AtomicBoolean;
|
||||
import java.util.concurrent.atomic.AtomicReference;
|
||||
import org.apache.calcite.avatica.util.ByteString;
|
||||
import org.apache.calcite.linq4j.AbstractEnumerable;
|
||||
import org.apache.calcite.linq4j.Enumerable;
|
||||
@@ -50,9 +57,19 @@ final class FederationExecutionSession implements AutoCloseable {
|
||||
private final FederationExecutionPolicy policy;
|
||||
private final FederationQueryMetricsTracker metrics;
|
||||
private final QueryDeadline deadline;
|
||||
private final FederationFragmentScheduler fragmentScheduler;
|
||||
private final Semaphore fragmentSlots;
|
||||
private final Set<FederationResultCursor> openFragmentCursors = ConcurrentHashMap.newKeySet();
|
||||
private final Set<Future<?>> fragmentTasks = ConcurrentHashMap.newKeySet();
|
||||
private final Map<String, CompletableFuture<PreparedFragment>> fragmentPreparations =
|
||||
new ConcurrentHashMap<>();
|
||||
private final Object preparationLock = new Object();
|
||||
private final ArrayDeque<CompiledFederationFragment> pendingPreparations =
|
||||
new ArrayDeque<>();
|
||||
private final AtomicBoolean closed = new AtomicBoolean();
|
||||
private final AtomicReference<Throwable> preparationFailure = new AtomicReference<>();
|
||||
private int runningPreparations;
|
||||
private boolean preparationFailed;
|
||||
|
||||
/**
|
||||
* 创建查询执行会话。
|
||||
@@ -64,6 +81,7 @@ final class FederationExecutionSession implements AutoCloseable {
|
||||
* @param policy 有效资源策略
|
||||
* @param metrics 指标跟踪器
|
||||
* @param deadline 请求级统一截止时间
|
||||
* @param fragmentScheduler Engine 级有界分片调度器
|
||||
*/
|
||||
FederationExecutionSession(
|
||||
DefaultFederationSqlPlan plan,
|
||||
@@ -72,7 +90,8 @@ final class FederationExecutionSession implements AutoCloseable {
|
||||
QueryCancellationRegistry.QueryRegistration registration,
|
||||
FederationExecutionPolicy policy,
|
||||
FederationQueryMetricsTracker metrics,
|
||||
QueryDeadline deadline
|
||||
QueryDeadline deadline,
|
||||
FederationFragmentScheduler fragmentScheduler
|
||||
) {
|
||||
this.plan = plan;
|
||||
this.context = context;
|
||||
@@ -81,9 +100,41 @@ final class FederationExecutionSession implements AutoCloseable {
|
||||
this.policy = policy;
|
||||
this.metrics = metrics;
|
||||
this.deadline = deadline;
|
||||
this.fragmentScheduler = fragmentScheduler;
|
||||
this.fragmentSlots = new Semaphore(policy.maximumConcurrentFragments(), true);
|
||||
}
|
||||
|
||||
/**
|
||||
* 异步启动多源物理分片,在查询级中间结果预算内完成有界物化。
|
||||
*/
|
||||
void prepareFragments() {
|
||||
if (plan.fragments().size() < 2 || policy.maximumConcurrentFragments() < 2) {
|
||||
return;
|
||||
}
|
||||
try {
|
||||
for (FederationFragmentPlan fragment : plan.fragments()) {
|
||||
CompiledFederationFragment compiled = plan.compiledFragments().get(
|
||||
fragment.fragmentId()
|
||||
);
|
||||
if (compiled == null) {
|
||||
throw new FederationSqlException(
|
||||
FederationSqlErrorCode.EXECUTION_FAILED,
|
||||
"compiled fragment is missing: " + fragment.fragmentId()
|
||||
);
|
||||
}
|
||||
CompletableFuture<PreparedFragment> preparation = new CompletableFuture<>();
|
||||
fragmentPreparations.put(fragment.fragmentId(), preparation);
|
||||
synchronized (preparationLock) {
|
||||
pendingPreparations.addLast(compiled);
|
||||
}
|
||||
}
|
||||
schedulePreparations();
|
||||
} catch (RuntimeException | Error exception) {
|
||||
cancelPreparation();
|
||||
throw exception;
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 为 Calcite ScannableTable 创建一个按需执行物理分片的 Enumerable。
|
||||
*
|
||||
@@ -101,6 +152,13 @@ final class FederationExecutionSession implements AutoCloseable {
|
||||
return new AbstractEnumerable<>() {
|
||||
@Override
|
||||
public Enumerator<Object[]> enumerator() {
|
||||
CompletableFuture<PreparedFragment> preparation =
|
||||
fragmentPreparations.get(fragmentId);
|
||||
if (preparation != null) {
|
||||
PreparedFragment prepared = awaitPrepared(preparation);
|
||||
fragmentPreparations.remove(fragmentId, preparation);
|
||||
return prepared.enumerator();
|
||||
}
|
||||
return openFragment(compiled);
|
||||
}
|
||||
};
|
||||
@@ -138,6 +196,7 @@ final class FederationExecutionSession implements AutoCloseable {
|
||||
}
|
||||
registration.ensureNotCancelled();
|
||||
}
|
||||
throwPreparationFailure();
|
||||
deadline.ensureAllowed();
|
||||
}
|
||||
|
||||
@@ -166,38 +225,260 @@ final class FederationExecutionSession implements AutoCloseable {
|
||||
private Enumerator<Object[]> openFragment(CompiledFederationFragment compiled) {
|
||||
ensureExecutionAllowed();
|
||||
acquireFragmentSlot();
|
||||
FederationFragmentPlan fragment = compiled.plan();
|
||||
SourceRuntime runtime = runtimeSnapshot.runtime(fragment.bindingName());
|
||||
try {
|
||||
FederationResultCursor cursor = runtime.adapter().fragmentExecutor().execute(
|
||||
new FederationFragmentExecutionContext(
|
||||
context.queryId(),
|
||||
fragment.executableSql(),
|
||||
remapParameters(fragment, context.parameters()),
|
||||
fragmentOptions(context.options()),
|
||||
runtime.handle().dataSource(),
|
||||
runtime.compatibility(),
|
||||
runtime.definition().adapterOptions(),
|
||||
registration.statementLifecycle(),
|
||||
metrics.fragmentObserver(fragment.fragmentId()),
|
||||
deadline
|
||||
)
|
||||
);
|
||||
if (cursor == null || !context.queryId().equals(cursor.queryId())) {
|
||||
closeQuietly(cursor);
|
||||
throw new FederationSqlException(
|
||||
FederationSqlErrorCode.EXECUTION_FAILED,
|
||||
"adapter returned an invalid fragment cursor"
|
||||
);
|
||||
}
|
||||
openFragmentCursors.add(cursor);
|
||||
return new FragmentEnumerator(compiled, cursor);
|
||||
return new FragmentEnumerator(compiled, createFragmentCursor(compiled), true);
|
||||
} catch (RuntimeException exception) {
|
||||
fragmentSlots.release();
|
||||
throw exception;
|
||||
}
|
||||
}
|
||||
|
||||
private void schedulePreparations() {
|
||||
List<PreparationTask> tasks = new ArrayList<>();
|
||||
synchronized (preparationLock) {
|
||||
while (!closed.get()
|
||||
&& !preparationFailed
|
||||
&& runningPreparations < policy.maximumConcurrentFragments()
|
||||
&& !pendingPreparations.isEmpty()) {
|
||||
CompiledFederationFragment compiled = pendingPreparations.removeFirst();
|
||||
CompletableFuture<PreparedFragment> preparation =
|
||||
fragmentPreparations.get(
|
||||
compiled.plan().fragmentId()
|
||||
);
|
||||
if (preparation == null || preparation.isCancelled()) {
|
||||
continue;
|
||||
}
|
||||
runningPreparations++;
|
||||
tasks.add(new PreparationTask(compiled, preparation));
|
||||
}
|
||||
}
|
||||
for (PreparationTask task : tasks) {
|
||||
submitPreparation(task);
|
||||
}
|
||||
}
|
||||
|
||||
private void submitPreparation(PreparationTask preparationTask) {
|
||||
FutureTask<Void> task = new FutureTask<>(() -> {
|
||||
runPreparation(preparationTask);
|
||||
return null;
|
||||
}) {
|
||||
@Override
|
||||
protected void done() {
|
||||
fragmentTasks.remove(this);
|
||||
}
|
||||
};
|
||||
fragmentTasks.add(task);
|
||||
try {
|
||||
fragmentScheduler.execute(task);
|
||||
} catch (RuntimeException | Error exception) {
|
||||
fragmentTasks.remove(task);
|
||||
task.cancel(true);
|
||||
preparationTask.preparation().completeExceptionally(exception);
|
||||
failPreparations(exception);
|
||||
preparationFinished();
|
||||
}
|
||||
}
|
||||
|
||||
private void runPreparation(PreparationTask task) {
|
||||
try {
|
||||
task.preparation().complete(prepareFragment(task.compiled()));
|
||||
} catch (RuntimeException | Error exception) {
|
||||
task.preparation().completeExceptionally(exception);
|
||||
failPreparations(exception);
|
||||
} finally {
|
||||
preparationFinished();
|
||||
}
|
||||
}
|
||||
|
||||
private void preparationFinished() {
|
||||
boolean scheduleNext;
|
||||
synchronized (preparationLock) {
|
||||
runningPreparations--;
|
||||
scheduleNext = !closed.get() && !preparationFailed;
|
||||
}
|
||||
if (scheduleNext) {
|
||||
schedulePreparations();
|
||||
}
|
||||
}
|
||||
|
||||
private void failPreparations(Throwable failure) {
|
||||
if (!preparationFailure.compareAndSet(null, failure)) {
|
||||
return;
|
||||
}
|
||||
synchronized (preparationLock) {
|
||||
preparationFailed = true;
|
||||
pendingPreparations.clear();
|
||||
}
|
||||
for (CompletableFuture<PreparedFragment> preparation : fragmentPreparations.values()) {
|
||||
preparation.completeExceptionally(failure);
|
||||
}
|
||||
registration.terminateStatementsForCleanup();
|
||||
}
|
||||
|
||||
private void throwPreparationFailure() {
|
||||
Throwable failure = preparationFailure.get();
|
||||
if (failure == null) {
|
||||
return;
|
||||
}
|
||||
if (failure instanceof RuntimeException runtimeException) {
|
||||
throw runtimeException;
|
||||
}
|
||||
if (failure instanceof Error error) {
|
||||
throw error;
|
||||
}
|
||||
throw new FederationSqlException(
|
||||
FederationSqlErrorCode.EXECUTION_FAILED,
|
||||
"federation fragment preparation failed",
|
||||
failure
|
||||
);
|
||||
}
|
||||
|
||||
private PreparedFragment awaitPrepared(Future<PreparedFragment> preparation) {
|
||||
try {
|
||||
while (true) {
|
||||
ensureExecutionAllowed();
|
||||
long remaining = deadline.remainingNanos();
|
||||
if (remaining <= 0L) {
|
||||
ensureExecutionAllowed();
|
||||
}
|
||||
try {
|
||||
return preparation.get(
|
||||
Math.min(remaining, TimeUnit.MILLISECONDS.toNanos(50)),
|
||||
TimeUnit.NANOSECONDS
|
||||
);
|
||||
} catch (TimeoutException ignored) {
|
||||
// 周期性复核共享 Deadline 和取消状态。
|
||||
}
|
||||
}
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
cancelPreparation();
|
||||
throw new FederationSqlException(
|
||||
FederationSqlErrorCode.EXECUTION_FAILED,
|
||||
"waiting for federation fragment preparation was interrupted",
|
||||
exception
|
||||
);
|
||||
} catch (ExecutionException exception) {
|
||||
cancelPreparation();
|
||||
Throwable cause = exception.getCause();
|
||||
if (cause instanceof RuntimeException runtimeException) {
|
||||
throw runtimeException;
|
||||
}
|
||||
if (cause instanceof Error error) {
|
||||
throw error;
|
||||
}
|
||||
throw new FederationSqlException(
|
||||
FederationSqlErrorCode.EXECUTION_FAILED,
|
||||
"federation fragment preparation failed",
|
||||
cause
|
||||
);
|
||||
} catch (RuntimeException | Error exception) {
|
||||
cancelPreparation();
|
||||
throw exception;
|
||||
}
|
||||
}
|
||||
|
||||
private PreparedFragment prepareFragment(CompiledFederationFragment compiled) {
|
||||
ensureExecutionAllowed();
|
||||
FederationResultCursor cursor = null;
|
||||
try {
|
||||
cursor = createFragmentCursor(compiled);
|
||||
List<Object[]> rows = new ArrayList<>();
|
||||
while (true) {
|
||||
ensureExecutionAllowed();
|
||||
if (!cursor.next()) {
|
||||
releaseFragment(compiled, cursor, true, false);
|
||||
cursor = null;
|
||||
return new PreparedFragment(rows);
|
||||
}
|
||||
Object[] row = normalizeRow(cursor, compiled.rowType());
|
||||
recordFragmentIntermediate(compiled, row);
|
||||
rows.add(row);
|
||||
}
|
||||
} catch (RuntimeException | Error exception) {
|
||||
if (cursor != null) {
|
||||
releaseFragment(compiled, cursor, false, false);
|
||||
}
|
||||
throw exception;
|
||||
}
|
||||
}
|
||||
|
||||
private FederationResultCursor createFragmentCursor(
|
||||
CompiledFederationFragment compiled
|
||||
) {
|
||||
FederationFragmentPlan fragment = compiled.plan();
|
||||
SourceRuntime runtime = runtimeSnapshot.runtime(fragment.bindingName());
|
||||
FederationResultCursor cursor = runtime.adapter().fragmentExecutor().execute(
|
||||
new FederationFragmentExecutionContext(
|
||||
context.queryId(),
|
||||
fragment.executableSql(),
|
||||
remapParameters(fragment, context.parameters()),
|
||||
fragmentOptions(context.options()),
|
||||
runtime.handle().dataSource(),
|
||||
runtime.compatibility(),
|
||||
runtime.definition().adapterOptions(),
|
||||
registration.statementLifecycle(),
|
||||
metrics.fragmentObserver(fragment.fragmentId()),
|
||||
deadline
|
||||
)
|
||||
);
|
||||
if (cursor == null || !context.queryId().equals(cursor.queryId())) {
|
||||
closeQuietly(cursor);
|
||||
throw new FederationSqlException(
|
||||
FederationSqlErrorCode.EXECUTION_FAILED,
|
||||
"adapter returned an invalid fragment cursor"
|
||||
);
|
||||
}
|
||||
openFragmentCursors.add(cursor);
|
||||
try {
|
||||
ensureExecutionAllowed();
|
||||
return cursor;
|
||||
} catch (RuntimeException | Error exception) {
|
||||
openFragmentCursors.remove(cursor);
|
||||
closeQuietly(cursor);
|
||||
throw exception;
|
||||
}
|
||||
}
|
||||
|
||||
private void recordFragmentIntermediate(
|
||||
CompiledFederationFragment compiled,
|
||||
Object[] row
|
||||
) {
|
||||
FederationQueryMetricsTracker.IntermediateUsage usage =
|
||||
metrics.recordIntermediate(compiled.plan().fragmentId(), row);
|
||||
if (usage.rows() > policy.maximumIntermediateRows()
|
||||
|| usage.bytes() > policy.maximumIntermediateBytes()) {
|
||||
throw new FederationSqlException(
|
||||
FederationSqlErrorCode.FEDERATION_RESOURCE_LIMIT_EXCEEDED,
|
||||
"federation intermediate result limit was exceeded"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
private void cancelPreparation() {
|
||||
synchronized (preparationLock) {
|
||||
preparationFailed = true;
|
||||
pendingPreparations.clear();
|
||||
}
|
||||
for (Future<?> task : List.copyOf(fragmentTasks)) {
|
||||
task.cancel(true);
|
||||
}
|
||||
fragmentTasks.clear();
|
||||
for (FederationResultCursor cursor : List.copyOf(openFragmentCursors)) {
|
||||
openFragmentCursors.remove(cursor);
|
||||
try {
|
||||
cursor.close();
|
||||
} catch (RuntimeException ignored) {
|
||||
// 保留触发取消的原始分片异常;Statement 取消通道继续收口剩余资源。
|
||||
}
|
||||
}
|
||||
for (CompletableFuture<PreparedFragment> preparation : fragmentPreparations.values()) {
|
||||
preparation.cancel(true);
|
||||
}
|
||||
fragmentPreparations.clear();
|
||||
registration.terminateStatementsForCleanup();
|
||||
}
|
||||
|
||||
private void acquireFragmentSlot() {
|
||||
while (true) {
|
||||
ensureExecutionAllowed();
|
||||
@@ -259,6 +540,21 @@ final class FederationExecutionSession implements AutoCloseable {
|
||||
if (!closed.compareAndSet(false, true)) {
|
||||
return;
|
||||
}
|
||||
boolean terminationRequired = !fragmentTasks.isEmpty()
|
||||
|| !openFragmentCursors.isEmpty();
|
||||
synchronized (preparationLock) {
|
||||
terminationRequired = terminationRequired || runningPreparations > 0;
|
||||
preparationFailed = true;
|
||||
pendingPreparations.clear();
|
||||
}
|
||||
for (Future<?> task : List.copyOf(fragmentTasks)) {
|
||||
task.cancel(true);
|
||||
}
|
||||
fragmentTasks.clear();
|
||||
for (CompletableFuture<PreparedFragment> preparation : fragmentPreparations.values()) {
|
||||
preparation.cancel(true);
|
||||
}
|
||||
fragmentPreparations.clear();
|
||||
RuntimeException failure = null;
|
||||
for (FederationResultCursor cursor : List.copyOf(openFragmentCursors)) {
|
||||
try {
|
||||
@@ -273,6 +569,9 @@ final class FederationExecutionSession implements AutoCloseable {
|
||||
openFragmentCursors.remove(cursor);
|
||||
}
|
||||
}
|
||||
if (terminationRequired) {
|
||||
registration.terminateStatementsForCleanup();
|
||||
}
|
||||
if (failure != null) {
|
||||
throw failure;
|
||||
}
|
||||
@@ -282,15 +581,18 @@ final class FederationExecutionSession implements AutoCloseable {
|
||||
|
||||
private final CompiledFederationFragment compiled;
|
||||
private final FederationResultCursor cursor;
|
||||
private final boolean fragmentSlotHeld;
|
||||
private final AtomicBoolean released = new AtomicBoolean();
|
||||
private Object[] current;
|
||||
|
||||
private FragmentEnumerator(
|
||||
CompiledFederationFragment compiled,
|
||||
FederationResultCursor cursor
|
||||
FederationResultCursor cursor,
|
||||
boolean fragmentSlotHeld
|
||||
) {
|
||||
this.compiled = compiled;
|
||||
this.cursor = cursor;
|
||||
this.fragmentSlotHeld = fragmentSlotHeld;
|
||||
}
|
||||
|
||||
/** {@inheritDoc} */
|
||||
@@ -309,15 +611,7 @@ final class FederationExecutionSession implements AutoCloseable {
|
||||
return false;
|
||||
}
|
||||
Object[] row = normalizeRow(cursor, compiled.rowType());
|
||||
FederationQueryMetricsTracker.IntermediateUsage usage =
|
||||
metrics.recordIntermediate(compiled.plan().fragmentId(), row);
|
||||
if (usage.rows() > policy.maximumIntermediateRows()
|
||||
|| usage.bytes() > policy.maximumIntermediateBytes()) {
|
||||
throw new FederationSqlException(
|
||||
FederationSqlErrorCode.FEDERATION_RESOURCE_LIMIT_EXCEEDED,
|
||||
"federation intermediate result limit was exceeded"
|
||||
);
|
||||
}
|
||||
recordFragmentIntermediate(compiled, row);
|
||||
current = row;
|
||||
return true;
|
||||
} catch (RuntimeException exception) {
|
||||
@@ -346,7 +640,9 @@ final class FederationExecutionSession implements AutoCloseable {
|
||||
try {
|
||||
cursor.close();
|
||||
} finally {
|
||||
fragmentSlots.release();
|
||||
if (fragmentSlotHeld) {
|
||||
fragmentSlots.release();
|
||||
}
|
||||
if (exhausted) {
|
||||
metrics.finishFragment(compiled.plan().fragmentId());
|
||||
}
|
||||
@@ -354,6 +650,78 @@ final class FederationExecutionSession implements AutoCloseable {
|
||||
}
|
||||
}
|
||||
|
||||
private static final class PreparedFragment {
|
||||
|
||||
private final List<Object[]> rows;
|
||||
|
||||
private PreparedFragment(List<Object[]> rows) {
|
||||
this.rows = rows;
|
||||
}
|
||||
|
||||
private Enumerator<Object[]> enumerator() {
|
||||
return new Enumerator<>() {
|
||||
|
||||
private int index = -1;
|
||||
private Object[] current;
|
||||
|
||||
@Override
|
||||
public Object[] current() {
|
||||
return current;
|
||||
}
|
||||
|
||||
@Override
|
||||
public boolean moveNext() {
|
||||
int next = index + 1;
|
||||
if (next >= rows.size()) {
|
||||
current = null;
|
||||
return false;
|
||||
}
|
||||
index = next;
|
||||
current = rows.set(index, null);
|
||||
return true;
|
||||
}
|
||||
|
||||
@Override
|
||||
public void reset() {
|
||||
throw new UnsupportedOperationException(
|
||||
"fragment cursor cannot be reset"
|
||||
);
|
||||
}
|
||||
|
||||
@Override
|
||||
public void close() {
|
||||
rows.clear();
|
||||
current = null;
|
||||
}
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
private record PreparationTask(
|
||||
CompiledFederationFragment compiled,
|
||||
CompletableFuture<PreparedFragment> preparation
|
||||
) {
|
||||
}
|
||||
|
||||
private void releaseFragment(
|
||||
CompiledFederationFragment compiled,
|
||||
FederationResultCursor cursor,
|
||||
boolean exhausted,
|
||||
boolean fragmentSlotHeld
|
||||
) {
|
||||
openFragmentCursors.remove(cursor);
|
||||
try {
|
||||
cursor.close();
|
||||
} finally {
|
||||
if (fragmentSlotHeld) {
|
||||
fragmentSlots.release();
|
||||
}
|
||||
if (exhausted) {
|
||||
metrics.finishFragment(compiled.plan().fragmentId());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private static Object[] normalizeRow(
|
||||
FederationResultCursor cursor,
|
||||
RelDataType rowType
|
||||
|
||||
@@ -0,0 +1,118 @@
|
||||
package com.easyagents.federation.sql.runtime;
|
||||
|
||||
import com.easyagents.federation.sql.api.FederationSqlErrorCode;
|
||||
import com.easyagents.federation.sql.api.FederationSqlException;
|
||||
import com.easyagents.federation.sql.federation.FederationExecutionPolicy;
|
||||
import java.util.concurrent.ArrayBlockingQueue;
|
||||
import java.util.concurrent.Executor;
|
||||
import java.util.concurrent.RejectedExecutionException;
|
||||
import java.util.concurrent.ThreadPoolExecutor;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.concurrent.atomic.AtomicBoolean;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
|
||||
/**
|
||||
* Engine 级有界分片工作线程池;查询级并发度仍由有效执行策略单独控制。
|
||||
*/
|
||||
final class FederationFragmentScheduler implements Executor, AutoCloseable {
|
||||
|
||||
private static final int MAXIMUM_WORKERS = 32;
|
||||
private static final int MINIMUM_QUEUE_CAPACITY = 64;
|
||||
private static final int QUEUE_CAPACITY_PER_WORKER = 8;
|
||||
|
||||
private final ThreadPoolExecutor executor;
|
||||
private final AtomicBoolean closed = new AtomicBoolean();
|
||||
|
||||
/**
|
||||
* 创建 Engine 独占的有界分片调度器。
|
||||
*
|
||||
* @param policy Engine 联邦执行硬上限
|
||||
*/
|
||||
FederationFragmentScheduler(FederationExecutionPolicy policy) {
|
||||
int availableProcessors = Math.max(1, Runtime.getRuntime().availableProcessors());
|
||||
int workers = Math.min(
|
||||
MAXIMUM_WORKERS,
|
||||
Math.max(policy.maximumConcurrentFragments(), availableProcessors)
|
||||
);
|
||||
int queueCapacity = Math.max(
|
||||
MINIMUM_QUEUE_CAPACITY,
|
||||
workers * QUEUE_CAPACITY_PER_WORKER
|
||||
);
|
||||
this.executor = executor(workers, queueCapacity);
|
||||
}
|
||||
|
||||
FederationFragmentScheduler(int workers, int queueCapacity) {
|
||||
if (workers <= 0) {
|
||||
throw new IllegalArgumentException("workers must be positive");
|
||||
}
|
||||
if (queueCapacity <= 0) {
|
||||
throw new IllegalArgumentException("queueCapacity must be positive");
|
||||
}
|
||||
this.executor = executor(workers, queueCapacity);
|
||||
}
|
||||
|
||||
private static ThreadPoolExecutor executor(int workers, int queueCapacity) {
|
||||
AtomicInteger threadSequence = new AtomicInteger();
|
||||
return new ThreadPoolExecutor(
|
||||
workers,
|
||||
workers,
|
||||
0L,
|
||||
TimeUnit.MILLISECONDS,
|
||||
new ArrayBlockingQueue<>(queueCapacity),
|
||||
runnable -> {
|
||||
Thread thread = new Thread(
|
||||
runnable,
|
||||
"easy-agents-federation-fragment-" + threadSequence.incrementAndGet()
|
||||
);
|
||||
thread.setDaemon(true);
|
||||
return thread;
|
||||
},
|
||||
new ThreadPoolExecutor.AbortPolicy()
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
* 提交分片任务;队列饱和时明确拒绝,禁止在调用线程绕过 Worker 上限执行 JDBC。
|
||||
*
|
||||
* @param command 分片任务
|
||||
*/
|
||||
@Override
|
||||
public void execute(Runnable command) {
|
||||
if (closed.get()) {
|
||||
throw engineClosed();
|
||||
}
|
||||
try {
|
||||
executor.execute(command);
|
||||
} catch (RejectedExecutionException exception) {
|
||||
if (closed.get() || executor.isShutdown()) {
|
||||
throw engineClosed();
|
||||
}
|
||||
throw schedulerOverloaded(exception);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 中断排队任务并关闭工作线程。
|
||||
*/
|
||||
@Override
|
||||
public void close() {
|
||||
if (closed.compareAndSet(false, true)) {
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
|
||||
private static FederationSqlException engineClosed() {
|
||||
return new FederationSqlException(
|
||||
FederationSqlErrorCode.ENGINE_CLOSED,
|
||||
"federation fragment scheduler is closed"
|
||||
);
|
||||
}
|
||||
|
||||
private static FederationSqlException schedulerOverloaded(Throwable cause) {
|
||||
return new FederationSqlException(
|
||||
FederationSqlErrorCode.QUERY_ADMISSION_TIMEOUT,
|
||||
"federation fragment scheduler is saturated",
|
||||
cause
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -301,6 +301,13 @@ final class QueryCancellationRegistry implements AutoCloseable {
|
||||
return state.cancellationRequested();
|
||||
}
|
||||
|
||||
/**
|
||||
* 不改变查询终态,仅终止当前已登记的 Statement,用于提前关闭异步分片。
|
||||
*/
|
||||
void terminateStatementsForCleanup() {
|
||||
state.terminateStatements();
|
||||
}
|
||||
|
||||
/**
|
||||
* 已取消时立即抛出稳定错误。
|
||||
*/
|
||||
@@ -485,13 +492,18 @@ final class QueryCancellationRegistry implements AutoCloseable {
|
||||
}
|
||||
|
||||
private void terminate(QueryId queryId) {
|
||||
if (isClosed()) {
|
||||
return;
|
||||
}
|
||||
SQLException failure = null;
|
||||
if (cancelIssued.compareAndSet(false, true)) {
|
||||
try {
|
||||
statement.cancel();
|
||||
} catch (SQLException exception) {
|
||||
cancelIssued.set(false);
|
||||
failure = exception;
|
||||
if (!isClosed()) {
|
||||
cancelIssued.set(false);
|
||||
failure = exception;
|
||||
}
|
||||
}
|
||||
}
|
||||
// close() 同时覆盖 register 与 executeQuery 之间的 JDBC 取消空窗。
|
||||
@@ -499,11 +511,13 @@ final class QueryCancellationRegistry implements AutoCloseable {
|
||||
try {
|
||||
statement.close();
|
||||
} catch (SQLException exception) {
|
||||
closeIssued.set(false);
|
||||
if (failure == null) {
|
||||
failure = exception;
|
||||
} else {
|
||||
failure.addSuppressed(exception);
|
||||
if (!isClosed()) {
|
||||
closeIssued.set(false);
|
||||
if (failure == null) {
|
||||
failure = exception;
|
||||
} else {
|
||||
failure.addSuppressed(exception);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -515,6 +529,14 @@ final class QueryCancellationRegistry implements AutoCloseable {
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
private boolean isClosed() {
|
||||
try {
|
||||
return statement.isClosed();
|
||||
} catch (SQLException ignored) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private static FederationSqlException cancelled(QueryId queryId, Throwable cause) {
|
||||
|
||||
@@ -0,0 +1,118 @@
|
||||
package com.easyagents.federation.sql.runtime;
|
||||
|
||||
import com.easyagents.federation.sql.api.FederationSqlErrorCode;
|
||||
import com.easyagents.federation.sql.api.FederationSqlException;
|
||||
import com.easyagents.federation.sql.federation.FederationExecutionPolicy;
|
||||
import java.util.concurrent.CountDownLatch;
|
||||
import java.util.concurrent.ExecutorCompletionService;
|
||||
import java.util.concurrent.Future;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.concurrent.atomic.AtomicBoolean;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
|
||||
/**
|
||||
* Engine 级有界分片调度器测试。
|
||||
*/
|
||||
public class FederationFragmentSchedulerTest {
|
||||
|
||||
/**
|
||||
* 验证两个分片任务能够真实重叠执行,且调度器关闭后拒绝新任务。
|
||||
*
|
||||
* @throws Exception 等待任务失败
|
||||
*/
|
||||
@Test
|
||||
public void shouldExecuteFragmentTasksConcurrentlyAndRejectAfterClose() throws Exception {
|
||||
FederationFragmentScheduler scheduler = new FederationFragmentScheduler(
|
||||
FederationExecutionPolicy.basic()
|
||||
);
|
||||
CountDownLatch started = new CountDownLatch(2);
|
||||
CountDownLatch release = new CountDownLatch(1);
|
||||
AtomicInteger active = new AtomicInteger();
|
||||
AtomicInteger maximumActive = new AtomicInteger();
|
||||
ExecutorCompletionService<Void> completions = new ExecutorCompletionService<>(scheduler);
|
||||
try {
|
||||
Future<Void> first = completions.submit(
|
||||
() -> runBlockingTask(started, release, active, maximumActive)
|
||||
);
|
||||
Future<Void> second = completions.submit(
|
||||
() -> runBlockingTask(started, release, active, maximumActive)
|
||||
);
|
||||
|
||||
Assert.assertTrue(started.await(2, TimeUnit.SECONDS));
|
||||
Assert.assertEquals(2, maximumActive.get());
|
||||
release.countDown();
|
||||
first.get(2, TimeUnit.SECONDS);
|
||||
second.get(2, TimeUnit.SECONDS);
|
||||
} finally {
|
||||
release.countDown();
|
||||
scheduler.close();
|
||||
}
|
||||
|
||||
try {
|
||||
scheduler.execute(() -> { });
|
||||
Assert.fail("closed scheduler should reject new tasks");
|
||||
} catch (FederationSqlException exception) {
|
||||
Assert.assertEquals(FederationSqlErrorCode.ENGINE_CLOSED, exception.errorCode());
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证队列饱和时明确拒绝,且不会在提交线程绕过 Worker 上限执行任务。
|
||||
*
|
||||
* @throws Exception 等待任务失败
|
||||
*/
|
||||
@Test
|
||||
public void shouldRejectSaturationWithoutRunningTaskOnCaller() throws Exception {
|
||||
FederationFragmentScheduler scheduler = new FederationFragmentScheduler(1, 1);
|
||||
CountDownLatch firstStarted = new CountDownLatch(1);
|
||||
CountDownLatch releaseFirst = new CountDownLatch(1);
|
||||
CountDownLatch queuedCompleted = new CountDownLatch(1);
|
||||
AtomicBoolean rejectedTaskRan = new AtomicBoolean();
|
||||
try {
|
||||
scheduler.execute(() -> {
|
||||
firstStarted.countDown();
|
||||
try {
|
||||
releaseFirst.await(2, TimeUnit.SECONDS);
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
}
|
||||
});
|
||||
Assert.assertTrue(firstStarted.await(2, TimeUnit.SECONDS));
|
||||
scheduler.execute(queuedCompleted::countDown);
|
||||
|
||||
try {
|
||||
scheduler.execute(() -> rejectedTaskRan.set(true));
|
||||
Assert.fail("saturated scheduler should reject the task");
|
||||
} catch (FederationSqlException exception) {
|
||||
Assert.assertEquals(
|
||||
FederationSqlErrorCode.QUERY_ADMISSION_TIMEOUT,
|
||||
exception.errorCode()
|
||||
);
|
||||
}
|
||||
Assert.assertFalse(rejectedTaskRan.get());
|
||||
} finally {
|
||||
releaseFirst.countDown();
|
||||
Assert.assertTrue(queuedCompleted.await(2, TimeUnit.SECONDS));
|
||||
scheduler.close();
|
||||
}
|
||||
}
|
||||
|
||||
private static Void runBlockingTask(
|
||||
CountDownLatch started,
|
||||
CountDownLatch release,
|
||||
AtomicInteger active,
|
||||
AtomicInteger maximumActive
|
||||
) throws InterruptedException {
|
||||
int current = active.incrementAndGet();
|
||||
maximumActive.accumulateAndGet(current, Math::max);
|
||||
started.countDown();
|
||||
try {
|
||||
Assert.assertTrue(release.await(2, TimeUnit.SECONDS));
|
||||
return null;
|
||||
} finally {
|
||||
active.decrementAndGet();
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -634,7 +634,8 @@ public class Chain {
|
||||
NodeStateField
|
||||
.EXECUTION_ATTEMPT_KEY);
|
||||
}
|
||||
if (node.getCondition() == null) {
|
||||
if (node.getJoinMode() == NodeJoinMode.ANY
|
||||
&& node.getCondition() == null) {
|
||||
s.recordTrigger(triggerEdgeId);
|
||||
fields.add(NodeStateField.TRIGGER_COUNT);
|
||||
fields.add(NodeStateField.TRIGGER_EDGE_IDS);
|
||||
@@ -795,7 +796,7 @@ public class Chain {
|
||||
|
||||
private boolean shouldSkipNode(Node node, String edgeId) {
|
||||
NodeCondition condition = node.getCondition();
|
||||
if (condition == null) {
|
||||
if (node.getJoinMode() == NodeJoinMode.ANY && condition == null) {
|
||||
return false;
|
||||
}
|
||||
return executeWithLock(stateInstanceId, 10, TimeUnit.SECONDS, () -> {
|
||||
@@ -805,7 +806,11 @@ public class Chain {
|
||||
});
|
||||
|
||||
Map<String, Object> prevResult = Collections.emptyMap();
|
||||
boolean shouldSkipNode = !condition.check(this, newState, prevResult);
|
||||
boolean joinPending = node.getJoinMode() == NodeJoinMode.ALL
|
||||
&& !newState.isUpstreamFullyExecuted();
|
||||
boolean shouldSkipNode = joinPending
|
||||
|| (condition != null
|
||||
&& !condition.check(this, newState, prevResult));
|
||||
if (shouldSkipNode) {
|
||||
updateStateSafely(state -> {
|
||||
return state.addUncheckedNodeId(node.id)
|
||||
@@ -1746,16 +1751,91 @@ public class Chain {
|
||||
stateInstanceId,
|
||||
10L,
|
||||
TimeUnit.SECONDS,
|
||||
() -> {
|
||||
ChainState current =
|
||||
chainStateRepository.load(stateInstanceId);
|
||||
if (current == null
|
||||
|| current.getStatus() != ChainStatus.SUSPEND) {
|
||||
return false;
|
||||
}
|
||||
resumeSuspended(variables);
|
||||
return true;
|
||||
});
|
||||
() -> resumeSuspended(variables));
|
||||
}
|
||||
|
||||
private void validateResumeVariables(
|
||||
ChainState state, Map<String, Object> variables) {
|
||||
List<Parameter> parameters = state.getSuspendForParameters();
|
||||
if (parameters == null || parameters.isEmpty()) {
|
||||
return;
|
||||
}
|
||||
Map<String, Object> submitted = variables == null
|
||||
? Collections.emptyMap()
|
||||
: variables;
|
||||
Set<String> expectedKeys = parameters.stream()
|
||||
.map(Parameter::getName)
|
||||
.filter(Objects::nonNull)
|
||||
.collect(java.util.stream.Collectors.toCollection(
|
||||
LinkedHashSet::new));
|
||||
Set<String> extraKeys = new LinkedHashSet<>(submitted.keySet());
|
||||
extraKeys.removeAll(expectedKeys);
|
||||
if (!extraKeys.isEmpty()) {
|
||||
throw new ChainResumeException(
|
||||
"确认参数包含未声明字段");
|
||||
}
|
||||
for (Parameter parameter : parameters) {
|
||||
String name = parameter.getName();
|
||||
Object value = submitted.get(name);
|
||||
String label = StringUtil.getFirstWithText(
|
||||
parameter.getFormLabel(), name);
|
||||
if (!submitted.containsKey(name) || isBlankResumeValue(value)) {
|
||||
throw new ChainResumeException(
|
||||
"确认参数[" + label + "]不能为空");
|
||||
}
|
||||
validateResumeOption(parameter, value, label);
|
||||
}
|
||||
}
|
||||
|
||||
private void validateResumeOption(
|
||||
Parameter parameter, Object value, String label) {
|
||||
List<ParameterOption> options = parameter.getOptions();
|
||||
if (options == null || options.isEmpty()) {
|
||||
return;
|
||||
}
|
||||
Set<String> allowedValues = options.stream()
|
||||
.map(ParameterOption::getValue)
|
||||
.filter(Objects::nonNull)
|
||||
.collect(java.util.stream.Collectors.toCollection(
|
||||
LinkedHashSet::new));
|
||||
if ("checkbox".equals(parameter.getFormType())) {
|
||||
if (!(value instanceof Collection<?> selected)) {
|
||||
throw new ChainResumeException(
|
||||
"确认参数[" + label + "]必须提交字符串数组");
|
||||
}
|
||||
Set<String> unique = new LinkedHashSet<>();
|
||||
for (Object item : selected) {
|
||||
if (!(item instanceof String selectedValue)
|
||||
|| !allowedValues.contains(selectedValue)) {
|
||||
throw new ChainResumeException(
|
||||
"确认参数[" + label + "]包含未配置选项");
|
||||
}
|
||||
if (!unique.add(selectedValue)) {
|
||||
throw new ChainResumeException(
|
||||
"确认参数[" + label + "]不能重复选择同一选项");
|
||||
}
|
||||
}
|
||||
return;
|
||||
}
|
||||
if (!(value instanceof String selectedValue)) {
|
||||
throw new ChainResumeException(
|
||||
"确认参数[" + label + "]必须提交单个字符串值");
|
||||
}
|
||||
if (!allowedValues.contains(selectedValue)) {
|
||||
throw new ChainResumeException(
|
||||
"确认参数[" + label + "]包含未配置选项");
|
||||
}
|
||||
}
|
||||
|
||||
private boolean isBlankResumeValue(Object value) {
|
||||
if (value == null) {
|
||||
return true;
|
||||
}
|
||||
if (value instanceof String text) {
|
||||
return text.trim().isEmpty();
|
||||
}
|
||||
return value instanceof Collection<?> collection
|
||||
&& collection.isEmpty();
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -1772,36 +1852,51 @@ public class Chain {
|
||||
*
|
||||
* @param variables 恢复时注入的变量
|
||||
*/
|
||||
private void resumeSuspended(Map<String, Object> variables) {
|
||||
ChainState newState = updateStateSafely(state -> {
|
||||
if (variables != null) {
|
||||
state.getMemory().putAll(variables);
|
||||
return EnumSet.of(ChainStateField.MEMORY);
|
||||
} else {
|
||||
private boolean resumeSuspended(Map<String, Object> variables) {
|
||||
AtomicBoolean resumed = new AtomicBoolean(false);
|
||||
AtomicReference<Set<String>> suspendedNodeIds =
|
||||
new AtomicReference<>(Collections.emptySet());
|
||||
updateStateSafely(state -> {
|
||||
resumed.set(false);
|
||||
suspendedNodeIds.set(Collections.emptySet());
|
||||
if (state.getStatus() != ChainStatus.SUSPEND) {
|
||||
return null;
|
||||
}
|
||||
validateResumeVariables(state, variables);
|
||||
if (state.getSuspendNodeIds() != null) {
|
||||
suspendedNodeIds.set(
|
||||
new LinkedHashSet<>(state.getSuspendNodeIds()));
|
||||
}
|
||||
EnumSet<ChainStateField> updatedFields = EnumSet.of(
|
||||
ChainStateField.STATUS,
|
||||
ChainStateField.SUSPEND_NODE_IDS,
|
||||
ChainStateField.SUSPEND_FOR_PARAMETERS);
|
||||
if (variables != null && !variables.isEmpty()) {
|
||||
state.getMemory().putAll(variables);
|
||||
updatedFields.add(ChainStateField.MEMORY);
|
||||
}
|
||||
state.setStatus(ChainStatus.RUNNING);
|
||||
state.setSuspendNodeIds(null);
|
||||
state.setSuspendForParameters(null);
|
||||
resumed.set(true);
|
||||
return updatedFields;
|
||||
});
|
||||
if (!resumed.get()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
notifyEvent(new ChainResumeEvent(this, variables));
|
||||
setStatusAndNotifyEvent(ChainStatus.RUNNING);
|
||||
notifyEvent(new ChainStatusChangeEvent(
|
||||
this, ChainStatus.RUNNING, ChainStatus.SUSPEND));
|
||||
|
||||
Set<String> suspendNodeIds = newState.getSuspendNodeIds();
|
||||
if (suspendNodeIds != null && !suspendNodeIds.isEmpty()) {
|
||||
// 移除 suspend 状态,方便二次 suspend 时,不带有旧数据
|
||||
updateStateSafely(state -> {
|
||||
state.setSuspendNodeIds(null);
|
||||
state.setSuspendForParameters(null);
|
||||
return EnumSet.of(ChainStateField.SUSPEND_NODE_IDS, ChainStateField.SUSPEND_FOR_PARAMETERS);
|
||||
});
|
||||
|
||||
for (String id : suspendNodeIds) {
|
||||
Node node = definition.getNodeById(id);
|
||||
if (node == null) {
|
||||
throw new ChainException("Node not found: " + id);
|
||||
}
|
||||
scheduleNode(node, null, TriggerType.RESUME, 0L);
|
||||
for (String id : suspendedNodeIds.get()) {
|
||||
Node node = definition.getNodeById(id);
|
||||
if (node == null) {
|
||||
throw new ChainException("Node not found: " + id);
|
||||
}
|
||||
scheduleNode(node, null, TriggerType.RESUME, 0L);
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
public void resume() {
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
/**
|
||||
* Copyright (c) 2025-2026, Michael Yang 杨福海 (fuhai999@gmail.com).
|
||||
* <p>
|
||||
* Licensed under the GNU Lesser General Public License (LGPL) ,Version 3.0.
|
||||
*/
|
||||
package com.easyagents.flow.core.chain;
|
||||
|
||||
/**
|
||||
* 工作流挂起参数不满足当前恢复请求时抛出的异常。
|
||||
*/
|
||||
public class ChainResumeException extends ChainException {
|
||||
|
||||
public ChainResumeException(String message) {
|
||||
super(message);
|
||||
}
|
||||
}
|
||||
@@ -42,6 +42,7 @@ public abstract class Node implements Serializable {
|
||||
|
||||
protected NodeCondition condition;
|
||||
protected NodeValidator validator;
|
||||
protected NodeJoinMode joinMode = NodeJoinMode.ANY;
|
||||
|
||||
// 循环执行相关属性
|
||||
protected boolean loopEnable = false; // 是否启用循环执行
|
||||
@@ -70,6 +71,10 @@ public abstract class Node implements Serializable {
|
||||
}
|
||||
|
||||
public void setParentId(String parentId) {
|
||||
if (StringUtil.hasText(parentId) && getJoinMode() == NodeJoinMode.ALL) {
|
||||
throw new IllegalArgumentException(
|
||||
"joinMode 'all' is not supported for loop child nodes");
|
||||
}
|
||||
this.parentId = parentId;
|
||||
}
|
||||
|
||||
@@ -121,6 +126,31 @@ public abstract class Node implements Serializable {
|
||||
this.validator = validator;
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取节点的直接入边汇聚模式。
|
||||
*
|
||||
* <p>旧序列化对象缺少该字段时返回 {@link NodeJoinMode#ANY}。</p>
|
||||
*
|
||||
* @return 汇聚模式
|
||||
*/
|
||||
public NodeJoinMode getJoinMode() {
|
||||
return joinMode == null ? NodeJoinMode.ANY : joinMode;
|
||||
}
|
||||
|
||||
/**
|
||||
* 设置节点的直接入边汇聚模式。
|
||||
*
|
||||
* @param joinMode 汇聚模式
|
||||
*/
|
||||
public void setJoinMode(NodeJoinMode joinMode) {
|
||||
NodeJoinMode resolved = joinMode == null ? NodeJoinMode.ANY : joinMode;
|
||||
if (resolved == NodeJoinMode.ALL && StringUtil.hasText(parentId)) {
|
||||
throw new IllegalArgumentException(
|
||||
"joinMode 'all' is not supported for loop child nodes");
|
||||
}
|
||||
this.joinMode = resolved;
|
||||
}
|
||||
|
||||
// protected void addOutwardEdge(Edge edge) {
|
||||
// if (this.outwardEdges == null) {
|
||||
// this.outwardEdges = new ArrayList<>();
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
/**
|
||||
* Copyright (c) 2025-2026, Michael Yang 杨福海 (fuhai999@gmail.com).
|
||||
* <p>
|
||||
* Licensed under the GNU Lesser General Public License (LGPL) ,Version 3.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
* <p>
|
||||
* http://www.gnu.org/licenses/lgpl-3.0.txt
|
||||
* <p>
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
package com.easyagents.flow.core.chain;
|
||||
|
||||
import java.util.Locale;
|
||||
|
||||
/**
|
||||
* 多入边节点的触发汇聚模式。
|
||||
*/
|
||||
public enum NodeJoinMode {
|
||||
|
||||
/** 任意一条直接入边到达即可执行。 */
|
||||
ANY("any"),
|
||||
/** 全部直接入边到达后才执行。 */
|
||||
ALL("all");
|
||||
|
||||
private final String value;
|
||||
|
||||
NodeJoinMode(String value) {
|
||||
this.value = value;
|
||||
}
|
||||
|
||||
public String getValue() {
|
||||
return value;
|
||||
}
|
||||
|
||||
/**
|
||||
* 按工作流 JSON 值解析汇聚模式。
|
||||
*
|
||||
* @param value 配置值
|
||||
* @return 汇聚模式
|
||||
* @throws IllegalArgumentException 配置为空或不受支持
|
||||
*/
|
||||
public static NodeJoinMode ofValue(String value) {
|
||||
if (value == null || value.trim().isEmpty()) {
|
||||
throw new IllegalArgumentException("joinMode must be 'any' or 'all'");
|
||||
}
|
||||
String normalized = value.trim().toLowerCase(Locale.ROOT);
|
||||
for (NodeJoinMode mode : values()) {
|
||||
if (mode.value.equals(normalized)) {
|
||||
return mode;
|
||||
}
|
||||
}
|
||||
throw new IllegalArgumentException(
|
||||
"Unsupported joinMode: " + value + "; expected 'any' or 'all'");
|
||||
}
|
||||
}
|
||||
@@ -49,6 +49,11 @@ public class Parameter implements Serializable, Cloneable {
|
||||
*/
|
||||
protected List<Object> enums;
|
||||
|
||||
/**
|
||||
* 显示文案与实际值分离的结构化选项。
|
||||
*/
|
||||
protected List<ParameterOption> options;
|
||||
|
||||
/**
|
||||
* 用户输入的表单类型,例如:"input" "textarea" "select" "radio" "checkbox" 等等
|
||||
*/
|
||||
@@ -242,6 +247,14 @@ public class Parameter implements Serializable, Cloneable {
|
||||
}
|
||||
}
|
||||
|
||||
public List<ParameterOption> getOptions() {
|
||||
return options;
|
||||
}
|
||||
|
||||
public void setOptions(List<ParameterOption> options) {
|
||||
this.options = options;
|
||||
}
|
||||
|
||||
public String getFormType() {
|
||||
return formType;
|
||||
}
|
||||
@@ -298,6 +311,7 @@ public class Parameter implements Serializable, Cloneable {
|
||||
", flattenAggregation=" + flattenAggregation +
|
||||
", children=" + children +
|
||||
", enums=" + enums +
|
||||
", options=" + options +
|
||||
", formType='" + formType + '\'' +
|
||||
", formLabel='" + formLabel + '\'' +
|
||||
", formPlaceholder='" + formPlaceholder + '\'' +
|
||||
@@ -320,6 +334,13 @@ public class Parameter implements Serializable, Cloneable {
|
||||
clone.enums = new ArrayList<>(this.enums.size());
|
||||
clone.enums.addAll(this.enums);
|
||||
}
|
||||
if (this.options != null) {
|
||||
clone.options = new ArrayList<>(this.options.size());
|
||||
for (ParameterOption option : this.options) {
|
||||
clone.options.add(new ParameterOption(
|
||||
option.getLabel(), option.getValue()));
|
||||
}
|
||||
}
|
||||
return clone;
|
||||
} catch (CloneNotSupportedException e) {
|
||||
throw new AssertionError();
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
/**
|
||||
* Copyright (c) 2025-2026, Michael Yang 杨福海 (fuhai999@gmail.com).
|
||||
* <p>
|
||||
* Licensed under the GNU Lesser General Public License (LGPL) ,Version 3.0.
|
||||
*/
|
||||
package com.easyagents.flow.core.chain;
|
||||
|
||||
import java.io.Serializable;
|
||||
|
||||
/**
|
||||
* 用户输入参数的结构化选项。
|
||||
*/
|
||||
public class ParameterOption implements Serializable {
|
||||
|
||||
private static final long serialVersionUID = 1L;
|
||||
|
||||
private String label;
|
||||
private String value;
|
||||
|
||||
public ParameterOption() {
|
||||
}
|
||||
|
||||
public ParameterOption(String label, String value) {
|
||||
setLabel(label);
|
||||
setValue(value);
|
||||
}
|
||||
|
||||
public String getLabel() {
|
||||
return label;
|
||||
}
|
||||
|
||||
public void setLabel(String label) {
|
||||
this.label = trim(label);
|
||||
}
|
||||
|
||||
public String getValue() {
|
||||
return value;
|
||||
}
|
||||
|
||||
public void setValue(String value) {
|
||||
this.value = trim(value);
|
||||
}
|
||||
|
||||
private static String trim(String value) {
|
||||
return value == null ? null : value.trim();
|
||||
}
|
||||
|
||||
@Override
|
||||
public String toString() {
|
||||
return "ParameterOption{" +
|
||||
"label='" + label + '\'' +
|
||||
", value='" + value + '\'' +
|
||||
'}';
|
||||
}
|
||||
}
|
||||
@@ -18,6 +18,7 @@ package com.easyagents.flow.core.knowledge;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
public class KnowledgeManager {
|
||||
|
||||
@@ -51,4 +52,20 @@ public class KnowledgeManager {
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
/**
|
||||
* 将完整知识检索请求交给首个能够处理它的 Provider。
|
||||
*
|
||||
* @param request 检索请求
|
||||
* @return 节点输出;没有 Provider 能处理时返回 null
|
||||
*/
|
||||
public Map<String, Object> search(KnowledgeSearchRequest request) {
|
||||
for (KnowledgeProvider provider : providers) {
|
||||
Map<String, Object> result = provider.search(request);
|
||||
if (result != null) {
|
||||
return result;
|
||||
}
|
||||
}
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -15,6 +15,37 @@
|
||||
*/
|
||||
package com.easyagents.flow.core.knowledge;
|
||||
|
||||
import com.easyagents.flow.core.util.Maps;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
public interface KnowledgeProvider {
|
||||
Knowledge getKnowledge(Object id);
|
||||
|
||||
/**
|
||||
* 执行完整的知识库节点检索请求。
|
||||
*
|
||||
* <p>默认实现保留单知识库兼容。需要跨知识库汇总的业务 Provider
|
||||
* 应覆盖本方法并返回完整节点输出。</p>
|
||||
*
|
||||
* @param request 检索请求
|
||||
* @return 节点输出;当前 Provider 不支持该请求时返回 null
|
||||
*/
|
||||
default Map<String, Object> search(KnowledgeSearchRequest request) {
|
||||
if (request == null || request.getKnowledgeIds().size() != 1) {
|
||||
return null;
|
||||
}
|
||||
Object knowledgeId = request.getKnowledgeIds().get(0);
|
||||
Knowledge knowledge = getKnowledge(knowledgeId);
|
||||
if (knowledge == null) {
|
||||
return null;
|
||||
}
|
||||
List<Map<String, Object>> documents = knowledge.search(
|
||||
request.getKeyword(),
|
||||
request.getLimit(),
|
||||
request.getKnowledgeNode(),
|
||||
request.getChain());
|
||||
return Maps.of("documents", documents);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
/**
|
||||
* Copyright (c) 2025-2026, Michael Yang 杨福海 (fuhai999@gmail.com).
|
||||
* <p>
|
||||
* Licensed under the GNU Lesser General Public License (LGPL) ,Version 3.0.
|
||||
*/
|
||||
package com.easyagents.flow.core.knowledge;
|
||||
|
||||
import com.easyagents.flow.core.chain.Chain;
|
||||
import com.easyagents.flow.core.node.KnowledgeNode;
|
||||
|
||||
import java.util.Collections;
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
|
||||
/**
|
||||
* 工作流知识库节点的完整检索请求。
|
||||
*/
|
||||
public class KnowledgeSearchRequest {
|
||||
|
||||
private final List<Object> knowledgeIds;
|
||||
private final String keyword;
|
||||
private final int limit;
|
||||
private final String retrievalMode;
|
||||
private final KnowledgeNode knowledgeNode;
|
||||
private final Chain chain;
|
||||
|
||||
public KnowledgeSearchRequest(
|
||||
List<Object> knowledgeIds,
|
||||
String keyword,
|
||||
int limit,
|
||||
String retrievalMode,
|
||||
KnowledgeNode knowledgeNode,
|
||||
Chain chain) {
|
||||
this.knowledgeIds = knowledgeIds == null
|
||||
? Collections.emptyList()
|
||||
: Collections.unmodifiableList(new ArrayList<>(knowledgeIds));
|
||||
this.keyword = keyword;
|
||||
this.limit = limit;
|
||||
this.retrievalMode = retrievalMode;
|
||||
this.knowledgeNode = knowledgeNode;
|
||||
this.chain = chain;
|
||||
}
|
||||
|
||||
public List<Object> getKnowledgeIds() {
|
||||
return knowledgeIds;
|
||||
}
|
||||
|
||||
public String getKeyword() {
|
||||
return keyword;
|
||||
}
|
||||
|
||||
public int getLimit() {
|
||||
return limit;
|
||||
}
|
||||
|
||||
public String getRetrievalMode() {
|
||||
return retrievalMode;
|
||||
}
|
||||
|
||||
public KnowledgeNode getKnowledgeNode() {
|
||||
return knowledgeNode;
|
||||
}
|
||||
|
||||
public Chain getChain() {
|
||||
return chain;
|
||||
}
|
||||
}
|
||||
@@ -15,21 +15,55 @@
|
||||
*/
|
||||
package com.easyagents.flow.core.node;
|
||||
|
||||
|
||||
import com.easyagents.flow.core.chain.Chain;
|
||||
import com.easyagents.flow.core.chain.ChainSuspendException;
|
||||
import com.easyagents.flow.core.chain.DataType;
|
||||
import com.easyagents.flow.core.chain.Parameter;
|
||||
import com.easyagents.flow.core.chain.ParameterOption;
|
||||
import com.easyagents.flow.core.chain.RefType;
|
||||
import com.easyagents.flow.core.chain.repository.ChainStateField;
|
||||
import com.easyagents.flow.core.chain.runtime.Trigger;
|
||||
import com.easyagents.flow.core.chain.runtime.TriggerContext;
|
||||
import com.easyagents.flow.core.chain.runtime.TriggerType;
|
||||
|
||||
import java.util.*;
|
||||
import java.util.ArrayList;
|
||||
import java.util.Collections;
|
||||
import java.util.EnumSet;
|
||||
import java.util.HashSet;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Objects;
|
||||
import java.util.Set;
|
||||
|
||||
public class ConfirmNode extends BaseNode {
|
||||
private static final long serialVersionUID = 1L;
|
||||
|
||||
public static final String DEFAULT_OUTPUT_NAME = "selection";
|
||||
public static final int MAX_OPTIONS = 100;
|
||||
public static final int MAX_MESSAGE_LENGTH = 2000;
|
||||
public static final int MAX_OPTION_LENGTH = 200;
|
||||
public static final Set<String> SUPPORTED_CONFIGURATION_KEYS = Set.of(
|
||||
"condition",
|
||||
"description",
|
||||
"expand",
|
||||
"joinMode",
|
||||
"loopBreakCondition",
|
||||
"loopEnable",
|
||||
"loopIntervalMs",
|
||||
"maxLoopCount",
|
||||
"maxRetryCount",
|
||||
"message",
|
||||
"multiple",
|
||||
"options",
|
||||
"outputDefs",
|
||||
"resetRetryCountAfterNormal",
|
||||
"retryEnable",
|
||||
"retryIntervalMs",
|
||||
"title");
|
||||
|
||||
private String message;
|
||||
private List<Parameter> confirms;
|
||||
private boolean multiple;
|
||||
private List<String> options;
|
||||
|
||||
public String getMessage() {
|
||||
return message;
|
||||
@@ -39,115 +73,152 @@ public class ConfirmNode extends BaseNode {
|
||||
this.message = message;
|
||||
}
|
||||
|
||||
public List<Parameter> getConfirms() {
|
||||
return confirms;
|
||||
public boolean isMultiple() {
|
||||
return multiple;
|
||||
}
|
||||
|
||||
public void setConfirms(List<Parameter> confirms) {
|
||||
if (confirms != null) {
|
||||
for (Parameter confirm : confirms) {
|
||||
confirm.setRefType(RefType.INPUT);
|
||||
confirm.setRequired(true); // 必填,才能正确通过 getParameterValuesOnly 获取参数值
|
||||
confirm.setName(confirm.getName());
|
||||
}
|
||||
}
|
||||
this.confirms = confirms;
|
||||
public void setMultiple(boolean multiple) {
|
||||
this.multiple = multiple;
|
||||
}
|
||||
|
||||
public List<String> getOptions() {
|
||||
return options;
|
||||
}
|
||||
|
||||
public void setOptions(List<String> options) {
|
||||
this.options = options;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Map<String, Object> execute(Chain chain) {
|
||||
validateConfiguration();
|
||||
String outputName = resolveOutputName();
|
||||
Parameter parameter = buildParameter();
|
||||
|
||||
List<Parameter> confirmParameters = new ArrayList<>();
|
||||
addConfirmParameter(confirmParameters);
|
||||
|
||||
if (confirms != null) {
|
||||
for (Parameter confirm : confirms) {
|
||||
Parameter clone = confirm.clone();
|
||||
clone.setName(confirm.getName() + "__" + getId());
|
||||
clone.setRefType(RefType.INPUT);
|
||||
confirmParameters.add(clone);
|
||||
}
|
||||
// 确认值只能来自经过 Chain.resumeIfSuspended 校验后创建的恢复触发器。
|
||||
// 启动参数与普通节点内存中的同名值均不能绕过人工确认。
|
||||
if (!isValidatedResumeTrigger(chain)) {
|
||||
chain.updateStateSafely(state -> {
|
||||
if (state.getMemory().remove(parameter.getName()) == null) {
|
||||
return null;
|
||||
}
|
||||
return EnumSet.of(ChainStateField.MEMORY);
|
||||
});
|
||||
}
|
||||
|
||||
Map<String, Object> values;
|
||||
try {
|
||||
values = chain.getExecutionState()
|
||||
.resolveParameters(this, confirmParameters);
|
||||
// 移除 confirm 参数,方便在其他节点二次确认,或者在 for 循环中第二次获取
|
||||
.resolveParameters(this, Collections.singletonList(parameter));
|
||||
chain.updateStateSafely(state -> {
|
||||
for (Parameter confirmParameter : confirmParameters) {
|
||||
state.getMemory().remove(confirmParameter.getName());
|
||||
if (!state.getMemory().containsKey(parameter.getName())) {
|
||||
return null;
|
||||
}
|
||||
state.getMemory().remove(parameter.getName());
|
||||
return EnumSet.of(ChainStateField.MEMORY);
|
||||
});
|
||||
} catch (ChainSuspendException e) {
|
||||
} catch (ChainSuspendException exception) {
|
||||
chain.updateStateSafely(state -> {
|
||||
state.setMessage(message);
|
||||
return EnumSet.of(ChainStateField.MESSAGE);
|
||||
});
|
||||
|
||||
if (confirms != null) {
|
||||
List<Parameter> newParameters = new ArrayList<>();
|
||||
for (Parameter confirm : confirms) {
|
||||
Parameter clone = confirm.clone();
|
||||
clone.setName(confirm.getName() + "__" + getId());
|
||||
clone.setRefType(RefType.REF); // 固定为 REF
|
||||
newParameters.add(clone);
|
||||
}
|
||||
|
||||
// 获取参数值,不会触发 ChainSuspendException 错误
|
||||
Map<String, Object> parameterValues =
|
||||
chain.getExecutionState().resolveParameters(
|
||||
this,
|
||||
newParameters,
|
||||
null,
|
||||
true);
|
||||
|
||||
// 设置 enums,方便前端给用户进行选择
|
||||
for (Parameter confirmParameter : confirmParameters) {
|
||||
if (confirmParameter.getEnums() == null) {
|
||||
Object enumsObject = parameterValues.get(confirmParameter.getName());
|
||||
confirmParameter.setEnumsObject(enumsObject);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
throw e;
|
||||
throw exception;
|
||||
}
|
||||
|
||||
return Collections.singletonMap(
|
||||
outputName,
|
||||
values.get(parameter.getName()));
|
||||
}
|
||||
|
||||
Map<String, Object> results = new HashMap<>(values.size());
|
||||
values.forEach((key, value) -> {
|
||||
int index = key.lastIndexOf("__");
|
||||
if (index >= 0) {
|
||||
results.put(key.substring(0, index), value);
|
||||
} else {
|
||||
results.put(key, value);
|
||||
/**
|
||||
* 获取并校验当前节点配置的唯一输出名称。
|
||||
*
|
||||
* @return 用户配置的输出名称
|
||||
*/
|
||||
public String resolveOutputName() {
|
||||
if (outputDefs == null || outputDefs.size() != 1
|
||||
|| outputDefs.get(0) == null) {
|
||||
throw new IllegalArgumentException(
|
||||
"用户确认节点必须配置一个输出参数");
|
||||
}
|
||||
Parameter output = outputDefs.get(0);
|
||||
String outputName = output.getName();
|
||||
if (outputName == null || outputName.trim().isEmpty()) {
|
||||
throw new IllegalArgumentException(
|
||||
"用户确认节点输出参数名称不能为空");
|
||||
}
|
||||
DataType expectedType = multiple
|
||||
? DataType.Array_String
|
||||
: DataType.String;
|
||||
if (output.getDataType() != expectedType) {
|
||||
throw new IllegalArgumentException(
|
||||
"用户确认节点输出参数类型必须与选择方式一致");
|
||||
}
|
||||
return outputName;
|
||||
}
|
||||
|
||||
private boolean isValidatedResumeTrigger(Chain chain) {
|
||||
Trigger trigger = TriggerContext.getCurrentTrigger();
|
||||
return trigger != null
|
||||
&& trigger.getType() == TriggerType.RESUME
|
||||
&& Objects.equals(
|
||||
chain.getStateInstanceId(),
|
||||
trigger.getStateInstanceId())
|
||||
&& Objects.equals(getId(), trigger.getNodeId());
|
||||
}
|
||||
|
||||
public void validateConfiguration() {
|
||||
requireText(message, MAX_MESSAGE_LENGTH, "确认提示内容");
|
||||
if (options == null || options.isEmpty()) {
|
||||
throw new IllegalArgumentException("用户确认节点至少需要一个选项");
|
||||
}
|
||||
if (options.size() > MAX_OPTIONS) {
|
||||
throw new IllegalArgumentException(
|
||||
"用户确认节点最多支持 " + MAX_OPTIONS + " 个选项");
|
||||
}
|
||||
|
||||
Set<String> normalizedOptions = new HashSet<>();
|
||||
for (String option : options) {
|
||||
String normalized = requireText(option, MAX_OPTION_LENGTH, "选项内容");
|
||||
if (!normalizedOptions.add(normalized)) {
|
||||
throw new IllegalArgumentException("用户确认节点选项内容重复: " + normalized);
|
||||
}
|
||||
});
|
||||
|
||||
return results;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
private void addConfirmParameter(List<Parameter> parameters) {
|
||||
// “确认 和 取消” 的参数
|
||||
private Parameter buildParameter() {
|
||||
Parameter parameter = new Parameter();
|
||||
parameter.setId(DEFAULT_OUTPUT_NAME);
|
||||
parameter.setName(DEFAULT_OUTPUT_NAME + "__" + getId());
|
||||
parameter.setDataType(multiple ? DataType.Array_String : DataType.String);
|
||||
parameter.setRefType(RefType.INPUT);
|
||||
parameter.setId("confirm");
|
||||
parameter.setName("confirm__" + getId());
|
||||
parameter.setRequired(true);
|
||||
|
||||
List<Object> selectionData = new ArrayList<>();
|
||||
selectionData.add("yes");
|
||||
selectionData.add("no");
|
||||
|
||||
parameter.setEnums(selectionData);
|
||||
parameter.setContentType("text");
|
||||
parameter.setFormType("confirm");
|
||||
parameters.add(parameter);
|
||||
parameter.setFormType(multiple ? "checkbox" : "radio");
|
||||
parameter.setFormLabel("选择内容");
|
||||
parameter.setOptions(buildRuntimeOptions());
|
||||
return parameter;
|
||||
}
|
||||
|
||||
private List<ParameterOption> buildRuntimeOptions() {
|
||||
List<ParameterOption> runtimeOptions = new ArrayList<>(options.size());
|
||||
for (String option : options) {
|
||||
String normalized = option.trim();
|
||||
runtimeOptions.add(new ParameterOption(normalized, normalized));
|
||||
}
|
||||
return runtimeOptions;
|
||||
}
|
||||
|
||||
private static String requireText(
|
||||
String value, int maxLength, String fieldName) {
|
||||
String normalized = value == null ? "" : value.trim();
|
||||
if (normalized.isEmpty()) {
|
||||
throw new IllegalArgumentException(fieldName + "不能为空");
|
||||
}
|
||||
if (normalized.length() > maxLength) {
|
||||
throw new IllegalArgumentException(
|
||||
fieldName + "不能超过 " + maxLength + " 个字符");
|
||||
}
|
||||
return normalized;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -17,15 +17,15 @@ package com.easyagents.flow.core.node;
|
||||
|
||||
import com.easyagents.flow.core.chain.Chain;
|
||||
import com.easyagents.flow.core.chain.ChainState;
|
||||
import com.easyagents.flow.core.knowledge.Knowledge;
|
||||
import com.easyagents.flow.core.knowledge.KnowledgeManager;
|
||||
import com.easyagents.flow.core.util.Maps;
|
||||
import com.easyagents.flow.core.knowledge.KnowledgeSearchRequest;
|
||||
import com.easyagents.flow.core.util.StringUtil;
|
||||
import com.easyagents.flow.core.util.TextTemplate;
|
||||
import org.slf4j.Logger;
|
||||
|
||||
import java.util.Arrays;
|
||||
import java.util.ArrayList;
|
||||
import java.util.Collections;
|
||||
import java.util.LinkedHashSet;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
@@ -36,6 +36,7 @@ public class KnowledgeNode extends BaseNode {
|
||||
private static final Logger logger = org.slf4j.LoggerFactory.getLogger(KnowledgeNode.class);
|
||||
|
||||
private Object knowledgeId;
|
||||
private List<Object> knowledgeIds = new ArrayList<>();
|
||||
private String keyword;
|
||||
private String limit;
|
||||
private String retrievalMode = "HYBRID";
|
||||
@@ -48,6 +49,32 @@ public class KnowledgeNode extends BaseNode {
|
||||
this.knowledgeId = knowledgeId;
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取规范化的知识库集合,兼容历史单值字段。
|
||||
*
|
||||
* @return 去重后的知识库 ID
|
||||
*/
|
||||
public List<Object> getKnowledgeIds() {
|
||||
if (knowledgeIds != null && !knowledgeIds.isEmpty()) {
|
||||
return Collections.unmodifiableList(knowledgeIds);
|
||||
}
|
||||
return knowledgeId == null
|
||||
? Collections.emptyList()
|
||||
: Collections.singletonList(knowledgeId);
|
||||
}
|
||||
|
||||
public void setKnowledgeIds(List<?> knowledgeIds) {
|
||||
LinkedHashSet<Object> normalized = new LinkedHashSet<>();
|
||||
if (knowledgeIds != null) {
|
||||
for (Object id : knowledgeIds) {
|
||||
if (id != null && StringUtil.hasText(String.valueOf(id))) {
|
||||
normalized.add(id);
|
||||
}
|
||||
}
|
||||
}
|
||||
this.knowledgeIds = new ArrayList<>(normalized);
|
||||
}
|
||||
|
||||
public String getKeyword() {
|
||||
return keyword;
|
||||
}
|
||||
@@ -88,25 +115,44 @@ public class KnowledgeNode extends BaseNode {
|
||||
if (StringUtil.hasText(realLimitString)) {
|
||||
try {
|
||||
realLimit = Integer.parseInt(realLimitString);
|
||||
} catch (Exception e) {
|
||||
logger.error(e.toString(), e);
|
||||
} catch (NumberFormatException exception) {
|
||||
throw new IllegalArgumentException(
|
||||
"知识库节点最终返回条数必须为正整数", exception);
|
||||
}
|
||||
}
|
||||
|
||||
Knowledge knowledge = KnowledgeManager.getInstance().getKnowledge(knowledgeId);
|
||||
|
||||
if (knowledge == null) {
|
||||
return Collections.emptyMap();
|
||||
if (realLimit <= 0) {
|
||||
throw new IllegalArgumentException(
|
||||
"知识库节点最终返回条数必须为正整数");
|
||||
}
|
||||
|
||||
List<Map<String, Object>> result = knowledge.search(realKeyword, realLimit, this, chain);
|
||||
return Maps.of("documents", result);
|
||||
List<Object> resolvedKnowledgeIds = getKnowledgeIds();
|
||||
if (resolvedKnowledgeIds.isEmpty()) {
|
||||
throw new IllegalArgumentException("知识库节点至少需要选择一个知识库");
|
||||
}
|
||||
if (resolvedKnowledgeIds.size() > 1
|
||||
&& !"VECTOR".equalsIgnoreCase(retrievalMode)) {
|
||||
throw new IllegalArgumentException("多知识库检索仅支持 VECTOR 模式");
|
||||
}
|
||||
|
||||
Map<String, Object> result = KnowledgeManager.getInstance().search(
|
||||
new KnowledgeSearchRequest(
|
||||
resolvedKnowledgeIds,
|
||||
realKeyword,
|
||||
realLimit,
|
||||
retrievalMode,
|
||||
this,
|
||||
chain));
|
||||
if (result == null) {
|
||||
throw new IllegalStateException("没有可用的知识库 Provider");
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
@Override
|
||||
public String toString() {
|
||||
return "KnowledgeNode{" +
|
||||
"knowledgeId=" + knowledgeId +
|
||||
", knowledgeIds=" + knowledgeIds +
|
||||
", keyword='" + keyword + '\'' +
|
||||
", limit='" + limit + '\'' +
|
||||
", retrievalMode='" + retrievalMode + '\'' +
|
||||
|
||||
@@ -21,6 +21,7 @@ import com.alibaba.fastjson.JSONObject;
|
||||
import com.easyagents.flow.core.chain.DataType;
|
||||
import com.easyagents.flow.core.chain.JsCodeCondition;
|
||||
import com.easyagents.flow.core.chain.Node;
|
||||
import com.easyagents.flow.core.chain.NodeJoinMode;
|
||||
import com.easyagents.flow.core.chain.Parameter;
|
||||
import com.easyagents.flow.core.chain.RefType;
|
||||
import com.easyagents.flow.core.node.BaseNode;
|
||||
@@ -123,6 +124,10 @@ public abstract class BaseNodeParser<T extends BaseNode> implements NodeParser<T
|
||||
|
||||
if (!data.isEmpty()) {
|
||||
|
||||
if (data.containsKey("joinMode")) {
|
||||
node.setJoinMode(NodeJoinMode.ofValue(data.getString("joinMode")));
|
||||
}
|
||||
|
||||
addParameters(node, data);
|
||||
addOutputDefs(node, data);
|
||||
|
||||
|
||||
@@ -15,11 +15,12 @@
|
||||
*/
|
||||
package com.easyagents.flow.core.parser.impl;
|
||||
|
||||
import com.alibaba.fastjson.JSONArray;
|
||||
import com.alibaba.fastjson.JSONObject;
|
||||
import com.easyagents.flow.core.chain.Parameter;
|
||||
import com.easyagents.flow.core.node.ConfirmNode;
|
||||
import com.easyagents.flow.core.parser.BaseNodeParser;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
|
||||
public class ConfirmNodeParser extends BaseNodeParser<ConfirmNode> {
|
||||
@@ -28,12 +29,36 @@ public class ConfirmNodeParser extends BaseNodeParser<ConfirmNode> {
|
||||
public ConfirmNode doParse(JSONObject root, JSONObject data, JSONObject chainJSONObject) {
|
||||
|
||||
ConfirmNode confirmNode = new ConfirmNode();
|
||||
confirmNode.setMessage(data.getString("message"));
|
||||
|
||||
List<Parameter> confirms = getParameters(data, "confirms");
|
||||
if (confirms != null && !confirms.isEmpty()) {
|
||||
confirmNode.setConfirms(confirms);
|
||||
for (String key : data.keySet()) {
|
||||
if (!ConfirmNode.SUPPORTED_CONFIGURATION_KEYS.contains(key)) {
|
||||
throw new IllegalArgumentException(
|
||||
"用户确认节点包含无效配置字段: " + key);
|
||||
}
|
||||
}
|
||||
Object message = data.get("message");
|
||||
if (!(message instanceof String)) {
|
||||
throw new IllegalArgumentException("用户确认节点提示内容必须为字符串");
|
||||
}
|
||||
confirmNode.setMessage((String) message);
|
||||
|
||||
Object multiple = data.get("multiple");
|
||||
if (!(multiple instanceof Boolean)) {
|
||||
throw new IllegalArgumentException("用户确认节点选择方式必须为布尔值");
|
||||
}
|
||||
confirmNode.setMultiple((Boolean) multiple);
|
||||
|
||||
Object optionsValue = data.get("options");
|
||||
if (!(optionsValue instanceof JSONArray options)) {
|
||||
throw new IllegalArgumentException("用户确认节点选项必须为数组");
|
||||
}
|
||||
List<String> confirmOptions = new ArrayList<>(options.size());
|
||||
for (Object option : options) {
|
||||
if (!(option instanceof String)) {
|
||||
throw new IllegalArgumentException("用户确认节点选项内容必须为字符串");
|
||||
}
|
||||
confirmOptions.add((String) option);
|
||||
}
|
||||
confirmNode.setOptions(confirmOptions);
|
||||
|
||||
return confirmNode;
|
||||
}
|
||||
|
||||
@@ -16,15 +16,41 @@
|
||||
package com.easyagents.flow.core.parser.impl;
|
||||
|
||||
import com.alibaba.fastjson.JSONObject;
|
||||
import com.alibaba.fastjson.JSONArray;
|
||||
import com.easyagents.flow.core.node.KnowledgeNode;
|
||||
import com.easyagents.flow.core.parser.BaseNodeParser;
|
||||
|
||||
import java.util.ArrayList;
|
||||
|
||||
public class KnowledgeNodeParser extends BaseNodeParser<KnowledgeNode> {
|
||||
|
||||
@Override
|
||||
public KnowledgeNode doParse(JSONObject root, JSONObject data, JSONObject chainJSONObject) {
|
||||
KnowledgeNode knowledgeNode = new KnowledgeNode();
|
||||
knowledgeNode.setKnowledgeId(data.get("knowledgeId"));
|
||||
if (data.containsKey("knowledgeIds")) {
|
||||
Object rawIds = data.get("knowledgeIds");
|
||||
if (!(rawIds instanceof JSONArray)) {
|
||||
throw new IllegalArgumentException("knowledgeIds 必须为数组");
|
||||
}
|
||||
JSONArray ids = (JSONArray) rawIds;
|
||||
if (ids.isEmpty()) {
|
||||
throw new IllegalArgumentException("knowledgeIds 不能为空");
|
||||
}
|
||||
java.util.LinkedHashSet<String> normalized =
|
||||
new java.util.LinkedHashSet<>();
|
||||
for (Object id : ids) {
|
||||
String value = id == null ? null : String.valueOf(id).trim();
|
||||
if (!com.easyagents.flow.core.util.StringUtil.hasText(value)) {
|
||||
throw new IllegalArgumentException("knowledgeIds 不能包含空值");
|
||||
}
|
||||
if (!normalized.add(value)) {
|
||||
throw new IllegalArgumentException("knowledgeIds 不能包含重复值");
|
||||
}
|
||||
}
|
||||
knowledgeNode.setKnowledgeIds(new ArrayList<>(normalized));
|
||||
} else {
|
||||
knowledgeNode.setKnowledgeId(data.get("knowledgeId"));
|
||||
}
|
||||
knowledgeNode.setLimit(data.getString("limit"));
|
||||
knowledgeNode.setKeyword(data.getString("keyword"));
|
||||
knowledgeNode.setRetrievalMode(data.getString("retrievalMode"));
|
||||
|
||||
@@ -15,10 +15,15 @@
|
||||
*/
|
||||
package com.easyagents.flow.core.test;
|
||||
|
||||
import com.alibaba.fastjson.JSONArray;
|
||||
import com.alibaba.fastjson.JSONObject;
|
||||
import com.easyagents.flow.core.chain.Chain;
|
||||
import com.easyagents.flow.core.chain.ChainDefinition;
|
||||
import com.easyagents.flow.core.chain.ChainStatus;
|
||||
import com.easyagents.flow.core.chain.DataType;
|
||||
import com.easyagents.flow.core.chain.Edge;
|
||||
import com.easyagents.flow.core.chain.Parameter;
|
||||
import com.easyagents.flow.core.chain.RefType;
|
||||
import com.easyagents.flow.core.chain.ChainState;
|
||||
import com.easyagents.flow.core.chain.event.ChainEndEvent;
|
||||
import com.easyagents.flow.core.chain.repository.ChainDefinitionSnapshotRepository;
|
||||
@@ -35,6 +40,7 @@ import com.easyagents.flow.core.node.EndNode;
|
||||
import com.easyagents.flow.core.node.BaseNode;
|
||||
import com.easyagents.flow.core.node.ConfirmNode;
|
||||
import com.easyagents.flow.core.node.StartNode;
|
||||
import com.easyagents.flow.core.parser.ChainParser;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
|
||||
@@ -61,6 +67,127 @@ import java.util.concurrent.atomic.AtomicReference;
|
||||
*/
|
||||
public class ChainExecutorConcurrencyTest {
|
||||
|
||||
/**
|
||||
* 验证启动变量不能伪造确认节点的恢复参数并绕过人工确认。
|
||||
*/
|
||||
@Test
|
||||
public void shouldSuspendConfirmNodeDespitePrefilledStartVariable()
|
||||
throws Exception {
|
||||
ScheduledExecutorService schedulerPool =
|
||||
Executors.newSingleThreadScheduledExecutor();
|
||||
ExecutorService workerPool = Executors.newFixedThreadPool(2);
|
||||
TriggerScheduler triggerScheduler = new TriggerScheduler(
|
||||
new InMemoryTriggerStore(), schedulerPool, workerPool, 10L);
|
||||
ChainDefinition definition = createConfirmDefinition();
|
||||
InMemoryChainStateRepository stateRepository =
|
||||
new InMemoryChainStateRepository();
|
||||
ChainExecutor executor = new ChainExecutor(
|
||||
ignored -> definition,
|
||||
stateRepository,
|
||||
new InMemoryNodeStateRepository(),
|
||||
triggerScheduler);
|
||||
|
||||
try {
|
||||
String instanceId = executor.executeAsync(
|
||||
definition.getId(),
|
||||
Map.of("selection__confirm", "未配置值"));
|
||||
|
||||
ChainState state = awaitStatus(
|
||||
stateRepository, instanceId, ChainStatus.SUSPEND);
|
||||
Assert.assertFalse(
|
||||
state.getMemory().containsKey("selection__confirm"));
|
||||
Assert.assertEquals(
|
||||
"selection__confirm",
|
||||
state.getSuspendForParameters().get(0).getName());
|
||||
} finally {
|
||||
triggerScheduler.shutdown();
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证设计器最终契约可解析、挂起、恢复,并把用户选择按配置名称交给结束节点。
|
||||
*/
|
||||
@Test
|
||||
public void shouldFlowConfiguredConfirmOutputToEndNode()
|
||||
throws Exception {
|
||||
ScheduledExecutorService schedulerPool =
|
||||
Executors.newSingleThreadScheduledExecutor();
|
||||
ExecutorService workerPool = Executors.newFixedThreadPool(2);
|
||||
TriggerScheduler triggerScheduler = new TriggerScheduler(
|
||||
new InMemoryTriggerStore(), schedulerPool, workerPool, 10L);
|
||||
ChainDefinition definition = createParsedConfirmDefinition();
|
||||
InMemoryChainStateRepository stateRepository =
|
||||
new InMemoryChainStateRepository();
|
||||
ChainExecutor executor = new ChainExecutor(
|
||||
ignored -> definition,
|
||||
stateRepository,
|
||||
new InMemoryNodeStateRepository(),
|
||||
triggerScheduler);
|
||||
|
||||
try {
|
||||
String instanceId = executor.executeAsync(
|
||||
definition.getId(), Collections.emptyMap());
|
||||
ChainState suspended = awaitStatus(
|
||||
stateRepository, instanceId, ChainStatus.SUSPEND);
|
||||
|
||||
Assert.assertEquals(
|
||||
"selection__confirm",
|
||||
suspended.getSuspendForParameters().get(0).getName());
|
||||
Assert.assertTrue(executor.resumeAsyncIfSuspended(
|
||||
instanceId,
|
||||
Map.of("selection__confirm", "继续")));
|
||||
|
||||
ChainState completed = awaitStatus(
|
||||
stateRepository, instanceId, ChainStatus.SUCCEEDED);
|
||||
Assert.assertEquals("继续", completed.getExecuteResult().get("result"));
|
||||
} finally {
|
||||
triggerScheduler.shutdown();
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证多选确认结果以字符串数组形式流转到结束节点。
|
||||
*/
|
||||
@Test
|
||||
public void shouldFlowMultipleConfirmOutputToEndNode()
|
||||
throws Exception {
|
||||
ScheduledExecutorService schedulerPool =
|
||||
Executors.newSingleThreadScheduledExecutor();
|
||||
ExecutorService workerPool = Executors.newFixedThreadPool(2);
|
||||
TriggerScheduler triggerScheduler = new TriggerScheduler(
|
||||
new InMemoryTriggerStore(), schedulerPool, workerPool, 10L);
|
||||
ChainDefinition definition = createParsedConfirmDefinition(true);
|
||||
InMemoryChainStateRepository stateRepository =
|
||||
new InMemoryChainStateRepository();
|
||||
ChainExecutor executor = new ChainExecutor(
|
||||
ignored -> definition,
|
||||
stateRepository,
|
||||
new InMemoryNodeStateRepository(),
|
||||
triggerScheduler);
|
||||
|
||||
try {
|
||||
String instanceId = executor.executeAsync(
|
||||
definition.getId(), Collections.emptyMap());
|
||||
ChainState suspended = awaitStatus(
|
||||
stateRepository, instanceId, ChainStatus.SUSPEND);
|
||||
|
||||
Assert.assertEquals(
|
||||
DataType.Array_String,
|
||||
suspended.getSuspendForParameters().get(0).getDataType());
|
||||
List<String> selection = List.of("继续", "停止");
|
||||
Assert.assertTrue(executor.resumeAsyncIfSuspended(
|
||||
instanceId,
|
||||
Map.of("selection__confirm", selection)));
|
||||
|
||||
ChainState completed = awaitStatus(
|
||||
stateRepository, instanceId, ChainStatus.SUCCEEDED);
|
||||
Assert.assertEquals(
|
||||
selection, completed.getExecuteResult().get("result"));
|
||||
} finally {
|
||||
triggerScheduler.shutdown();
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证同步 Tool 入口遇到人工挂起会快速失败,不会无限占用调用线程。
|
||||
*
|
||||
@@ -563,6 +690,12 @@ public class ChainExecutorConcurrencyTest {
|
||||
start.setId("start");
|
||||
ConfirmNode confirm = new ConfirmNode();
|
||||
confirm.setId("confirm");
|
||||
confirm.setMessage("请选择是否继续");
|
||||
confirm.setOptions(List.of("继续", "停止"));
|
||||
confirm.setOutputDefs(Collections.singletonList(
|
||||
new Parameter(
|
||||
ConfirmNode.DEFAULT_OUTPUT_NAME,
|
||||
DataType.String)));
|
||||
EndNode end = new EndNode();
|
||||
end.setId("end");
|
||||
Edge first = new Edge();
|
||||
@@ -581,6 +714,91 @@ public class ChainExecutorConcurrencyTest {
|
||||
return definition;
|
||||
}
|
||||
|
||||
private ChainDefinition createParsedConfirmDefinition() {
|
||||
return createParsedConfirmDefinition(false);
|
||||
}
|
||||
|
||||
private ChainDefinition createParsedConfirmDefinition(boolean multiple) {
|
||||
JSONArray nodes = new JSONArray();
|
||||
nodes.add(nodeJson("start", "startNode", new JSONObject()));
|
||||
|
||||
JSONObject confirmData = new JSONObject();
|
||||
confirmData.put("message", "请选择是否继续");
|
||||
confirmData.put("multiple", multiple);
|
||||
confirmData.put("options", new JSONArray(List.of("继续", "停止")));
|
||||
String outputType = multiple ? "Array<String>" : "String";
|
||||
confirmData.put("outputDefs", new JSONArray(List.of(
|
||||
parameterJson("templateType", outputType, null))));
|
||||
nodes.add(nodeJson("confirm", "confirmNode", confirmData));
|
||||
|
||||
JSONObject endData = new JSONObject();
|
||||
endData.put("outputDefs", new JSONArray(List.of(
|
||||
parameterJson(
|
||||
"result", outputType, "confirm.templateType"))));
|
||||
nodes.add(nodeJson("end", "endNode", endData));
|
||||
|
||||
JSONArray edges = new JSONArray();
|
||||
edges.add(edgeJson("start-to-confirm", "start", "confirm"));
|
||||
edges.add(edgeJson("confirm-to-end", "confirm", "end"));
|
||||
JSONObject flow = new JSONObject();
|
||||
flow.put("nodes", nodes);
|
||||
flow.put("edges", edges);
|
||||
|
||||
ChainDefinition definition = ChainParser.builder()
|
||||
.withDefaultParsers(true)
|
||||
.build()
|
||||
.parse(flow.toJSONString());
|
||||
definition.setId("confirm-output-flow-test");
|
||||
return definition;
|
||||
}
|
||||
|
||||
private JSONObject nodeJson(
|
||||
String id, String type, JSONObject data) {
|
||||
JSONObject node = new JSONObject();
|
||||
node.put("id", id);
|
||||
node.put("type", type);
|
||||
node.put("data", data);
|
||||
return node;
|
||||
}
|
||||
|
||||
private JSONObject edgeJson(
|
||||
String id, String source, String target) {
|
||||
JSONObject edge = new JSONObject();
|
||||
edge.put("id", id);
|
||||
edge.put("source", source);
|
||||
edge.put("target", target);
|
||||
return edge;
|
||||
}
|
||||
|
||||
private JSONObject parameterJson(
|
||||
String name, String dataType, String ref) {
|
||||
JSONObject parameter = new JSONObject();
|
||||
parameter.put("name", name);
|
||||
parameter.put("dataType", dataType);
|
||||
if (ref != null) {
|
||||
parameter.put("ref", ref);
|
||||
parameter.put("refType", RefType.REF.toString());
|
||||
}
|
||||
return parameter;
|
||||
}
|
||||
|
||||
private ChainState awaitStatus(
|
||||
InMemoryChainStateRepository repository,
|
||||
String instanceId,
|
||||
ChainStatus expected) throws InterruptedException {
|
||||
long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(3);
|
||||
ChainState state;
|
||||
do {
|
||||
state = repository.load(instanceId);
|
||||
if (state != null && state.getStatus() == expected) {
|
||||
return state;
|
||||
}
|
||||
Thread.sleep(10L);
|
||||
} while (System.nanoTime() < deadline);
|
||||
Assert.fail("workflow did not reach status " + expected);
|
||||
return state;
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建用于取消传播验证的工作流。
|
||||
*
|
||||
|
||||
@@ -2,15 +2,29 @@ package com.easyagents.flow.core.test;
|
||||
|
||||
import com.easyagents.flow.core.chain.Chain;
|
||||
import com.easyagents.flow.core.chain.ChainDefinition;
|
||||
import com.easyagents.flow.core.chain.ChainResumeException;
|
||||
import com.easyagents.flow.core.chain.ChainState;
|
||||
import com.easyagents.flow.core.chain.ChainStatus;
|
||||
import com.easyagents.flow.core.chain.EventManager;
|
||||
import com.easyagents.flow.core.chain.Parameter;
|
||||
import com.easyagents.flow.core.chain.ParameterOption;
|
||||
import com.easyagents.flow.core.chain.repository.ChainStateField;
|
||||
import com.easyagents.flow.core.chain.repository.InMemoryChainStateRepository;
|
||||
import com.easyagents.flow.core.chain.repository.InMemoryNodeStateRepository;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
|
||||
import java.util.Arrays;
|
||||
import java.util.Collections;
|
||||
import java.util.EnumSet;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.concurrent.CountDownLatch;
|
||||
import java.util.concurrent.ExecutorService;
|
||||
import java.util.concurrent.Executors;
|
||||
import java.util.concurrent.Future;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
|
||||
/**
|
||||
* {@link Chain} 暂停恢复状态守卫测试。
|
||||
@@ -72,6 +86,167 @@ public class ChainResumeGuardTest {
|
||||
.containsKey("unexpected"));
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证确认选项只能按挂起时声明的字段和值恢复。
|
||||
*/
|
||||
@Test
|
||||
public void shouldValidateDeclaredResumeOptionsBeforeStateTransition() {
|
||||
InMemoryChainStateRepository stateRepository =
|
||||
new InMemoryChainStateRepository();
|
||||
Chain chain = createChain(stateRepository, "resume-options");
|
||||
Parameter single = optionParameter(
|
||||
"templateType__confirm", "会议类型", "radio");
|
||||
Parameter multiple = optionParameter(
|
||||
"participants__confirm", "参会人员", "checkbox");
|
||||
chain.getExecutionState().setSuspendForParameters(
|
||||
Arrays.asList(single, multiple));
|
||||
chain.suspend();
|
||||
|
||||
ChainResumeException invalidOption = assertRejected(
|
||||
chain,
|
||||
Map.of(
|
||||
"templateType__confirm", "UNKNOWN",
|
||||
"participants__confirm", List.of("REVIEW")));
|
||||
ChainResumeException extraField = assertRejected(
|
||||
chain,
|
||||
Map.of(
|
||||
"templateType__confirm", "AGENDA",
|
||||
"participants__confirm", List.of("REVIEW", "REVIEW")));
|
||||
assertRejected(
|
||||
chain,
|
||||
Map.of(
|
||||
"templateType__confirm", "AGENDA",
|
||||
"participants__confirm", List.of("REVIEW"),
|
||||
"extra", "value"));
|
||||
Assert.assertFalse(invalidOption.getMessage().contains("UNKNOWN"));
|
||||
Assert.assertFalse(extraField.getMessage().contains("extra"));
|
||||
|
||||
Assert.assertEquals(
|
||||
ChainStatus.SUSPEND,
|
||||
stateRepository.load("resume-options").getStatus());
|
||||
Assert.assertTrue(
|
||||
stateRepository.load("resume-options").getMemory().isEmpty());
|
||||
|
||||
boolean resumed = chain.resumeIfSuspended(Map.of(
|
||||
"templateType__confirm", "AGENDA",
|
||||
"participants__confirm", List.of("REVIEW", "BRIEFING")));
|
||||
|
||||
Assert.assertTrue(resumed);
|
||||
Assert.assertEquals(
|
||||
"AGENDA",
|
||||
stateRepository.load("resume-options")
|
||||
.getMemory()
|
||||
.get("templateType__confirm"));
|
||||
Assert.assertEquals(
|
||||
List.of("REVIEW", "BRIEFING"),
|
||||
stateRepository.load("resume-options")
|
||||
.getMemory()
|
||||
.get("participants__confirm"));
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证并发恢复同一暂停实例时,只有一个请求可以完成状态转换。
|
||||
*/
|
||||
@Test
|
||||
public void shouldAllowOnlyOneConcurrentResume() throws Exception {
|
||||
InMemoryChainStateRepository stateRepository =
|
||||
new InMemoryChainStateRepository();
|
||||
Chain first = createChain(stateRepository, "resume-concurrent");
|
||||
Chain second = createChain(stateRepository, "resume-concurrent");
|
||||
Parameter parameter = optionParameter(
|
||||
"templateType__confirm", "会议类型", "radio");
|
||||
first.getExecutionState().setSuspendForParameters(
|
||||
List.of(parameter));
|
||||
first.suspend();
|
||||
|
||||
CountDownLatch ready = new CountDownLatch(2);
|
||||
CountDownLatch start = new CountDownLatch(1);
|
||||
ExecutorService executor = Executors.newFixedThreadPool(2);
|
||||
try {
|
||||
Future<Boolean> firstResult = executor.submit(() -> {
|
||||
ready.countDown();
|
||||
start.await();
|
||||
return first.resumeIfSuspended(
|
||||
Map.of("templateType__confirm", "AGENDA"));
|
||||
});
|
||||
Future<Boolean> secondResult = executor.submit(() -> {
|
||||
ready.countDown();
|
||||
start.await();
|
||||
return second.resumeIfSuspended(
|
||||
Map.of("templateType__confirm", "REVIEW"));
|
||||
});
|
||||
|
||||
Assert.assertTrue(ready.await(5, TimeUnit.SECONDS));
|
||||
start.countDown();
|
||||
int resumedCount = (firstResult.get(5, TimeUnit.SECONDS) ? 1 : 0)
|
||||
+ (secondResult.get(5, TimeUnit.SECONDS) ? 1 : 0);
|
||||
|
||||
Assert.assertEquals(1, resumedCount);
|
||||
Assert.assertEquals(
|
||||
ChainStatus.RUNNING,
|
||||
stateRepository.load("resume-concurrent").getStatus());
|
||||
Object selected = stateRepository.load("resume-concurrent")
|
||||
.getMemory()
|
||||
.get("templateType__confirm");
|
||||
Assert.assertTrue(
|
||||
"AGENDA".equals(selected) || "REVIEW".equals(selected));
|
||||
} finally {
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证恢复变量、状态和暂停上下文通过一次原子更新完成。
|
||||
*/
|
||||
@Test
|
||||
public void shouldCommitResumeTransitionInSingleStateUpdate() {
|
||||
CountingChainStateRepository stateRepository =
|
||||
new CountingChainStateRepository();
|
||||
Chain chain = createChain(stateRepository, "resume-single-update");
|
||||
chain.getExecutionState().setSuspendForParameters(List.of(
|
||||
optionParameter("templateType__confirm", "会议类型", "radio")));
|
||||
chain.suspend();
|
||||
stateRepository.resetUpdateCount();
|
||||
|
||||
boolean resumed = chain.resumeIfSuspended(
|
||||
Map.of("templateType__confirm", "AGENDA"));
|
||||
|
||||
ChainState state = stateRepository.load("resume-single-update");
|
||||
Assert.assertTrue(resumed);
|
||||
Assert.assertEquals(1, stateRepository.getUpdateCount());
|
||||
Assert.assertEquals(ChainStatus.RUNNING, state.getStatus());
|
||||
Assert.assertNull(state.getSuspendNodeIds());
|
||||
Assert.assertNull(state.getSuspendForParameters());
|
||||
Assert.assertEquals("AGENDA", state.getMemory().get(
|
||||
"templateType__confirm"));
|
||||
}
|
||||
|
||||
private ChainResumeException assertRejected(
|
||||
Chain chain, Map<String, Object> variables) {
|
||||
try {
|
||||
chain.resumeIfSuspended(variables);
|
||||
Assert.fail("invalid resume variables must be rejected");
|
||||
return null;
|
||||
} catch (ChainResumeException expected) {
|
||||
Assert.assertNotNull(expected.getMessage());
|
||||
return expected;
|
||||
}
|
||||
}
|
||||
|
||||
private Parameter optionParameter(
|
||||
String name, String label, String formType) {
|
||||
Parameter parameter = new Parameter();
|
||||
parameter.setName(name);
|
||||
parameter.setFormLabel(label);
|
||||
parameter.setFormType(formType);
|
||||
parameter.setRequired(true);
|
||||
parameter.setOptions(List.of(
|
||||
new ParameterOption("第一议题", "AGENDA"),
|
||||
new ParameterOption("审议类", "REVIEW"),
|
||||
new ParameterOption("听取类", "BRIEFING")));
|
||||
return parameter;
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建使用进程内状态仓储的最小工作流。
|
||||
*
|
||||
@@ -93,4 +268,25 @@ public class ChainResumeGuardTest {
|
||||
chain.setEventManager(new EventManager());
|
||||
return chain;
|
||||
}
|
||||
|
||||
private static class CountingChainStateRepository
|
||||
extends InMemoryChainStateRepository {
|
||||
private final AtomicInteger updateCount = new AtomicInteger();
|
||||
|
||||
@Override
|
||||
public boolean tryUpdate(
|
||||
ChainState chainState,
|
||||
EnumSet<ChainStateField> fields) {
|
||||
updateCount.incrementAndGet();
|
||||
return super.tryUpdate(chainState, fields);
|
||||
}
|
||||
|
||||
private int getUpdateCount() {
|
||||
return updateCount.get();
|
||||
}
|
||||
|
||||
private void resetUpdateCount() {
|
||||
updateCount.set(0);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,235 @@
|
||||
/**
|
||||
* Copyright (c) 2025-2026, Michael Yang 杨福海 (fuhai999@gmail.com).
|
||||
* <p>
|
||||
* Licensed under the GNU Lesser General Public License (LGPL) ,Version 3.0.
|
||||
*/
|
||||
package com.easyagents.flow.core.test;
|
||||
|
||||
import com.alibaba.fastjson.JSONArray;
|
||||
import com.alibaba.fastjson.JSONObject;
|
||||
import com.easyagents.flow.core.chain.Chain;
|
||||
import com.easyagents.flow.core.chain.ChainDefinition;
|
||||
import com.easyagents.flow.core.chain.ChainSuspendException;
|
||||
import com.easyagents.flow.core.chain.DataType;
|
||||
import com.easyagents.flow.core.chain.EventManager;
|
||||
import com.easyagents.flow.core.chain.Parameter;
|
||||
import com.easyagents.flow.core.chain.repository.InMemoryChainStateRepository;
|
||||
import com.easyagents.flow.core.chain.repository.InMemoryNodeStateRepository;
|
||||
import com.easyagents.flow.core.chain.runtime.Trigger;
|
||||
import com.easyagents.flow.core.chain.runtime.TriggerContext;
|
||||
import com.easyagents.flow.core.chain.runtime.TriggerType;
|
||||
import com.easyagents.flow.core.node.ConfirmNode;
|
||||
import com.easyagents.flow.core.parser.impl.ConfirmNodeParser;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
|
||||
import java.util.Collections;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
/**
|
||||
* 用户确认节点选择与输出契约测试。
|
||||
*/
|
||||
public class ConfirmNodeTest {
|
||||
|
||||
@Test
|
||||
public void shouldBuildSingleChoiceAndReturnSelectedContent() {
|
||||
ConfirmNode node = parse(false);
|
||||
node.setId("confirm-1");
|
||||
Chain chain = createChain();
|
||||
|
||||
Parameter parameter = suspend(node, chain);
|
||||
Assert.assertEquals("selection__confirm-1", parameter.getName());
|
||||
Assert.assertEquals("radio", parameter.getFormType());
|
||||
Assert.assertEquals("选择内容", parameter.getFormLabel());
|
||||
Assert.assertEquals(DataType.String, parameter.getDataType());
|
||||
Assert.assertEquals("第一议题", parameter.getOptions().get(0).getLabel());
|
||||
Assert.assertEquals("第一议题", parameter.getOptions().get(0).getValue());
|
||||
|
||||
chain.getExecutionState().getMemory().put(parameter.getName(), "审议类");
|
||||
Map<String, Object> result = executeAsResume(node, chain);
|
||||
|
||||
Assert.assertEquals(Collections.singletonMap("selection", "审议类"), result);
|
||||
Assert.assertFalse(chain.getExecutionState().getMemory()
|
||||
.containsKey(parameter.getName()));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldBuildMultipleChoiceAndReturnSelectedContents() {
|
||||
ConfirmNode node = parse(true);
|
||||
node.setId("confirm-1");
|
||||
Chain chain = createChain();
|
||||
|
||||
Parameter parameter = suspend(node, chain);
|
||||
Assert.assertEquals("checkbox", parameter.getFormType());
|
||||
Assert.assertEquals(DataType.Array_String, parameter.getDataType());
|
||||
|
||||
List<String> selected = List.of("第一议题", "听取类");
|
||||
chain.getExecutionState().getMemory().put(parameter.getName(), selected);
|
||||
|
||||
Assert.assertEquals(selected, executeAsResume(node, chain).get("selection"));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldUseConfiguredOutputNameWithoutChangingSuspendParameter() {
|
||||
ConfirmNode node = parse(false, "templateChoice");
|
||||
node.setId("confirm-1");
|
||||
Chain chain = createChain();
|
||||
|
||||
Parameter parameter = suspend(node, chain);
|
||||
Assert.assertEquals("selection__confirm-1", parameter.getName());
|
||||
|
||||
chain.getExecutionState().getMemory().put(
|
||||
parameter.getName(), "第一议题");
|
||||
Assert.assertEquals(
|
||||
Collections.singletonMap("templateChoice", "第一议题"),
|
||||
executeAsResume(node, chain));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldIgnorePrefilledValueWithoutResumeTrigger() {
|
||||
ConfirmNode node = parse(false);
|
||||
node.setId("confirm-1");
|
||||
Chain chain = createChain();
|
||||
chain.getExecutionState().getMemory().put(
|
||||
"selection__confirm-1", "未配置值");
|
||||
|
||||
Parameter parameter = suspend(node, chain);
|
||||
|
||||
Assert.assertEquals("selection__confirm-1", parameter.getName());
|
||||
Assert.assertFalse(chain.getExecutionState().getMemory()
|
||||
.containsKey(parameter.getName()));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRejectDuplicateNormalizedOptionContents() {
|
||||
ConfirmNode node = parse(false);
|
||||
node.setOptions(List.of("审议类", " 审议类 "));
|
||||
|
||||
try {
|
||||
node.validateConfiguration();
|
||||
Assert.fail("duplicate option contents must be rejected");
|
||||
} catch (IllegalArgumentException expected) {
|
||||
Assert.assertTrue(expected.getMessage().contains("选项内容重复"));
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRejectNonStringOptionDuringParsing() {
|
||||
JSONObject data = data(false);
|
||||
data.getJSONArray("options").add(1);
|
||||
|
||||
try {
|
||||
new ConfirmNodeParser().doParse(
|
||||
new JSONObject(), data, new JSONObject());
|
||||
Assert.fail("non-string option must be rejected");
|
||||
} catch (IllegalArgumentException expected) {
|
||||
Assert.assertTrue(expected.getMessage().contains("必须为字符串"));
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRejectStringEncodedOptionsDuringParsing() {
|
||||
JSONObject data = data(false);
|
||||
data.put("options", "[\"第一议题\"]");
|
||||
|
||||
try {
|
||||
new ConfirmNodeParser().doParse(
|
||||
new JSONObject(), data, new JSONObject());
|
||||
Assert.fail("string encoded options must be rejected");
|
||||
} catch (IllegalArgumentException expected) {
|
||||
Assert.assertTrue(expected.getMessage().contains("必须为数组"));
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRejectImplicitModeConversionDuringParsing() {
|
||||
JSONObject data = data(false);
|
||||
data.put("multiple", "false");
|
||||
|
||||
try {
|
||||
new ConfirmNodeParser().doParse(
|
||||
new JSONObject(), data, new JSONObject());
|
||||
Assert.fail("non-boolean mode must be rejected");
|
||||
} catch (IllegalArgumentException expected) {
|
||||
Assert.assertTrue(expected.getMessage().contains("必须为布尔值"));
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRejectUnknownConfigurationFieldDuringParsing() {
|
||||
JSONObject data = data(false);
|
||||
data.put("async", true);
|
||||
|
||||
try {
|
||||
new ConfirmNodeParser().doParse(
|
||||
new JSONObject(), data, new JSONObject());
|
||||
Assert.fail("unknown confirm configuration must be rejected");
|
||||
} catch (IllegalArgumentException expected) {
|
||||
Assert.assertTrue(expected.getMessage().contains("无效配置字段"));
|
||||
}
|
||||
}
|
||||
|
||||
private static ConfirmNode parse(boolean multiple) {
|
||||
return parse(multiple, ConfirmNode.DEFAULT_OUTPUT_NAME);
|
||||
}
|
||||
|
||||
private static ConfirmNode parse(boolean multiple, String outputName) {
|
||||
ConfirmNode node = new ConfirmNodeParser().doParse(
|
||||
new JSONObject(), data(multiple), new JSONObject());
|
||||
Parameter output = new Parameter();
|
||||
output.setName(outputName);
|
||||
output.setDataType(multiple
|
||||
? DataType.Array_String
|
||||
: DataType.String);
|
||||
node.setOutputDefs(Collections.singletonList(output));
|
||||
node.validateConfiguration();
|
||||
return node;
|
||||
}
|
||||
|
||||
private static JSONObject data(boolean multiple) {
|
||||
JSONObject data = new JSONObject();
|
||||
data.put("message", "请选择会议纪要模板");
|
||||
data.put("multiple", multiple);
|
||||
JSONArray options = new JSONArray();
|
||||
options.addAll(List.of("第一议题", "审议类", "听取类"));
|
||||
data.put("options", options);
|
||||
return data;
|
||||
}
|
||||
|
||||
private static Parameter suspend(ConfirmNode node, Chain chain) {
|
||||
try {
|
||||
node.execute(chain);
|
||||
throw new AssertionError("confirm node must suspend");
|
||||
} catch (ChainSuspendException expected) {
|
||||
Assert.assertEquals(1, expected.getSuspendParameters().size());
|
||||
return expected.getSuspendParameters().get(0);
|
||||
}
|
||||
}
|
||||
|
||||
private static Map<String, Object> executeAsResume(
|
||||
ConfirmNode node, Chain chain) {
|
||||
Trigger trigger = new Trigger();
|
||||
trigger.setType(TriggerType.RESUME);
|
||||
trigger.setStateInstanceId(chain.getStateInstanceId());
|
||||
trigger.setNodeId(node.getId());
|
||||
TriggerContext.setCurrentTrigger(trigger);
|
||||
try {
|
||||
return node.execute(chain);
|
||||
} finally {
|
||||
TriggerContext.clearCurrentTrigger();
|
||||
}
|
||||
}
|
||||
|
||||
private static Chain createChain() {
|
||||
ChainDefinition definition = new ChainDefinition();
|
||||
definition.setId("confirm-node-test");
|
||||
definition.setNodes(Collections.emptyList());
|
||||
definition.setEdges(Collections.emptyList());
|
||||
Chain chain = new Chain(definition, "confirm-node-instance");
|
||||
chain.setChainStateRepository(new InMemoryChainStateRepository());
|
||||
chain.setNodeStateRepository(new InMemoryNodeStateRepository());
|
||||
chain.setEventManager(new EventManager());
|
||||
return chain;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,217 @@
|
||||
/**
|
||||
* Copyright (c) 2025-2026, Michael Yang 杨福海 (fuhai999@gmail.com).
|
||||
* <p>
|
||||
* Licensed under the GNU Lesser General Public License (LGPL) ,Version 3.0.
|
||||
*/
|
||||
package com.easyagents.flow.core.test;
|
||||
|
||||
import com.alibaba.fastjson.JSONArray;
|
||||
import com.alibaba.fastjson.JSONObject;
|
||||
import com.easyagents.flow.core.chain.Chain;
|
||||
import com.easyagents.flow.core.chain.ChainDefinition;
|
||||
import com.easyagents.flow.core.chain.ChainState;
|
||||
import com.easyagents.flow.core.knowledge.Knowledge;
|
||||
import com.easyagents.flow.core.knowledge.KnowledgeManager;
|
||||
import com.easyagents.flow.core.knowledge.KnowledgeProvider;
|
||||
import com.easyagents.flow.core.knowledge.KnowledgeSearchRequest;
|
||||
import com.easyagents.flow.core.node.KnowledgeNode;
|
||||
import com.easyagents.flow.core.parser.impl.KnowledgeNodeParser;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
|
||||
/**
|
||||
* 知识库节点多来源契约测试。
|
||||
*/
|
||||
public class KnowledgeNodeTest {
|
||||
|
||||
@Test
|
||||
public void shouldParseLegacyKnowledgeId() {
|
||||
JSONObject data = baseData();
|
||||
data.put("knowledgeId", "101");
|
||||
|
||||
KnowledgeNode node = parse(data);
|
||||
|
||||
Assert.assertEquals(List.of("101"), node.getKnowledgeIds());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldPreferKnowledgeIdsAndKeepOrder() {
|
||||
JSONObject data = baseData();
|
||||
data.put("knowledgeId", "legacy");
|
||||
JSONArray ids = new JSONArray();
|
||||
ids.addAll(List.of("201", "202"));
|
||||
data.put("knowledgeIds", ids);
|
||||
|
||||
KnowledgeNode node = parse(data);
|
||||
|
||||
Assert.assertEquals(List.of("201", "202"), node.getKnowledgeIds());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldNormalizeKnowledgeIdsBeforeStoringThem() {
|
||||
JSONObject data = baseData();
|
||||
JSONArray ids = new JSONArray();
|
||||
ids.addAll(List.of(" 201 ", "202"));
|
||||
data.put("knowledgeIds", ids);
|
||||
|
||||
KnowledgeNode node = parse(data);
|
||||
|
||||
Assert.assertEquals(List.of("201", "202"), node.getKnowledgeIds());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRejectInvalidKnowledgeIds() {
|
||||
JSONObject data = baseData();
|
||||
data.put("knowledgeIds", "[201,202]");
|
||||
assertParseFailure(data, "必须为数组");
|
||||
|
||||
JSONArray duplicateIds = new JSONArray();
|
||||
duplicateIds.addAll(List.of("201", "201"));
|
||||
data.put("knowledgeIds", duplicateIds);
|
||||
assertParseFailure(data, "重复值");
|
||||
}
|
||||
|
||||
@Test
|
||||
public void defaultProviderShouldKeepSingleKnowledgeCompatibility() {
|
||||
KnowledgeProvider provider = id ->
|
||||
(keyword, limit, node, chain) -> List.of(Map.of(
|
||||
"knowledgeId", id,
|
||||
"content", keyword));
|
||||
KnowledgeNode node = new KnowledgeNode();
|
||||
node.setKnowledgeId("301");
|
||||
|
||||
Map<String, Object> output = provider.search(
|
||||
new KnowledgeSearchRequest(
|
||||
node.getKnowledgeIds(),
|
||||
"问题",
|
||||
3,
|
||||
"HYBRID",
|
||||
node,
|
||||
null));
|
||||
|
||||
Assert.assertNotNull(output);
|
||||
Assert.assertEquals(1, ((List<?>) output.get("documents")).size());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void defaultProviderShouldDeclineMultiKnowledgeRequest() {
|
||||
KnowledgeProvider provider = id -> null;
|
||||
KnowledgeNode node = new KnowledgeNode();
|
||||
node.setKnowledgeIds(List.of("401", "402"));
|
||||
|
||||
Assert.assertNull(provider.search(new KnowledgeSearchRequest(
|
||||
node.getKnowledgeIds(),
|
||||
"问题",
|
||||
3,
|
||||
"VECTOR",
|
||||
node,
|
||||
null)));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldResolveVariableLimitAndDefaultBlankValueAtRuntime() {
|
||||
Assert.assertEquals(7, executeAndCaptureLimit("{{start.limit}}", "7"));
|
||||
Assert.assertEquals(10, executeAndCaptureLimit("{{start.limit}}", " "));
|
||||
Assert.assertEquals(10, executeAndCaptureLimit("{{start.limit ?? }}", null));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRejectInvalidResolvedVariableLimitAtRuntime() {
|
||||
assertRuntimeLimitFailure("abc");
|
||||
assertRuntimeLimitFailure("0");
|
||||
assertRuntimeLimitFailure("-2");
|
||||
}
|
||||
|
||||
private static KnowledgeNode parse(JSONObject data) {
|
||||
return new KnowledgeNodeParser().doParse(
|
||||
new JSONObject(), data, new JSONObject());
|
||||
}
|
||||
|
||||
private static JSONObject baseData() {
|
||||
JSONObject data = new JSONObject();
|
||||
data.put("keyword", "问题");
|
||||
data.put("limit", "5");
|
||||
data.put("retrievalMode", "VECTOR");
|
||||
return data;
|
||||
}
|
||||
|
||||
private static void assertParseFailure(
|
||||
JSONObject data, String expectedMessage) {
|
||||
try {
|
||||
parse(data);
|
||||
Assert.fail("invalid knowledgeIds must be rejected");
|
||||
} catch (IllegalArgumentException expected) {
|
||||
Assert.assertTrue(expected.getMessage().contains(expectedMessage));
|
||||
}
|
||||
}
|
||||
|
||||
private static int executeAndCaptureLimit(
|
||||
String limitTemplate,
|
||||
String runtimeValue) {
|
||||
AtomicInteger capturedLimit = new AtomicInteger(-1);
|
||||
KnowledgeProvider provider = new KnowledgeProvider() {
|
||||
@Override
|
||||
public Knowledge getKnowledge(Object id) {
|
||||
return null;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Map<String, Object> search(KnowledgeSearchRequest request) {
|
||||
capturedLimit.set(request.getLimit());
|
||||
return Map.of("documents", List.of());
|
||||
}
|
||||
};
|
||||
KnowledgeManager.getInstance().registerProvider(provider);
|
||||
try {
|
||||
KnowledgeNode node = runtimeNode(limitTemplate);
|
||||
ChainState state = new ChainState();
|
||||
if (runtimeValue != null) {
|
||||
state.getMemory().put("start.limit", runtimeValue);
|
||||
}
|
||||
node.execute(new FixedStateChain(state));
|
||||
return capturedLimit.get();
|
||||
} finally {
|
||||
KnowledgeManager.getInstance().removeProvider(provider);
|
||||
}
|
||||
}
|
||||
|
||||
private static void assertRuntimeLimitFailure(String runtimeValue) {
|
||||
KnowledgeNode node = runtimeNode("{{start.limit}}");
|
||||
ChainState state = new ChainState();
|
||||
state.getMemory().put("start.limit", runtimeValue);
|
||||
|
||||
IllegalArgumentException exception = Assert.assertThrows(
|
||||
IllegalArgumentException.class,
|
||||
() -> node.execute(new FixedStateChain(state)));
|
||||
|
||||
Assert.assertTrue(exception.getMessage().contains("必须为正整数"));
|
||||
}
|
||||
|
||||
private static KnowledgeNode runtimeNode(String limitTemplate) {
|
||||
KnowledgeNode node = new KnowledgeNode();
|
||||
node.setKnowledgeIds(List.of("501", "502"));
|
||||
node.setKeyword("问题");
|
||||
node.setLimit(limitTemplate);
|
||||
node.setRetrievalMode("VECTOR");
|
||||
return node;
|
||||
}
|
||||
|
||||
private static final class FixedStateChain extends Chain {
|
||||
|
||||
private final ChainState state;
|
||||
|
||||
private FixedStateChain(ChainState state) {
|
||||
super(new ChainDefinition(), "knowledge-node-limit-test");
|
||||
this.state = state;
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChainState getExecutionState() {
|
||||
return state;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,407 @@
|
||||
package com.easyagents.flow.core.test;
|
||||
|
||||
import com.alibaba.fastjson.JSONObject;
|
||||
import com.easyagents.flow.core.chain.Chain;
|
||||
import com.easyagents.flow.core.chain.ChainDefinition;
|
||||
import com.easyagents.flow.core.chain.ChainState;
|
||||
import com.easyagents.flow.core.chain.ChainStatus;
|
||||
import com.easyagents.flow.core.chain.Edge;
|
||||
import com.easyagents.flow.core.chain.Node;
|
||||
import com.easyagents.flow.core.chain.NodeCondition;
|
||||
import com.easyagents.flow.core.chain.NodeJoinMode;
|
||||
import com.easyagents.flow.core.chain.NodeState;
|
||||
import com.easyagents.flow.core.chain.Parameter;
|
||||
import com.easyagents.flow.core.chain.RefType;
|
||||
import com.easyagents.flow.core.chain.repository.InMemoryChainStateRepository;
|
||||
import com.easyagents.flow.core.chain.repository.InMemoryNodeStateRepository;
|
||||
import com.easyagents.flow.core.chain.runtime.ChainExecutor;
|
||||
import com.easyagents.flow.core.chain.runtime.InMemoryTriggerStore;
|
||||
import com.easyagents.flow.core.chain.runtime.TriggerScheduler;
|
||||
import com.easyagents.flow.core.node.BaseNode;
|
||||
import com.easyagents.flow.core.node.EndNode;
|
||||
import com.easyagents.flow.core.node.StartNode;
|
||||
import com.easyagents.flow.core.parser.ChainParser;
|
||||
import com.easyagents.flow.core.parser.impl.EndNodeParser;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
|
||||
import java.util.Collections;
|
||||
import java.util.HashMap;
|
||||
import java.util.Map;
|
||||
import java.util.concurrent.CountDownLatch;
|
||||
import java.util.concurrent.ExecutorService;
|
||||
import java.util.concurrent.Executors;
|
||||
import java.util.concurrent.ScheduledExecutorService;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
import java.util.function.BooleanSupplier;
|
||||
|
||||
/**
|
||||
* 验证普通节点的直接入边汇聚模式。
|
||||
*/
|
||||
public class NodeJoinModeTest {
|
||||
|
||||
@Test
|
||||
public void shouldParseJoinModeAndRejectInvalidValues() {
|
||||
ChainParser parser = ChainParser.builder()
|
||||
.withDefaultParsers(true)
|
||||
.build();
|
||||
|
||||
Assert.assertEquals(
|
||||
NodeJoinMode.ANY,
|
||||
parseEndNode(parser, null, null).getJoinMode());
|
||||
Assert.assertEquals(
|
||||
NodeJoinMode.ALL,
|
||||
parseEndNode(parser, "all", null).getJoinMode());
|
||||
Assert.assertEquals(
|
||||
NodeJoinMode.ANY,
|
||||
parseEndNode(parser, "ANY", null).getJoinMode());
|
||||
|
||||
assertInvalidJoinMode(() -> parseEndNode(parser, "first", null));
|
||||
assertInvalidJoinMode(() -> parseEndNode(parser, "", null));
|
||||
assertInvalidJoinMode(() -> parseEndNode(parser, "all", "loop"));
|
||||
|
||||
ProbeJoinNode node = new ProbeJoinNode(new AtomicInteger());
|
||||
node.setJoinMode(NodeJoinMode.ALL);
|
||||
assertInvalidJoinMode(() -> node.setParentId("loop"));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldWaitForEveryInboundEdgeBeforeExecutingAndCheckingCondition()
|
||||
throws Exception {
|
||||
JoinFixture fixture = createJoinFixture(NodeJoinMode.ALL, true);
|
||||
String instanceId = null;
|
||||
try {
|
||||
instanceId = fixture.executor.executeAsync(
|
||||
fixture.definition.getId(), Collections.emptyMap());
|
||||
Assert.assertTrue(fixture.branchACompleted.await(2, TimeUnit.SECONDS));
|
||||
Assert.assertTrue(fixture.branchBStarted.await(2, TimeUnit.SECONDS));
|
||||
|
||||
String currentInstanceId = instanceId;
|
||||
await(() -> hasTriggerEdge(
|
||||
fixture.nodeStateRepository,
|
||||
currentInstanceId,
|
||||
"join",
|
||||
"a-join"));
|
||||
Assert.assertEquals(0, fixture.joinExecutions.get());
|
||||
Assert.assertEquals(0, fixture.conditionChecks.get());
|
||||
|
||||
fixture.releaseBranchB.countDown();
|
||||
ChainState finalState = awaitTerminal(
|
||||
fixture.chainStateRepository, instanceId);
|
||||
|
||||
Assert.assertEquals(1, fixture.joinExecutions.get());
|
||||
Assert.assertEquals(1, fixture.conditionChecks.get());
|
||||
Assert.assertEquals("A+B", finalState.getExecuteResult().get("combined"));
|
||||
Assert.assertEquals(Boolean.TRUE, finalState.getExecuteResult().get("sawBoth"));
|
||||
} finally {
|
||||
fixture.releaseBranchB.countDown();
|
||||
fixture.scheduler.shutdown();
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldKeepAnyModeFirstArrivalBehavior() throws Exception {
|
||||
JoinFixture fixture = createJoinFixture(NodeJoinMode.ANY, false);
|
||||
try {
|
||||
String instanceId = fixture.executor.executeAsync(
|
||||
fixture.definition.getId(), Collections.emptyMap());
|
||||
Assert.assertTrue(fixture.branchACompleted.await(2, TimeUnit.SECONDS));
|
||||
Assert.assertTrue(fixture.branchBStarted.await(2, TimeUnit.SECONDS));
|
||||
|
||||
await(() -> fixture.joinExecutions.get() > 0);
|
||||
ChainState finalState = awaitTerminal(
|
||||
fixture.chainStateRepository, instanceId);
|
||||
|
||||
Assert.assertEquals(1, fixture.joinExecutions.get());
|
||||
Assert.assertEquals(Boolean.FALSE, finalState.getExecuteResult().get("sawBoth"));
|
||||
} finally {
|
||||
fixture.releaseBranchB.countDown();
|
||||
fixture.scheduler.shutdown();
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldKeepDefaultAnyRetryAndLoopBehavior() throws Exception {
|
||||
ScheduledExecutorService schedulerPool =
|
||||
Executors.newSingleThreadScheduledExecutor();
|
||||
ExecutorService workerPool = Executors.newFixedThreadPool(3);
|
||||
TriggerScheduler scheduler = new TriggerScheduler(
|
||||
new InMemoryTriggerStore(), schedulerPool, workerPool, 1000L);
|
||||
AtomicInteger executions = new AtomicInteger();
|
||||
ChainDefinition definition = new ChainDefinition();
|
||||
definition.setId("join-mode-retry-loop");
|
||||
StartNode start = new StartNode();
|
||||
start.setId("start");
|
||||
RetryLoopNode worker = new RetryLoopNode(executions);
|
||||
worker.setId("worker");
|
||||
worker.setRetryEnable(true);
|
||||
worker.setMaxRetryCount(1);
|
||||
worker.setRetryIntervalMs(0L);
|
||||
worker.setLoopEnable(true);
|
||||
worker.setMaxLoopCount(2);
|
||||
worker.setLoopIntervalMs(0L);
|
||||
EndNode end = endNode("worker", "count", "count");
|
||||
definition.addNode(start);
|
||||
definition.addNode(worker);
|
||||
definition.addNode(end);
|
||||
definition.addEdge(edge("start-worker", "start", "worker"));
|
||||
definition.addEdge(edge("worker-end", "worker", "end"));
|
||||
|
||||
ChainExecutor executor = new ChainExecutor(
|
||||
ignored -> definition,
|
||||
new InMemoryChainStateRepository(),
|
||||
new InMemoryNodeStateRepository(),
|
||||
scheduler);
|
||||
try {
|
||||
Map<String, Object> result = executor.execute(
|
||||
definition.getId(), Collections.emptyMap(), 5L, TimeUnit.SECONDS);
|
||||
|
||||
Assert.assertEquals(NodeJoinMode.ANY, worker.getJoinMode());
|
||||
Assert.assertEquals(3, executions.get());
|
||||
Assert.assertEquals(3, result.get("count"));
|
||||
} finally {
|
||||
scheduler.shutdown();
|
||||
}
|
||||
}
|
||||
|
||||
private JoinFixture createJoinFixture(
|
||||
NodeJoinMode joinMode, boolean withCondition) {
|
||||
JoinFixture fixture = new JoinFixture();
|
||||
fixture.schedulerPool = Executors.newScheduledThreadPool(2);
|
||||
fixture.workerPool = Executors.newFixedThreadPool(4);
|
||||
fixture.scheduler = new TriggerScheduler(
|
||||
new InMemoryTriggerStore(),
|
||||
fixture.schedulerPool,
|
||||
fixture.workerPool,
|
||||
1000L);
|
||||
fixture.chainStateRepository = new InMemoryChainStateRepository();
|
||||
fixture.nodeStateRepository = new InMemoryNodeStateRepository();
|
||||
fixture.definition = new ChainDefinition();
|
||||
fixture.definition.setId("join-mode-" + joinMode.getValue());
|
||||
|
||||
StartNode start = new StartNode();
|
||||
start.setId("start");
|
||||
BranchNode branchA = new BranchNode(
|
||||
"A", fixture.branchACompleted, null, null);
|
||||
branchA.setId("a");
|
||||
BranchNode branchB = new BranchNode(
|
||||
"B", null, fixture.branchBStarted, fixture.releaseBranchB);
|
||||
branchB.setId("b");
|
||||
ProbeJoinNode join = new ProbeJoinNode(fixture.joinExecutions);
|
||||
join.setId("join");
|
||||
join.setJoinMode(joinMode);
|
||||
if (withCondition) {
|
||||
join.setCondition(new BothOutputsCondition(fixture.conditionChecks));
|
||||
}
|
||||
EndNode end = endNode("join", "combined", "combined");
|
||||
end.addOutputDef(outputRef("sawBoth", "join.sawBoth"));
|
||||
|
||||
fixture.definition.addNode(start);
|
||||
fixture.definition.addNode(branchA);
|
||||
fixture.definition.addNode(branchB);
|
||||
fixture.definition.addNode(join);
|
||||
fixture.definition.addNode(end);
|
||||
fixture.definition.addEdge(edge("start-a", "start", "a"));
|
||||
fixture.definition.addEdge(edge("start-b", "start", "b"));
|
||||
fixture.definition.addEdge(edge("a-join", "a", "join"));
|
||||
fixture.definition.addEdge(edge("b-join", "b", "join"));
|
||||
fixture.definition.addEdge(edge("join-end", "join", "end"));
|
||||
|
||||
fixture.executor = new ChainExecutor(
|
||||
ignored -> fixture.definition,
|
||||
fixture.chainStateRepository,
|
||||
fixture.nodeStateRepository,
|
||||
fixture.scheduler);
|
||||
return fixture;
|
||||
}
|
||||
|
||||
private Node parseEndNode(
|
||||
ChainParser parser, String joinMode, String parentId) {
|
||||
JSONObject data = new JSONObject();
|
||||
if (joinMode != null) {
|
||||
data.put("joinMode", joinMode);
|
||||
}
|
||||
JSONObject nodeJson = new JSONObject();
|
||||
nodeJson.put("id", "end");
|
||||
nodeJson.put("type", "endNode");
|
||||
nodeJson.put("parentId", parentId);
|
||||
nodeJson.put("data", data);
|
||||
return new EndNodeParser().parse(
|
||||
nodeJson, new JSONObject(), parser);
|
||||
}
|
||||
|
||||
private void assertInvalidJoinMode(Runnable action) {
|
||||
try {
|
||||
action.run();
|
||||
Assert.fail("Expected invalid join mode");
|
||||
} catch (IllegalArgumentException expected) {
|
||||
Assert.assertTrue(expected.getMessage().contains("joinMode"));
|
||||
}
|
||||
}
|
||||
|
||||
private static boolean hasTriggerEdge(
|
||||
InMemoryNodeStateRepository repository,
|
||||
String instanceId,
|
||||
String nodeId,
|
||||
String edgeId) {
|
||||
NodeState state = repository.load(instanceId, nodeId);
|
||||
return state != null && state.getTriggerEdgeIds().contains(edgeId);
|
||||
}
|
||||
|
||||
private static ChainState awaitTerminal(
|
||||
InMemoryChainStateRepository repository,
|
||||
String instanceId) throws Exception {
|
||||
await(() -> {
|
||||
ChainState state = repository.load(instanceId);
|
||||
return state != null
|
||||
&& state.getStatus() != null
|
||||
&& state.getStatus().isTerminal();
|
||||
});
|
||||
ChainState state = repository.load(instanceId);
|
||||
Assert.assertEquals(ChainStatus.SUCCEEDED, state.getStatus());
|
||||
return state;
|
||||
}
|
||||
|
||||
private static void await(BooleanSupplier condition) throws Exception {
|
||||
long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(3L);
|
||||
while (!condition.getAsBoolean() && System.nanoTime() < deadline) {
|
||||
Thread.sleep(10L);
|
||||
}
|
||||
Assert.assertTrue("condition was not met before timeout", condition.getAsBoolean());
|
||||
}
|
||||
|
||||
private static EndNode endNode(
|
||||
String sourceNodeId, String sourceName, String outputName) {
|
||||
EndNode end = new EndNode();
|
||||
end.setId("end");
|
||||
end.addOutputDef(outputRef(
|
||||
outputName, sourceNodeId + "." + sourceName));
|
||||
return end;
|
||||
}
|
||||
|
||||
private static Parameter outputRef(String name, String ref) {
|
||||
Parameter parameter = new Parameter();
|
||||
parameter.setName(name);
|
||||
parameter.setRef(ref);
|
||||
parameter.setRefType(RefType.REF);
|
||||
return parameter;
|
||||
}
|
||||
|
||||
private static Edge edge(String id, String source, String target) {
|
||||
Edge edge = new Edge();
|
||||
edge.setId(id);
|
||||
edge.setSource(source);
|
||||
edge.setTarget(target);
|
||||
return edge;
|
||||
}
|
||||
|
||||
private static final class BranchNode extends BaseNode {
|
||||
private final String value;
|
||||
private final CountDownLatch completed;
|
||||
private final CountDownLatch started;
|
||||
private final CountDownLatch release;
|
||||
|
||||
private BranchNode(
|
||||
String value,
|
||||
CountDownLatch completed,
|
||||
CountDownLatch started,
|
||||
CountDownLatch release) {
|
||||
this.value = value;
|
||||
this.completed = completed;
|
||||
this.started = started;
|
||||
this.release = release;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Map<String, Object> execute(Chain chain) {
|
||||
if (started != null) {
|
||||
started.countDown();
|
||||
}
|
||||
if (release != null) {
|
||||
try {
|
||||
if (!release.await(3L, TimeUnit.SECONDS)) {
|
||||
throw new IllegalStateException("branch release timed out");
|
||||
}
|
||||
} catch (InterruptedException error) {
|
||||
Thread.currentThread().interrupt();
|
||||
throw new IllegalStateException("branch interrupted", error);
|
||||
}
|
||||
}
|
||||
if (completed != null) {
|
||||
completed.countDown();
|
||||
}
|
||||
return Collections.singletonMap("value", value);
|
||||
}
|
||||
}
|
||||
|
||||
private static final class ProbeJoinNode extends BaseNode {
|
||||
private final AtomicInteger executions;
|
||||
|
||||
private ProbeJoinNode(AtomicInteger executions) {
|
||||
this.executions = executions;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Map<String, Object> execute(Chain chain) {
|
||||
executions.incrementAndGet();
|
||||
Object a = chain.getExecutionState().getMemory().get("a.value");
|
||||
Object b = chain.getExecutionState().getMemory().get("b.value");
|
||||
Map<String, Object> result = new HashMap<>();
|
||||
result.put("combined", String.valueOf(a) + "+" + String.valueOf(b));
|
||||
result.put("sawBoth", a != null && b != null);
|
||||
return result;
|
||||
}
|
||||
}
|
||||
|
||||
private static final class BothOutputsCondition implements NodeCondition {
|
||||
private static final long serialVersionUID = 1L;
|
||||
private final AtomicInteger checks;
|
||||
|
||||
private BothOutputsCondition(AtomicInteger checks) {
|
||||
this.checks = checks;
|
||||
}
|
||||
|
||||
@Override
|
||||
public boolean check(
|
||||
Chain chain,
|
||||
NodeState context,
|
||||
Map<String, Object> executeResult) {
|
||||
checks.incrementAndGet();
|
||||
Map<String, Object> memory = chain.getExecutionState().getMemory();
|
||||
return memory.containsKey("a.value") && memory.containsKey("b.value");
|
||||
}
|
||||
}
|
||||
|
||||
private static final class RetryLoopNode extends BaseNode {
|
||||
private final AtomicInteger executions;
|
||||
|
||||
private RetryLoopNode(AtomicInteger executions) {
|
||||
this.executions = executions;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Map<String, Object> execute(Chain chain) {
|
||||
int count = executions.incrementAndGet();
|
||||
if (count == 1) {
|
||||
throw new IllegalStateException("retry once");
|
||||
}
|
||||
return Collections.singletonMap("count", count);
|
||||
}
|
||||
}
|
||||
|
||||
private static final class JoinFixture {
|
||||
private final CountDownLatch branchACompleted = new CountDownLatch(1);
|
||||
private final CountDownLatch branchBStarted = new CountDownLatch(1);
|
||||
private final CountDownLatch releaseBranchB = new CountDownLatch(1);
|
||||
private final AtomicInteger joinExecutions = new AtomicInteger();
|
||||
private final AtomicInteger conditionChecks = new AtomicInteger();
|
||||
private ScheduledExecutorService schedulerPool;
|
||||
private ExecutorService workerPool;
|
||||
private TriggerScheduler scheduler;
|
||||
private InMemoryChainStateRepository chainStateRepository;
|
||||
private InMemoryNodeStateRepository nodeStateRepository;
|
||||
private ChainDefinition definition;
|
||||
private ChainExecutor executor;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
package com.easyagents.scheduler;
|
||||
|
||||
import java.time.Duration;
|
||||
|
||||
/**
|
||||
* 请求调度提供方持久化重试同一次逻辑触发的异常。
|
||||
*
|
||||
* <p>仅用于尚未产生业务副作用、且调用方必须保证至少一次登记的短 Handler。
|
||||
* {@code maxRefires=0} 表示不限制持久化重试次数;Provider 正常情况下应释放当前执行线程,
|
||||
* 并通过持久化的延迟触发保留原 fire 上下文。仅当 JobStore 无法写入持久重试时,
|
||||
* Provider 可按相同退避暂时占用当前线程并原地重试,以避免正常完成造成 fire 丢失。</p>
|
||||
*/
|
||||
public class ScheduleRefireException extends RuntimeException {
|
||||
|
||||
private static final int DEFAULT_MAX_REFIRES = 0;
|
||||
private static final Duration DEFAULT_BASE_DELAY = Duration.ofMillis(250);
|
||||
private final int maxRefires;
|
||||
private final Duration baseDelay;
|
||||
|
||||
public ScheduleRefireException(String message, Throwable cause) {
|
||||
this(message, cause, DEFAULT_MAX_REFIRES, DEFAULT_BASE_DELAY);
|
||||
}
|
||||
|
||||
public ScheduleRefireException(String message, Throwable cause,
|
||||
int maxRefires, Duration baseDelay) {
|
||||
super(requireMessage(message), cause);
|
||||
if (maxRefires < 0) throw new IllegalArgumentException("maxRefires must not be negative");
|
||||
if (baseDelay == null || baseDelay.isZero() || baseDelay.isNegative()) {
|
||||
throw new IllegalArgumentException("baseDelay must be positive");
|
||||
}
|
||||
this.maxRefires = maxRefires;
|
||||
this.baseDelay = baseDelay;
|
||||
}
|
||||
|
||||
/**
|
||||
* @return 最大重新执行次数,0 表示不限制
|
||||
*/
|
||||
public int maxRefires() {
|
||||
return maxRefires;
|
||||
}
|
||||
|
||||
public Duration delayFor(int refireNumber) {
|
||||
int shift = Math.min(4, Math.max(0, refireNumber - 1));
|
||||
return baseDelay.multipliedBy(1L << shift);
|
||||
}
|
||||
|
||||
private static String requireMessage(String message) {
|
||||
if (message == null || message.isBlank()) {
|
||||
throw new IllegalArgumentException("message must not be blank");
|
||||
}
|
||||
return message;
|
||||
}
|
||||
}
|
||||
@@ -29,6 +29,9 @@ abstract class AbstractDispatchJob implements InterruptableJob {
|
||||
throw new JobExecutionException("easy-agents scheduler runtime is not available");
|
||||
}
|
||||
quartzRuntime.execute(context);
|
||||
} catch (JobExecutionException exception) {
|
||||
// JobExecutionException 继承 SchedulerException,必须保留 refire 等控制语义。
|
||||
throw exception;
|
||||
} catch (SchedulerException exception) {
|
||||
throw new JobExecutionException("failed to access scheduler runtime", exception, false);
|
||||
} finally {
|
||||
|
||||
@@ -7,9 +7,16 @@ import com.easyagents.scheduler.ScheduleExecutionListener;
|
||||
import com.easyagents.scheduler.ScheduleFireContext;
|
||||
import com.easyagents.scheduler.ScheduleHandler;
|
||||
import com.easyagents.scheduler.ScheduleId;
|
||||
import com.easyagents.scheduler.ScheduleRefireException;
|
||||
import org.quartz.JobDataMap;
|
||||
import org.quartz.JobExecutionContext;
|
||||
import org.quartz.JobExecutionException;
|
||||
import org.quartz.ObjectAlreadyExistsException;
|
||||
import org.quartz.SchedulerException;
|
||||
import org.quartz.SimpleScheduleBuilder;
|
||||
import org.quartz.Trigger;
|
||||
import org.quartz.TriggerBuilder;
|
||||
import org.quartz.TriggerKey;
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
|
||||
@@ -19,6 +26,10 @@ import java.util.Collection;
|
||||
import java.util.LinkedHashMap;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.nio.charset.StandardCharsets;
|
||||
import java.util.Date;
|
||||
import java.util.UUID;
|
||||
import java.util.concurrent.locks.LockSupport;
|
||||
|
||||
/**
|
||||
* 当前 Scheduler 节点的 Handler 和监听器运行时。
|
||||
@@ -55,15 +66,27 @@ final class QuartzRuntime {
|
||||
void execute(JobExecutionContext quartzContext) throws JobExecutionException {
|
||||
JobDataMap data = quartzContext.getMergedJobDataMap();
|
||||
ScheduleDefinition definition = QuartzScheduleMapper.toDefinition(data);
|
||||
Instant actualFireTime = toInstant(quartzContext.getFireTime(), Instant.now());
|
||||
Instant observedActualFireTime = toInstant(quartzContext.getFireTime(), Instant.now());
|
||||
Instant scheduledFireTime = instantValue(
|
||||
data,
|
||||
QuartzScheduleMapper.KEY_RETRY_SCHEDULED_FIRE_TIME,
|
||||
toInstant(quartzContext.getScheduledFireTime(), observedActualFireTime)
|
||||
);
|
||||
Instant actualFireTime = instantValue(
|
||||
data,
|
||||
QuartzScheduleMapper.KEY_RETRY_ACTUAL_FIRE_TIME,
|
||||
observedActualFireTime
|
||||
);
|
||||
ScheduleFireContext context = new ScheduleFireContext(
|
||||
definition.id(),
|
||||
definition.handlerCode(),
|
||||
toInstant(quartzContext.getScheduledFireTime(), actualFireTime),
|
||||
scheduledFireTime,
|
||||
actualFireTime,
|
||||
quartzContext.getFireInstanceId(),
|
||||
stringValue(data, QuartzScheduleMapper.KEY_RETRY_FIRE_INSTANCE_ID,
|
||||
quartzContext.getFireInstanceId()),
|
||||
stringValue(data, QuartzScheduleMapper.KEY_INVOCATION),
|
||||
quartzContext.isRecovering(),
|
||||
booleanValue(data, QuartzScheduleMapper.KEY_RETRY_RECOVERING,
|
||||
quartzContext.isRecovering()),
|
||||
definition.parameters()
|
||||
);
|
||||
ScheduleHandler handler = handlers.get(definition.handlerCode());
|
||||
@@ -82,10 +105,80 @@ final class QuartzRuntime {
|
||||
notifySucceeded(context, elapsed(startedAt));
|
||||
} catch (Exception exception) {
|
||||
notifyFailed(context, elapsed(startedAt), exception);
|
||||
if (exception instanceof ScheduleRefireException retry) {
|
||||
int attempt = Math.max(
|
||||
intValue(data, QuartzScheduleMapper.KEY_RETRY_ATTEMPT, 0),
|
||||
quartzContext.getRefireCount()
|
||||
);
|
||||
if (retry.maxRefires() == 0 || attempt < retry.maxRefires()) {
|
||||
try {
|
||||
persistRetry(quartzContext, context, retry, attempt + 1);
|
||||
return;
|
||||
} catch (SchedulerException retryFailure) {
|
||||
exception.addSuppressed(retryFailure);
|
||||
log.error("Failed to persist schedule handler retry: scheduleId={}, attempt={}",
|
||||
context.scheduleId(), attempt + 1, retryFailure);
|
||||
if (canRefireInPlace(quartzContext, retry, attempt + 1)) {
|
||||
// JobStore 暂时不可写时,正常完成原 fire 会造成登记丢失。
|
||||
// 仅在这条降级路径短时占用当前 worker;一旦持久重试落库即释放。
|
||||
throw new JobExecutionException(exception, true);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
log.error("Schedule handler retry limit exhausted: scheduleId={}, retries={}",
|
||||
context.scheduleId(), retry.maxRefires(), exception);
|
||||
}
|
||||
}
|
||||
// 不在 Quartz worker 内无限 refire。持久化重试失败时保留原异常,
|
||||
// 由 requestRecovery 和集群故障恢复处理未完成的 fired trigger。
|
||||
throw new JobExecutionException(exception, false);
|
||||
}
|
||||
}
|
||||
|
||||
private static boolean canRefireInPlace(
|
||||
JobExecutionContext context,
|
||||
ScheduleRefireException retry,
|
||||
int attempt
|
||||
) {
|
||||
LockSupport.parkNanos(retry.delayFor(attempt).toNanos());
|
||||
if (Thread.currentThread().isInterrupted()) return false;
|
||||
try {
|
||||
return !context.getScheduler().isShutdown();
|
||||
} catch (SchedulerException exception) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
private static void persistRetry(
|
||||
JobExecutionContext quartzContext,
|
||||
ScheduleFireContext context,
|
||||
ScheduleRefireException retry,
|
||||
int attempt
|
||||
) throws SchedulerException {
|
||||
String source = context.scheduleId() + "|" + context.fireInstanceId() + "|"
|
||||
+ String.valueOf(context.invocationId()) + "|" + attempt;
|
||||
String retryName = "retry-" + UUID.nameUUIDFromBytes(
|
||||
source.getBytes(StandardCharsets.UTF_8));
|
||||
TriggerKey retryKey = new TriggerKey(
|
||||
retryName,
|
||||
QuartzScheduleMapper.GROUP_PREFIX + "retry." + context.scheduleId().namespace()
|
||||
);
|
||||
Trigger trigger = TriggerBuilder.newTrigger()
|
||||
.withIdentity(retryKey)
|
||||
.forJob(quartzContext.getJobDetail().getKey())
|
||||
.usingJobData(QuartzScheduleMapper.retryData(context, attempt))
|
||||
.startAt(Date.from(Instant.now().plus(retry.delayFor(attempt))))
|
||||
.withSchedule(SimpleScheduleBuilder.simpleSchedule()
|
||||
.withRepeatCount(0)
|
||||
.withMisfireHandlingInstructionFireNow())
|
||||
.build();
|
||||
try {
|
||||
quartzContext.getScheduler().scheduleJob(trigger);
|
||||
} catch (ObjectAlreadyExistsException ignored) {
|
||||
// 同一原始 fire/attempt 的确定性 key 已落库,即视为持久化成功。
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 向监听器发布 Misfire 事件。
|
||||
*
|
||||
@@ -175,4 +268,24 @@ final class QuartzRuntime {
|
||||
Object value = data.get(key);
|
||||
return value == null ? null : value.toString();
|
||||
}
|
||||
|
||||
private static String stringValue(JobDataMap data, String key, String fallback) {
|
||||
String value = stringValue(data, key);
|
||||
return value == null ? fallback : value;
|
||||
}
|
||||
|
||||
private static Instant instantValue(JobDataMap data, String key, Instant fallback) {
|
||||
String value = stringValue(data, key);
|
||||
return value == null ? fallback : Instant.ofEpochMilli(Long.parseLong(value));
|
||||
}
|
||||
|
||||
private static boolean booleanValue(JobDataMap data, String key, boolean fallback) {
|
||||
String value = stringValue(data, key);
|
||||
return value == null ? fallback : Boolean.parseBoolean(value);
|
||||
}
|
||||
|
||||
private static int intValue(JobDataMap data, String key, int fallback) {
|
||||
String value = stringValue(data, key);
|
||||
return value == null ? fallback : Integer.parseInt(value);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@ import com.easyagents.scheduler.OnceSchedulePlan;
|
||||
import com.easyagents.scheduler.ScheduleDefinition;
|
||||
import com.easyagents.scheduler.ScheduleErrorCode;
|
||||
import com.easyagents.scheduler.ScheduleException;
|
||||
import com.easyagents.scheduler.ScheduleFireContext;
|
||||
import com.easyagents.scheduler.ScheduleId;
|
||||
import com.easyagents.scheduler.SchedulePlan;
|
||||
import org.quartz.CronScheduleBuilder;
|
||||
@@ -47,6 +48,11 @@ final class QuartzScheduleMapper {
|
||||
static final String KEY_DESCRIPTION = "ea.description";
|
||||
static final String KEY_INVOCATION = "ea.invocationId";
|
||||
static final String KEY_IMMEDIATE_PARAMETER_SNAPSHOT = "ea.immediateParameterSnapshot";
|
||||
static final String KEY_RETRY_ATTEMPT = "ea.retry.attempt";
|
||||
static final String KEY_RETRY_SCHEDULED_FIRE_TIME = "ea.retry.scheduledFireTime";
|
||||
static final String KEY_RETRY_ACTUAL_FIRE_TIME = "ea.retry.actualFireTime";
|
||||
static final String KEY_RETRY_FIRE_INSTANCE_ID = "ea.retry.fireInstanceId";
|
||||
static final String KEY_RETRY_RECOVERING = "ea.retry.recovering";
|
||||
static final String PARAMETER_PREFIX = "ea.parameter.";
|
||||
static final String IMMEDIATE_PARAMETER_PREFIX = "ea.immediateParameter.";
|
||||
static final String GROUP_PREFIX = "ea.scheduler.";
|
||||
@@ -153,6 +159,30 @@ final class QuartzScheduleMapper {
|
||||
return data;
|
||||
}
|
||||
|
||||
/**
|
||||
* 固化持久化重试所需的原始 fire 上下文。
|
||||
*/
|
||||
static JobDataMap retryData(ScheduleFireContext context, int attempt) {
|
||||
JobDataMap data = identityData(context.scheduleId());
|
||||
// Trigger 数据覆盖 JobDetail;任务定义被 replace 后,既有 fire 的重试仍须
|
||||
// 派发给首次触发时的 Handler,而不是意外切换到新 Handler。
|
||||
data.put(KEY_HANDLER, context.handlerCode());
|
||||
data.put(KEY_RETRY_ATTEMPT, Integer.toString(attempt));
|
||||
data.put(KEY_RETRY_SCHEDULED_FIRE_TIME,
|
||||
Long.toString(context.scheduledFireTime().toEpochMilli()));
|
||||
data.put(KEY_RETRY_ACTUAL_FIRE_TIME,
|
||||
Long.toString(context.actualFireTime().toEpochMilli()));
|
||||
data.put(KEY_RETRY_FIRE_INSTANCE_ID, context.fireInstanceId());
|
||||
data.put(KEY_RETRY_RECOVERING, Boolean.toString(context.recovering()));
|
||||
if (context.invocationId() != null) {
|
||||
data.put(KEY_INVOCATION, context.invocationId());
|
||||
}
|
||||
// 重试必须沿用首次 fire 的参数快照,不能读取期间被替换的新定义参数。
|
||||
data.put(KEY_IMMEDIATE_PARAMETER_SNAPSHOT, Boolean.TRUE.toString());
|
||||
putParameters(data, context.parameters(), IMMEDIATE_PARAMETER_PREFIX);
|
||||
return data;
|
||||
}
|
||||
|
||||
/**
|
||||
* 从持久 JobData 恢复公共定义。
|
||||
*
|
||||
|
||||
@@ -335,7 +335,8 @@ public final class QuartzScheduleService implements ScheduleService, AutoCloseab
|
||||
}
|
||||
try {
|
||||
if (!scheduler.isShutdown()) {
|
||||
// false 保证 close 本身有界;Quartz 会向内部 InterruptableJob 发送中断。
|
||||
// Factory 启用 interruptJobsOnShutdown,Quartz 会先中断内部
|
||||
// InterruptableJob;false 只表示不再无界等待忽略中断的 Handler。
|
||||
scheduler.shutdown(false);
|
||||
}
|
||||
} catch (SchedulerException exception) {
|
||||
|
||||
@@ -10,6 +10,8 @@ package com.easyagents.scheduler.quartz;
|
||||
* @param clustered 是否启用 JDBC 集群
|
||||
* @param threadCount Quartz Worker 线程数
|
||||
* @param threadPriority Quartz Worker 线程优先级
|
||||
* @param batchTriggerAcquisitionMaxCount 单次批量获取 Trigger 的最大数量
|
||||
* @param batchTriggerAcquisitionFireAheadTimeWindowMillis 可提前纳入批量的时间窗口,单位毫秒
|
||||
* @param clusterCheckinIntervalMillis 集群心跳间隔,单位毫秒
|
||||
* @param misfireThresholdMillis Misfire 判定阈值,单位毫秒
|
||||
* @param waitForJobsToCompleteOnShutdown 关闭时是否等待运行中任务完成
|
||||
@@ -24,6 +26,8 @@ public record QuartzSchedulerConfig(
|
||||
boolean clustered,
|
||||
int threadCount,
|
||||
int threadPriority,
|
||||
int batchTriggerAcquisitionMaxCount,
|
||||
long batchTriggerAcquisitionFireAheadTimeWindowMillis,
|
||||
long clusterCheckinIntervalMillis,
|
||||
long misfireThresholdMillis,
|
||||
boolean waitForJobsToCompleteOnShutdown,
|
||||
@@ -55,6 +59,18 @@ public record QuartzSchedulerConfig(
|
||||
if (threadPriority < Thread.MIN_PRIORITY || threadPriority > Thread.MAX_PRIORITY) {
|
||||
throw new IllegalArgumentException("threadPriority must be between 1 and 10");
|
||||
}
|
||||
if (batchTriggerAcquisitionMaxCount < 1
|
||||
|| batchTriggerAcquisitionMaxCount > threadCount) {
|
||||
throw new IllegalArgumentException(
|
||||
"batchTriggerAcquisitionMaxCount must be between 1 and threadCount"
|
||||
);
|
||||
}
|
||||
if (batchTriggerAcquisitionFireAheadTimeWindowMillis < 0
|
||||
|| batchTriggerAcquisitionFireAheadTimeWindowMillis > 60_000L) {
|
||||
throw new IllegalArgumentException(
|
||||
"batchTriggerAcquisitionFireAheadTimeWindowMillis must be between 0 and 60000"
|
||||
);
|
||||
}
|
||||
if (clusterCheckinIntervalMillis < 1000) {
|
||||
throw new IllegalArgumentException(
|
||||
"clusterCheckinIntervalMillis must be at least 1000"
|
||||
@@ -85,6 +101,8 @@ public record QuartzSchedulerConfig(
|
||||
true,
|
||||
8,
|
||||
Thread.NORM_PRIORITY,
|
||||
1,
|
||||
0L,
|
||||
15_000L,
|
||||
60_000L,
|
||||
true,
|
||||
|
||||
@@ -154,7 +154,7 @@ public final class QuartzSchedulerFactory {
|
||||
* @param dataSourceName Quartz 内部 DataSource 名称
|
||||
* @return Quartz 属性
|
||||
*/
|
||||
private static Properties properties(
|
||||
static Properties properties(
|
||||
QuartzSchedulerConfig config,
|
||||
String dataSourceName
|
||||
) {
|
||||
@@ -173,6 +173,14 @@ public final class QuartzSchedulerFactory {
|
||||
"org.quartz.threadPool.threadPriority",
|
||||
Integer.toString(config.threadPriority())
|
||||
);
|
||||
properties.setProperty(
|
||||
"org.quartz.scheduler.batchTriggerAcquisitionMaxCount",
|
||||
Integer.toString(config.batchTriggerAcquisitionMaxCount())
|
||||
);
|
||||
properties.setProperty(
|
||||
"org.quartz.scheduler.batchTriggerAcquisitionFireAheadTimeWindow",
|
||||
Long.toString(config.batchTriggerAcquisitionFireAheadTimeWindowMillis())
|
||||
);
|
||||
properties.setProperty(
|
||||
"org.quartz.jobStore.class",
|
||||
"org.quartz.impl.jdbcjobstore.JobStoreTX"
|
||||
@@ -182,6 +190,10 @@ public final class QuartzSchedulerFactory {
|
||||
config.driverDelegateClass()
|
||||
);
|
||||
properties.setProperty("org.quartz.jobStore.useProperties", "true");
|
||||
properties.setProperty(
|
||||
"org.quartz.jobStore.acquireTriggersWithinLock",
|
||||
Boolean.toString(config.batchTriggerAcquisitionMaxCount() > 1)
|
||||
);
|
||||
properties.setProperty("org.quartz.jobStore.dataSource", dataSourceName);
|
||||
properties.setProperty("org.quartz.jobStore.tablePrefix", config.tablePrefix());
|
||||
properties.setProperty(
|
||||
|
||||
@@ -9,6 +9,7 @@ import com.easyagents.scheduler.ScheduleException;
|
||||
import com.easyagents.scheduler.ScheduleFireContext;
|
||||
import com.easyagents.scheduler.ScheduleHandler;
|
||||
import com.easyagents.scheduler.ScheduleId;
|
||||
import com.easyagents.scheduler.ScheduleRefireException;
|
||||
import org.h2.jdbcx.JdbcDataSource;
|
||||
import org.h2.tools.RunScript;
|
||||
import org.junit.Test;
|
||||
@@ -20,9 +21,11 @@ import java.sql.Connection;
|
||||
import java.sql.ResultSet;
|
||||
import java.sql.Statement;
|
||||
import java.time.Instant;
|
||||
import java.time.Duration;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.UUID;
|
||||
import java.util.concurrent.CopyOnWriteArrayList;
|
||||
import java.util.concurrent.CountDownLatch;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
|
||||
@@ -218,6 +221,83 @@ public class QuartzJdbcIntegrationTest {
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证延迟重试先落入 JDBC JobStore,关闭并重建节点后仍使用原 fire 上下文执行。
|
||||
*/
|
||||
@Test
|
||||
public void shouldRecoverPersistedRetryAfterSchedulerRestart() throws Exception {
|
||||
JdbcDataSource dataSource = dataSource();
|
||||
executeSchema(dataSource);
|
||||
QuartzSchedulerConfig config = jdbcConfig();
|
||||
CountDownLatch firstFailed = new CountDownLatch(1);
|
||||
CountDownLatch retryCompleted = new CountDownLatch(1);
|
||||
CopyOnWriteArrayList<ScheduleFireContext> contexts = new CopyOnWriteArrayList<>();
|
||||
ScheduleHandler retryingHandler = new ScheduleHandler() {
|
||||
@Override
|
||||
public String code() {
|
||||
return "persistent-retry-handler";
|
||||
}
|
||||
|
||||
@Override
|
||||
public void execute(ScheduleFireContext context) {
|
||||
contexts.add(context);
|
||||
if (contexts.size() == 1) {
|
||||
firstFailed.countDown();
|
||||
throw new ScheduleRefireException("temporary database failure",
|
||||
new IllegalStateException("unavailable"), 2,
|
||||
Duration.ofSeconds(2));
|
||||
}
|
||||
retryCompleted.countDown();
|
||||
}
|
||||
};
|
||||
ScheduleId scheduleId = new ScheduleId("jdbc", "persistent-retry");
|
||||
QuartzScheduleService first = QuartzSchedulerFactory.createJdbc(
|
||||
dataSource, config, List.of(retryingHandler), List.of());
|
||||
first.start();
|
||||
first.create(new ScheduleDefinition(
|
||||
scheduleId,
|
||||
retryingHandler.code(),
|
||||
new OnceSchedulePlan(Instant.parse("2099-01-01T00:00:00Z")),
|
||||
MisfirePolicy.FIRE_ONCE_NOW,
|
||||
ConcurrencyPolicy.DISALLOW,
|
||||
true,
|
||||
Map.of("snapshot", "original"),
|
||||
"persistent retry"
|
||||
));
|
||||
first.triggerNow(scheduleId, "persistent-invocation", Map.of());
|
||||
assertTrue(firstFailed.await(5, TimeUnit.SECONDS));
|
||||
awaitRetryTrigger(dataSource);
|
||||
first.close();
|
||||
|
||||
QuartzScheduleService restarted = QuartzSchedulerFactory.createJdbc(
|
||||
dataSource, config, List.of(retryingHandler), List.of());
|
||||
try {
|
||||
restarted.start();
|
||||
assertTrue("persisted retry did not execute after restart",
|
||||
retryCompleted.await(8, TimeUnit.SECONDS));
|
||||
assertEquals(contexts.get(0).scheduledFireTime(), contexts.get(1).scheduledFireTime());
|
||||
assertEquals(contexts.get(0).actualFireTime(), contexts.get(1).actualFireTime());
|
||||
assertEquals(contexts.get(0).fireInstanceId(), contexts.get(1).fireInstanceId());
|
||||
assertEquals("persistent-invocation", contexts.get(1).invocationId());
|
||||
} finally {
|
||||
restarted.close();
|
||||
}
|
||||
}
|
||||
|
||||
private static void awaitRetryTrigger(JdbcDataSource dataSource) throws Exception {
|
||||
long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(5);
|
||||
do {
|
||||
try (Connection connection = dataSource.getConnection();
|
||||
Statement statement = connection.createStatement();
|
||||
ResultSet result = statement.executeQuery(
|
||||
"SELECT COUNT(*) FROM QRTZ_TRIGGERS WHERE TRIGGER_GROUP LIKE 'ea.scheduler.retry.%'")) {
|
||||
if (result.next() && result.getInt(1) > 0) return;
|
||||
}
|
||||
Thread.sleep(25L);
|
||||
} while (System.nanoTime() < deadline);
|
||||
fail("persistent retry trigger was not stored");
|
||||
}
|
||||
|
||||
private static JdbcDataSource dataSource() {
|
||||
JdbcDataSource dataSource = new JdbcDataSource();
|
||||
dataSource.setURL("jdbc:h2:mem:scheduler-" + UUID.randomUUID() + ";DB_CLOSE_DELAY=-1");
|
||||
@@ -259,6 +339,8 @@ public class QuartzJdbcIntegrationTest {
|
||||
false,
|
||||
2,
|
||||
Thread.NORM_PRIORITY,
|
||||
1,
|
||||
0L,
|
||||
15_000L,
|
||||
1_000L,
|
||||
true,
|
||||
|
||||
@@ -0,0 +1,105 @@
|
||||
package com.easyagents.scheduler.quartz;
|
||||
|
||||
import com.easyagents.scheduler.ConcurrencyPolicy;
|
||||
import com.easyagents.scheduler.MisfirePolicy;
|
||||
import com.easyagents.scheduler.OnceSchedulePlan;
|
||||
import com.easyagents.scheduler.ScheduleDefinition;
|
||||
import com.easyagents.scheduler.ScheduleHandler;
|
||||
import com.easyagents.scheduler.ScheduleId;
|
||||
import com.easyagents.scheduler.ScheduleRefireException;
|
||||
import org.junit.Test;
|
||||
import org.quartz.JobDetail;
|
||||
import org.quartz.JobExecutionContext;
|
||||
import org.quartz.JobExecutionException;
|
||||
import org.quartz.Scheduler;
|
||||
import org.quartz.SchedulerException;
|
||||
import org.quartz.Trigger;
|
||||
|
||||
import java.lang.reflect.Proxy;
|
||||
import java.time.Duration;
|
||||
import java.time.Instant;
|
||||
import java.util.Date;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
import static org.junit.Assert.assertTrue;
|
||||
import static org.junit.Assert.fail;
|
||||
|
||||
/** {@link QuartzRuntime} 失败控制语义测试。 */
|
||||
public class QuartzRuntimeTest {
|
||||
|
||||
/** JobStore 无法保存延迟触发时不得把当前 fire 当作正常完成。 */
|
||||
@Test
|
||||
public void shouldRefireInPlaceWhenPersistentRetryCannotBeStored() throws Exception {
|
||||
ScheduleDefinition definition = new ScheduleDefinition(
|
||||
new ScheduleId("test", "retry-store-failure"),
|
||||
"handler",
|
||||
new OnceSchedulePlan(Instant.parse("2099-01-01T00:00:00Z")),
|
||||
MisfirePolicy.FIRE_ONCE_NOW,
|
||||
ConcurrencyPolicy.DISALLOW,
|
||||
true,
|
||||
Map.of(),
|
||||
"retry store failure"
|
||||
);
|
||||
JobDetail job = QuartzScheduleMapper.toJobDetail(definition);
|
||||
Scheduler scheduler = proxy(Scheduler.class, (proxy, method, arguments) -> {
|
||||
if ("scheduleJob".equals(method.getName())
|
||||
&& arguments != null && arguments.length == 1
|
||||
&& arguments[0] instanceof Trigger) {
|
||||
throw new SchedulerException("job store unavailable");
|
||||
}
|
||||
if ("isShutdown".equals(method.getName())) return false;
|
||||
return defaultValue(method.getReturnType());
|
||||
});
|
||||
Date fireTime = new Date();
|
||||
JobExecutionContext context = proxy(JobExecutionContext.class,
|
||||
(proxy, method, arguments) -> switch (method.getName()) {
|
||||
case "getMergedJobDataMap" -> job.getJobDataMap();
|
||||
case "getJobDetail" -> job;
|
||||
case "getScheduler" -> scheduler;
|
||||
case "getFireTime", "getScheduledFireTime" -> fireTime;
|
||||
case "getFireInstanceId" -> "fire-1";
|
||||
case "isRecovering" -> false;
|
||||
case "getRefireCount" -> 0;
|
||||
default -> defaultValue(method.getReturnType());
|
||||
});
|
||||
QuartzRuntime runtime = new QuartzRuntime(List.of(new ScheduleHandler() {
|
||||
@Override
|
||||
public String code() {
|
||||
return "handler";
|
||||
}
|
||||
|
||||
@Override
|
||||
public void execute(com.easyagents.scheduler.ScheduleFireContext context) {
|
||||
throw new ScheduleRefireException("registration unavailable",
|
||||
new IllegalStateException("database unavailable"), 0,
|
||||
Duration.ofMillis(1));
|
||||
}
|
||||
}), List.of());
|
||||
|
||||
try {
|
||||
runtime.execute(context);
|
||||
fail("expected refire request");
|
||||
} catch (JobExecutionException exception) {
|
||||
assertTrue(exception.refireImmediately());
|
||||
}
|
||||
}
|
||||
|
||||
@SuppressWarnings("unchecked")
|
||||
private static <T> T proxy(Class<T> type, java.lang.reflect.InvocationHandler handler) {
|
||||
return (T) Proxy.newProxyInstance(type.getClassLoader(), new Class<?>[]{type}, handler);
|
||||
}
|
||||
|
||||
private static Object defaultValue(Class<?> type) {
|
||||
if (!type.isPrimitive()) return null;
|
||||
if (type == boolean.class) return false;
|
||||
if (type == byte.class) return (byte) 0;
|
||||
if (type == short.class) return (short) 0;
|
||||
if (type == int.class) return 0;
|
||||
if (type == long.class) return 0L;
|
||||
if (type == float.class) return 0F;
|
||||
if (type == double.class) return 0D;
|
||||
if (type == char.class) return '\0';
|
||||
return null;
|
||||
}
|
||||
}
|
||||
@@ -10,20 +10,25 @@ import com.easyagents.scheduler.ScheduleException;
|
||||
import com.easyagents.scheduler.ScheduleFireContext;
|
||||
import com.easyagents.scheduler.ScheduleHandler;
|
||||
import com.easyagents.scheduler.ScheduleId;
|
||||
import com.easyagents.scheduler.ScheduleRefireException;
|
||||
import com.easyagents.scheduler.ScheduleStatus;
|
||||
import org.junit.After;
|
||||
import org.junit.Test;
|
||||
import org.quartz.JobDetail;
|
||||
import org.quartz.Scheduler;
|
||||
import org.quartz.impl.StdSchedulerFactory;
|
||||
import org.quartz.impl.matchers.GroupMatcher;
|
||||
|
||||
import java.nio.charset.StandardCharsets;
|
||||
import java.time.Duration;
|
||||
import java.time.Instant;
|
||||
import java.time.ZoneId;
|
||||
import java.util.Base64;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Properties;
|
||||
import java.util.UUID;
|
||||
import java.util.concurrent.CopyOnWriteArrayList;
|
||||
import java.util.concurrent.CountDownLatch;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
@@ -112,6 +117,108 @@ public class QuartzScheduleServiceTest {
|
||||
assertEquals("value", captured.get().parameters().get("stable"));
|
||||
}
|
||||
|
||||
/** 验证可恢复失败会持久化延迟重试,并保留原始 fire 上下文。 */
|
||||
@Test
|
||||
public void shouldRefireRetryableHandlerFailure() throws Exception {
|
||||
AtomicInteger attempts = new AtomicInteger();
|
||||
CountDownLatch completed = new CountDownLatch(1);
|
||||
CopyOnWriteArrayList<ScheduleFireContext> contexts = new CopyOnWriteArrayList<>();
|
||||
service = newRamService(context -> {
|
||||
contexts.add(context);
|
||||
if (attempts.incrementAndGet() == 1) {
|
||||
throw new ScheduleRefireException("temporary registration failure",
|
||||
new IllegalStateException("database unavailable"), 2,
|
||||
Duration.ofMillis(100));
|
||||
}
|
||||
completed.countDown();
|
||||
});
|
||||
ScheduleDefinition definition = onceDefinition(
|
||||
"retryable-immediate",
|
||||
MisfirePolicy.FIRE_ONCE_NOW,
|
||||
ConcurrencyPolicy.DISALLOW,
|
||||
Instant.parse("2099-01-01T00:00:00Z")
|
||||
);
|
||||
service.create(definition);
|
||||
|
||||
service.triggerNow(definition.id(), "retryable-invocation", Map.of());
|
||||
|
||||
assertTrue("retryable handler was not refired", completed.await(5, TimeUnit.SECONDS));
|
||||
assertEquals(2, attempts.get());
|
||||
assertEquals(contexts.get(0).scheduledFireTime(), contexts.get(1).scheduledFireTime());
|
||||
assertEquals(contexts.get(0).actualFireTime(), contexts.get(1).actualFireTime());
|
||||
assertEquals(contexts.get(0).fireInstanceId(), contexts.get(1).fireInstanceId());
|
||||
assertEquals("retryable-invocation", contexts.get(1).invocationId());
|
||||
}
|
||||
|
||||
/** 持久化延迟重试不得占用当前 Quartz worker。 */
|
||||
@Test
|
||||
public void shouldReleaseWorkerWhilePersistentRetryIsDelayed() throws Exception {
|
||||
CountDownLatch retryScheduled = new CountDownLatch(1);
|
||||
CountDownLatch healthyCompleted = new CountDownLatch(1);
|
||||
CountDownLatch retryCompleted = new CountDownLatch(1);
|
||||
AtomicInteger retryAttempts = new AtomicInteger();
|
||||
service = newRamService(context -> {
|
||||
if (context.scheduleId().name().equals("delayed-retry")) {
|
||||
if (retryAttempts.incrementAndGet() == 1) {
|
||||
retryScheduled.countDown();
|
||||
throw new ScheduleRefireException("temporary registration failure",
|
||||
new IllegalStateException("database unavailable"), 1,
|
||||
Duration.ofSeconds(1));
|
||||
}
|
||||
retryCompleted.countDown();
|
||||
} else {
|
||||
healthyCompleted.countDown();
|
||||
}
|
||||
}, 30_000L, 1);
|
||||
ScheduleDefinition retry = onceDefinition(
|
||||
"delayed-retry", MisfirePolicy.FIRE_ONCE_NOW,
|
||||
ConcurrencyPolicy.DISALLOW, Instant.parse("2099-01-01T00:00:00Z"));
|
||||
ScheduleDefinition healthy = onceDefinition(
|
||||
"healthy-during-retry", MisfirePolicy.FIRE_ONCE_NOW,
|
||||
ConcurrencyPolicy.DISALLOW, Instant.parse("2099-01-01T00:00:00Z"));
|
||||
service.create(retry);
|
||||
service.create(healthy);
|
||||
|
||||
service.triggerNow(retry.id(), "retry-1", Map.of());
|
||||
assertTrue(retryScheduled.await(5, TimeUnit.SECONDS));
|
||||
service.triggerNow(healthy.id(), "healthy-1", Map.of());
|
||||
|
||||
assertTrue("single Quartz worker remained occupied by delayed retry",
|
||||
healthyCompleted.await(750, TimeUnit.MILLISECONDS));
|
||||
assertTrue(retryCompleted.await(5, TimeUnit.SECONDS));
|
||||
}
|
||||
|
||||
/** 已持久化 fire 的重试不得因 replace 而切换到新 Handler。 */
|
||||
@Test
|
||||
public void shouldKeepOriginalHandlerWhenScheduleIsReplacedDuringRetry() throws Exception {
|
||||
AtomicInteger oldAttempts = new AtomicInteger();
|
||||
AtomicInteger newAttempts = new AtomicInteger();
|
||||
CountDownLatch oldRetryCompleted = new CountDownLatch(1);
|
||||
service = newRamServiceWithHandlers(List.of(
|
||||
handler("old-handler", context -> {
|
||||
if (oldAttempts.incrementAndGet() == 1) {
|
||||
throw new ScheduleRefireException("temporary registration failure",
|
||||
new IllegalStateException("database unavailable"), 1,
|
||||
Duration.ofMillis(400));
|
||||
}
|
||||
oldRetryCompleted.countDown();
|
||||
}),
|
||||
handler("new-handler", context -> newAttempts.incrementAndGet())
|
||||
), 30_000L, 2);
|
||||
ScheduleId id = new ScheduleId("test", "replace-during-retry");
|
||||
ScheduleDefinition original = definition(id, "old-handler");
|
||||
service.create(original);
|
||||
service.triggerNow(id, "replace-retry-1", Map.of());
|
||||
awaitPersistentRetryTrigger();
|
||||
|
||||
service.replace(definition(id, "new-handler"));
|
||||
|
||||
assertTrue("old handler retry did not complete",
|
||||
oldRetryCompleted.await(5, TimeUnit.SECONDS));
|
||||
assertEquals(2, oldAttempts.get());
|
||||
assertEquals(0, newAttempts.get());
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证立即触发会在返回回执前校验基础参数与覆盖参数的合并结果。
|
||||
*/
|
||||
@@ -385,6 +492,23 @@ public class QuartzScheduleServiceTest {
|
||||
private QuartzScheduleService newRamService(
|
||||
Consumer<ScheduleFireContext> handler,
|
||||
long shutdownWaitTimeoutMillis
|
||||
) throws Exception {
|
||||
return newRamService(handler, shutdownWaitTimeoutMillis, 2);
|
||||
}
|
||||
|
||||
private QuartzScheduleService newRamService(
|
||||
Consumer<ScheduleFireContext> handler,
|
||||
long shutdownWaitTimeoutMillis,
|
||||
int threadCount
|
||||
) throws Exception {
|
||||
return newRamServiceWithHandlers(List.of(handler("handler", handler)),
|
||||
shutdownWaitTimeoutMillis, threadCount);
|
||||
}
|
||||
|
||||
private QuartzScheduleService newRamServiceWithHandlers(
|
||||
List<ScheduleHandler> handlers,
|
||||
long shutdownWaitTimeoutMillis,
|
||||
int threadCount
|
||||
) throws Exception {
|
||||
Properties properties = new Properties();
|
||||
properties.setProperty(
|
||||
@@ -394,7 +518,7 @@ public class QuartzScheduleServiceTest {
|
||||
properties.setProperty("org.quartz.scheduler.instanceId", "NON_CLUSTERED");
|
||||
properties.setProperty("org.quartz.scheduler.interruptJobsOnShutdown", "true");
|
||||
properties.setProperty("org.quartz.threadPool.class", "org.quartz.simpl.SimpleThreadPool");
|
||||
properties.setProperty("org.quartz.threadPool.threadCount", "2");
|
||||
properties.setProperty("org.quartz.threadPool.threadCount", Integer.toString(threadCount));
|
||||
properties.setProperty("org.quartz.jobStore.class", "org.quartz.simpl.RAMJobStore");
|
||||
properties.setProperty("org.quartz.jobStore.misfireThreshold", "100");
|
||||
Scheduler scheduler = new StdSchedulerFactory(properties).getScheduler();
|
||||
@@ -402,17 +526,7 @@ public class QuartzScheduleServiceTest {
|
||||
scheduler,
|
||||
true,
|
||||
shutdownWaitTimeoutMillis,
|
||||
java.util.List.of(new ScheduleHandler() {
|
||||
@Override
|
||||
public String code() {
|
||||
return "handler";
|
||||
}
|
||||
|
||||
@Override
|
||||
public void execute(ScheduleFireContext context) {
|
||||
handler.accept(context);
|
||||
}
|
||||
}),
|
||||
handlers,
|
||||
java.util.List.of()
|
||||
);
|
||||
result.start();
|
||||
@@ -420,6 +534,50 @@ public class QuartzScheduleServiceTest {
|
||||
return result;
|
||||
}
|
||||
|
||||
private void awaitPersistentRetryTrigger() throws Exception {
|
||||
long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(5);
|
||||
while (System.nanoTime() < deadline) {
|
||||
boolean found = service.quartzScheduler()
|
||||
.getTriggerKeys(GroupMatcher.anyTriggerGroup())
|
||||
.stream()
|
||||
.anyMatch(key -> key.getGroup().startsWith(
|
||||
QuartzScheduleMapper.GROUP_PREFIX + "retry."));
|
||||
if (found) return;
|
||||
Thread.sleep(10L);
|
||||
}
|
||||
fail("persistent retry trigger was not created");
|
||||
}
|
||||
|
||||
private static ScheduleHandler handler(
|
||||
String code,
|
||||
Consumer<ScheduleFireContext> consumer
|
||||
) {
|
||||
return new ScheduleHandler() {
|
||||
@Override
|
||||
public String code() {
|
||||
return code;
|
||||
}
|
||||
|
||||
@Override
|
||||
public void execute(ScheduleFireContext context) {
|
||||
consumer.accept(context);
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
private static ScheduleDefinition definition(ScheduleId id, String handlerCode) {
|
||||
return new ScheduleDefinition(
|
||||
id,
|
||||
handlerCode,
|
||||
new OnceSchedulePlan(Instant.parse("2099-01-01T00:00:00Z")),
|
||||
MisfirePolicy.FIRE_ONCE_NOW,
|
||||
ConcurrencyPolicy.DISALLOW,
|
||||
true,
|
||||
Map.of(),
|
||||
id.name()
|
||||
);
|
||||
}
|
||||
|
||||
private static ScheduleDefinition cronDefinition(String handlerCode, String expression) {
|
||||
return new ScheduleDefinition(
|
||||
new ScheduleId("test", "lifecycle"),
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
package com.easyagents.scheduler.quartz;
|
||||
|
||||
import org.junit.Test;
|
||||
|
||||
import java.util.Properties;
|
||||
|
||||
import static org.junit.Assert.assertEquals;
|
||||
|
||||
/** {@link QuartzSchedulerFactory} 原生属性映射测试。 */
|
||||
public class QuartzSchedulerFactoryTest {
|
||||
|
||||
@Test
|
||||
public void batchAcquisitionMustRunWithinJobStoreLock() {
|
||||
QuartzSchedulerConfig config = new QuartzSchedulerConfig(
|
||||
"batch-scheduler",
|
||||
"NON_CLUSTERED",
|
||||
"QRTZ_",
|
||||
QuartzSchedulerConfig.STANDARD_JDBC_DELEGATE,
|
||||
false,
|
||||
8,
|
||||
Thread.NORM_PRIORITY,
|
||||
8,
|
||||
1_000L,
|
||||
15_000L,
|
||||
60_000L,
|
||||
true,
|
||||
30_000L,
|
||||
false
|
||||
);
|
||||
|
||||
Properties properties = QuartzSchedulerFactory.properties(config, "testDs");
|
||||
|
||||
assertEquals("8", properties.getProperty(
|
||||
"org.quartz.scheduler.batchTriggerAcquisitionMaxCount"));
|
||||
assertEquals("1000", properties.getProperty(
|
||||
"org.quartz.scheduler.batchTriggerAcquisitionFireAheadTimeWindow"));
|
||||
assertEquals("true", properties.getProperty(
|
||||
"org.quartz.jobStore.acquireTriggersWithinLock"));
|
||||
}
|
||||
}
|
||||
@@ -53,6 +53,11 @@
|
||||
<artifactId>h2</artifactId>
|
||||
<scope>test</scope>
|
||||
</dependency>
|
||||
<dependency>
|
||||
<groupId>org.springframework</groupId>
|
||||
<artifactId>spring-jdbc</artifactId>
|
||||
<scope>test</scope>
|
||||
</dependency>
|
||||
<dependency>
|
||||
<groupId>junit</groupId>
|
||||
<artifactId>junit</artifactId>
|
||||
|
||||
@@ -18,6 +18,7 @@ import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean
|
||||
import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty;
|
||||
import org.springframework.boot.autoconfigure.jdbc.DataSourceAutoConfiguration;
|
||||
import org.springframework.boot.context.properties.EnableConfigurationProperties;
|
||||
import org.springframework.boot.sql.init.dependency.DependsOnDatabaseInitialization;
|
||||
import org.springframework.context.annotation.Bean;
|
||||
import org.springframework.context.annotation.Configuration;
|
||||
|
||||
@@ -61,6 +62,7 @@ public class EasyAgentsSchedulerAutoConfiguration {
|
||||
*/
|
||||
@Bean(name = SCHEDULER_BEAN_NAME, initMethod = "start", destroyMethod = "close")
|
||||
@ConditionalOnMissingBean(ScheduleService.class)
|
||||
@DependsOnDatabaseInitialization
|
||||
public QuartzScheduleService easyAgentsQuartzScheduleService(
|
||||
EasyAgentsSchedulerProperties properties,
|
||||
ListableBeanFactory beanFactory,
|
||||
@@ -83,6 +85,8 @@ public class EasyAgentsSchedulerAutoConfiguration {
|
||||
quartz.isClustered(),
|
||||
quartz.getThreadCount(),
|
||||
quartz.getThreadPriority(),
|
||||
quartz.getBatchTriggerAcquisitionMaxCount(),
|
||||
quartz.getBatchTriggerAcquisitionFireAheadTimeWindowMillis(),
|
||||
quartz.getClusterCheckinIntervalMillis(),
|
||||
quartz.getMisfireThresholdMillis(),
|
||||
quartz.isWaitForJobsToCompleteOnShutdown(),
|
||||
|
||||
@@ -104,6 +104,12 @@ public class EasyAgentsSchedulerProperties {
|
||||
/** Quartz Worker 线程优先级。 */
|
||||
private int threadPriority = Thread.NORM_PRIORITY;
|
||||
|
||||
/** 单次批量获取 Trigger 的最大数量。 */
|
||||
private int batchTriggerAcquisitionMaxCount = 1;
|
||||
|
||||
/** 可提前纳入批量的时间窗口,单位毫秒。 */
|
||||
private long batchTriggerAcquisitionFireAheadTimeWindowMillis;
|
||||
|
||||
/** 集群心跳间隔,单位毫秒。 */
|
||||
private long clusterCheckinIntervalMillis = 15_000L;
|
||||
|
||||
@@ -251,6 +257,45 @@ public class EasyAgentsSchedulerProperties {
|
||||
this.threadPriority = threadPriority;
|
||||
}
|
||||
|
||||
/**
|
||||
* 返回单次批量获取 Trigger 的最大数量。
|
||||
*
|
||||
* @return 批量上限
|
||||
*/
|
||||
public int getBatchTriggerAcquisitionMaxCount() {
|
||||
return batchTriggerAcquisitionMaxCount;
|
||||
}
|
||||
|
||||
/**
|
||||
* 设置单次批量获取 Trigger 的最大数量。
|
||||
*
|
||||
* @param batchTriggerAcquisitionMaxCount 批量上限
|
||||
*/
|
||||
public void setBatchTriggerAcquisitionMaxCount(
|
||||
int batchTriggerAcquisitionMaxCount) {
|
||||
this.batchTriggerAcquisitionMaxCount = batchTriggerAcquisitionMaxCount;
|
||||
}
|
||||
|
||||
/**
|
||||
* 返回可提前纳入批量的时间窗口。
|
||||
*
|
||||
* @return 毫秒窗口
|
||||
*/
|
||||
public long getBatchTriggerAcquisitionFireAheadTimeWindowMillis() {
|
||||
return batchTriggerAcquisitionFireAheadTimeWindowMillis;
|
||||
}
|
||||
|
||||
/**
|
||||
* 设置可提前纳入批量的时间窗口。
|
||||
*
|
||||
* @param batchTriggerAcquisitionFireAheadTimeWindowMillis 毫秒窗口
|
||||
*/
|
||||
public void setBatchTriggerAcquisitionFireAheadTimeWindowMillis(
|
||||
long batchTriggerAcquisitionFireAheadTimeWindowMillis) {
|
||||
this.batchTriggerAcquisitionFireAheadTimeWindowMillis =
|
||||
batchTriggerAcquisitionFireAheadTimeWindowMillis;
|
||||
}
|
||||
|
||||
/**
|
||||
* 返回集群心跳间隔。
|
||||
*
|
||||
|
||||
@@ -114,6 +114,47 @@ public class EasyAgentsSchedulerAutoConfigurationTest {
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证调度器等待 Spring Boot 数据库脚本初始化完成后再启动。
|
||||
*/
|
||||
@Test
|
||||
public void shouldWaitForDatabaseInitializationBeforeStartingScheduler() throws Exception {
|
||||
JdbcDataSource dataSource = dataSource();
|
||||
Map<String, Object> properties = enabledProperties();
|
||||
properties.put("easy-agents.scheduler.data-source-bean-name", "schedulerDataSource");
|
||||
properties.put("spring.sql.init.mode", "always");
|
||||
properties.put(
|
||||
"spring.sql.init.schema-locations",
|
||||
"classpath:quartz-schema/h2-2.5.2.sql"
|
||||
);
|
||||
|
||||
SpringApplication application = new SpringApplication(AutoDiscoveryApplication.class);
|
||||
application.setWebApplicationType(WebApplicationType.NONE);
|
||||
application.setDefaultProperties(properties);
|
||||
application.addInitializers(applicationContext -> {
|
||||
GenericApplicationContext genericContext =
|
||||
(GenericApplicationContext) applicationContext;
|
||||
genericContext.registerBean(
|
||||
"schedulerDataSource",
|
||||
DataSource.class,
|
||||
() -> dataSource
|
||||
);
|
||||
});
|
||||
|
||||
try (ConfigurableApplicationContext context = application.run()) {
|
||||
assertNotNull(context.getBean(ScheduleService.class));
|
||||
try (
|
||||
Connection connection = dataSource.getConnection();
|
||||
Statement statement = connection.createStatement();
|
||||
ResultSet resultSet = statement.executeQuery(
|
||||
"SELECT COUNT(*) FROM QRTZ_SCHEDULER_STATE"
|
||||
)
|
||||
) {
|
||||
assertTrue(resultSet.next());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证多个 DataSource 未明确选择时启动失败并提供可操作信息。
|
||||
*/
|
||||
@@ -162,6 +203,13 @@ public class EasyAgentsSchedulerAutoConfigurationTest {
|
||||
properties.put("easy-agents.scheduler.quartz.instance-id", "NON_CLUSTERED");
|
||||
properties.put("easy-agents.scheduler.quartz.clustered", "false");
|
||||
properties.put("easy-agents.scheduler.quartz.thread-count", "2");
|
||||
properties.put(
|
||||
"easy-agents.scheduler.quartz.batch-trigger-acquisition-max-count", "2"
|
||||
);
|
||||
properties.put(
|
||||
"easy-agents.scheduler.quartz.batch-trigger-acquisition-fire-ahead-time-window-millis",
|
||||
"500"
|
||||
);
|
||||
properties.put("easy-agents.scheduler.quartz.misfire-threshold-millis", "1000");
|
||||
return properties;
|
||||
}
|
||||
|
||||
@@ -199,10 +199,13 @@ public class LuceneSearcher implements DocumentSearcher, AutoCloseable {
|
||||
@Override
|
||||
public List<Document> searchDocuments(KeywordSearchRequest request) {
|
||||
List<Document> results = new ArrayList<>();
|
||||
if (request == null || request.getKeyword() == null || request.getKeyword().trim().isEmpty()) {
|
||||
return results;
|
||||
}
|
||||
try (IndexReader reader = DirectoryReader.open(directory)) {
|
||||
IndexSearcher searcher = new IndexSearcher(reader);
|
||||
Query query = buildQuery(request);
|
||||
TopDocs topDocs = searcher.search(query, request == null ? 10 : request.getCount());
|
||||
TopDocs topDocs = searcher.search(query, request.getCount());
|
||||
for (ScoreDoc scoreDoc : topDocs.scoreDocs) {
|
||||
org.apache.lucene.document.Document doc = searcher.doc(scoreDoc.doc);
|
||||
Document resultDoc = new Document();
|
||||
@@ -224,29 +227,24 @@ public class LuceneSearcher implements DocumentSearcher, AutoCloseable {
|
||||
return results;
|
||||
}
|
||||
|
||||
Query buildQuery(KeywordSearchRequest request) {
|
||||
try {
|
||||
String keyword = request == null ? null : request.getKeyword();
|
||||
Query buildQuery(KeywordSearchRequest request) throws ParseException {
|
||||
String escapedKeyword = QueryParser.escape(request.getKeyword());
|
||||
|
||||
QueryParser titleQueryParser = new QueryParser("title", analyzer);
|
||||
Query titleQuery = titleQueryParser.parse(keyword);
|
||||
BooleanClause titleBooleanClause = new BooleanClause(titleQuery, BooleanClause.Occur.SHOULD);
|
||||
QueryParser titleQueryParser = new QueryParser("title", analyzer);
|
||||
Query titleQuery = titleQueryParser.parse(escapedKeyword);
|
||||
BooleanClause titleBooleanClause = new BooleanClause(titleQuery, BooleanClause.Occur.SHOULD);
|
||||
|
||||
QueryParser contentQueryParser = new QueryParser("content", analyzer);
|
||||
Query contentQuery = contentQueryParser.parse(keyword);
|
||||
BooleanClause contentBooleanClause = new BooleanClause(contentQuery, BooleanClause.Occur.SHOULD);
|
||||
QueryParser contentQueryParser = new QueryParser("content", analyzer);
|
||||
Query contentQuery = contentQueryParser.parse(escapedKeyword);
|
||||
BooleanClause contentBooleanClause = new BooleanClause(contentQuery, BooleanClause.Occur.SHOULD);
|
||||
|
||||
BooleanQuery.Builder builder = new BooleanQuery.Builder();
|
||||
builder.add(titleBooleanClause)
|
||||
.add(contentBooleanClause);
|
||||
if (request != null && request.getKnowledgeId() != null && !request.getKnowledgeId().trim().isEmpty()) {
|
||||
builder.add(new TermQuery(new Term(KeywordSearchMetadataKeys.KNOWLEDGE_ID, request.getKnowledgeId().trim())), BooleanClause.Occur.MUST);
|
||||
}
|
||||
return builder.build();
|
||||
} catch (ParseException e) {
|
||||
LOG.error(e.toString(), e);
|
||||
BooleanQuery.Builder builder = new BooleanQuery.Builder();
|
||||
builder.add(titleBooleanClause)
|
||||
.add(contentBooleanClause);
|
||||
if (request.getKnowledgeId() != null && !request.getKnowledgeId().trim().isEmpty()) {
|
||||
builder.add(new TermQuery(new Term(KeywordSearchMetadataKeys.KNOWLEDGE_ID, request.getKnowledgeId().trim())), BooleanClause.Occur.MUST);
|
||||
}
|
||||
return null;
|
||||
return builder.build();
|
||||
}
|
||||
|
||||
private static Analyzer createAnalyzer() {
|
||||
|
||||
@@ -57,6 +57,45 @@ public class LuceneSearcherTest {
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证用户输入中的 Lucene 特殊字符按普通文本检索。
|
||||
*
|
||||
* @throws Exception 临时目录或 Lucene 资源操作失败时抛出
|
||||
*/
|
||||
@Test
|
||||
public void shouldTreatLuceneSpecialCharactersAsPlainText() throws Exception {
|
||||
Path tempDir = Files.createTempDirectory("lucene-searcher-special-character-test");
|
||||
LuceneConfig config = new LuceneConfig();
|
||||
config.setIndexDirPath(tempDir.toString());
|
||||
try (LuceneSearcher searcher = new LuceneSearcher(config)) {
|
||||
Document document = new Document();
|
||||
document.setId("special-character");
|
||||
document.setContent("A/C.\\nPLEASE");
|
||||
Assert.assertTrue(searcher.addDocument(document));
|
||||
|
||||
List<Document> results = searcher.searchDocuments("A/C.\\nPLEASE", 10);
|
||||
|
||||
Assert.assertEquals(1, results.size());
|
||||
Assert.assertEquals("special-character", String.valueOf(results.get(0).getId()));
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证空查询直接返回空结果。
|
||||
*
|
||||
* @throws Exception 临时目录或 Lucene 资源操作失败时抛出
|
||||
*/
|
||||
@Test
|
||||
public void shouldReturnEmptyForMissingKeyword() throws Exception {
|
||||
Path tempDir = Files.createTempDirectory("lucene-searcher-empty-keyword-test");
|
||||
LuceneConfig config = new LuceneConfig();
|
||||
config.setIndexDirPath(tempDir.toString());
|
||||
try (LuceneSearcher searcher = new LuceneSearcher(config)) {
|
||||
Assert.assertTrue(searcher.searchDocuments((KeywordSearchRequest) null).isEmpty());
|
||||
Assert.assertTrue(searcher.searchDocuments(" ", 10).isEmpty());
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证多个导入线程可共享同一个 IndexWriter 完成批量写入。
|
||||
*
|
||||
|
||||
@@ -25,7 +25,6 @@
|
||||
<dependency>
|
||||
<groupId>io.milvus</groupId>
|
||||
<artifactId>milvus-sdk-java</artifactId>
|
||||
<version>2.4.1</version>
|
||||
</dependency>
|
||||
<dependency>
|
||||
<groupId>junit</groupId>
|
||||
|
||||
@@ -0,0 +1,617 @@
|
||||
/*
|
||||
* Copyright (c) 2023-2026, Easy-Agents (fuhai999@gmail.com).
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
*/
|
||||
package com.easyagents.store.milvus;
|
||||
|
||||
import com.easyagents.core.util.StringUtil;
|
||||
import io.grpc.Context;
|
||||
import io.milvus.pool.MilvusClientV2Pool;
|
||||
import io.milvus.pool.PoolConfig;
|
||||
import io.milvus.v2.client.ConnectConfig;
|
||||
import io.milvus.v2.client.MilvusClientV2;
|
||||
import io.milvus.v2.client.RetryConfig;
|
||||
|
||||
import java.net.URI;
|
||||
import java.nio.charset.StandardCharsets;
|
||||
import java.security.MessageDigest;
|
||||
import java.security.NoSuchAlgorithmException;
|
||||
import java.time.Duration;
|
||||
import java.util.Collections;
|
||||
import java.util.HashSet;
|
||||
import java.util.Set;
|
||||
import java.util.concurrent.Callable;
|
||||
import java.util.concurrent.CancellationException;
|
||||
import java.util.concurrent.CompletableFuture;
|
||||
import java.util.concurrent.ConcurrentHashMap;
|
||||
import java.util.concurrent.ConcurrentMap;
|
||||
import java.util.concurrent.locks.ReentrantReadWriteLock;
|
||||
import java.util.function.Function;
|
||||
|
||||
/**
|
||||
* Shared Milvus client pool and collection state.
|
||||
*/
|
||||
public class MilvusClientManager implements AutoCloseable {
|
||||
|
||||
private static final String POOL_KEY = "default";
|
||||
private static final RetryConfig SINGLE_ATTEMPT_RETRY_CONFIG = RetryConfig.builder()
|
||||
.maxRetryTimes(1)
|
||||
.retryOnRateLimit(false)
|
||||
.maxRetryTimeoutMs(0L)
|
||||
.build();
|
||||
|
||||
private final ReentrantReadWriteLock lifecycleLock = new ReentrantReadWriteLock();
|
||||
private final Set<String> initializedCollections =
|
||||
Collections.synchronizedSet(new HashSet<String>());
|
||||
private final Set<String> loadedCollections =
|
||||
Collections.synchronizedSet(new HashSet<String>());
|
||||
private final ConcurrentMap<String, CollectionLoadTicket> collectionLoads =
|
||||
new ConcurrentHashMap<String, CollectionLoadTicket>();
|
||||
private final Set<Context.CancellableContext> activeContexts =
|
||||
ConcurrentHashMap.newKeySet();
|
||||
private final ConcurrentMap<Thread, Integer> activeOperations =
|
||||
new ConcurrentHashMap<Thread, Integer>();
|
||||
private volatile ManagedMilvusClientV2Pool pool;
|
||||
private volatile String poolFingerprint;
|
||||
private volatile long poolGeneration;
|
||||
private volatile boolean acceptingOperations = true;
|
||||
private volatile boolean closed;
|
||||
|
||||
public MilvusClientManager(MilvusVectorStoreConfig config) {
|
||||
PoolSettings settings = PoolSettings.from(config);
|
||||
this.pool = createPool(settings);
|
||||
this.poolFingerprint = fingerprint(settings);
|
||||
}
|
||||
|
||||
private static ManagedMilvusClientV2Pool createPool(PoolSettings settings) {
|
||||
ConnectConfig connectConfig = buildConnectConfig(settings);
|
||||
PoolConfig poolConfig = PoolConfig.builder()
|
||||
.maxTotal(settings.poolMaxTotal())
|
||||
.maxTotalPerKey(settings.poolMaxTotalPerKey())
|
||||
.maxIdlePerKey(settings.poolMaxIdlePerKey())
|
||||
.minIdlePerKey(settings.poolMinIdlePerKey())
|
||||
.blockWhenExhausted(true)
|
||||
.maxBlockWaitDuration(Duration.ofMillis(settings.poolMaxWaitMillis()))
|
||||
.evictionPollingInterval(Duration.ofMillis(settings.poolEvictionIntervalMillis()))
|
||||
.minEvictableIdleDuration(Duration.ofMillis(settings.poolMinEvictableIdleMillis()))
|
||||
.testOnBorrow(true)
|
||||
.testOnReturn(false)
|
||||
.build();
|
||||
try {
|
||||
return new ManagedMilvusClientV2Pool(poolConfig, connectConfig);
|
||||
} catch (ReflectiveOperationException exception) {
|
||||
throw new IllegalStateException("Unable to initialize Milvus client pool", exception);
|
||||
}
|
||||
}
|
||||
|
||||
public <T> T withClient(Function<MilvusClientV2, T> operation) {
|
||||
return withClient(null, operation);
|
||||
}
|
||||
|
||||
public <T> T withClient(
|
||||
Duration maxWait,
|
||||
Function<MilvusClientV2, T> operation
|
||||
) {
|
||||
Thread operationThread = registerActiveOperation();
|
||||
lifecycleLock.readLock().lock();
|
||||
try {
|
||||
ManagedMilvusClientV2Pool currentPool = requireOpenPool();
|
||||
MilvusClientV2 client = maxWait == null
|
||||
? currentPool.getClient(POOL_KEY)
|
||||
: currentPool.getClient(POOL_KEY, maxWait);
|
||||
if (client == null) {
|
||||
throw new IllegalStateException(
|
||||
"Milvus client pool is exhausted or unavailable"
|
||||
);
|
||||
}
|
||||
Throwable operationFailure = null;
|
||||
try {
|
||||
client.retryConfig(SINGLE_ATTEMPT_RETRY_CONFIG);
|
||||
Context.CancellableContext operationContext =
|
||||
Context.current().withCancellation();
|
||||
try {
|
||||
return withRequestContext(operationContext,
|
||||
() -> operation.apply(client));
|
||||
} catch (RuntimeException | Error exception) {
|
||||
throw exception;
|
||||
} catch (Exception exception) {
|
||||
throw new IllegalStateException(
|
||||
"Milvus client operation failed", exception);
|
||||
} finally {
|
||||
operationContext.cancel(null);
|
||||
}
|
||||
} catch (RuntimeException | Error exception) {
|
||||
operationFailure = exception;
|
||||
throw exception;
|
||||
} finally {
|
||||
RuntimeException cleanupFailure = null;
|
||||
try {
|
||||
releaseClient(currentPool, client);
|
||||
} catch (RuntimeException exception) {
|
||||
cleanupFailure = exception;
|
||||
}
|
||||
if (cleanupFailure != null) {
|
||||
if (operationFailure == null) {
|
||||
throw cleanupFailure;
|
||||
}
|
||||
operationFailure.addSuppressed(cleanupFailure);
|
||||
}
|
||||
}
|
||||
} finally {
|
||||
unregisterActiveOperation(operationThread);
|
||||
lifecycleLock.readLock().unlock();
|
||||
}
|
||||
}
|
||||
|
||||
private Thread registerActiveOperation() {
|
||||
ensureAcceptingOperations();
|
||||
Thread currentThread = Thread.currentThread();
|
||||
activeOperations.merge(currentThread, 1, Integer::sum);
|
||||
if (!acceptingOperations) {
|
||||
unregisterActiveOperation(currentThread);
|
||||
ensureAcceptingOperations();
|
||||
}
|
||||
return currentThread;
|
||||
}
|
||||
|
||||
private void unregisterActiveOperation(Thread operationThread) {
|
||||
activeOperations.computeIfPresent(operationThread,
|
||||
(thread, depth) -> depth <= 1 ? null : depth - 1);
|
||||
}
|
||||
|
||||
private void releaseClient(
|
||||
ManagedMilvusClientV2Pool currentPool,
|
||||
MilvusClientV2 client
|
||||
) {
|
||||
RuntimeException readinessFailure = null;
|
||||
boolean reusable = false;
|
||||
try {
|
||||
reusable = client.clientIsReady();
|
||||
} catch (RuntimeException exception) {
|
||||
readinessFailure = exception;
|
||||
}
|
||||
try {
|
||||
if (reusable) {
|
||||
currentPool.returnClient(POOL_KEY, client);
|
||||
} else {
|
||||
discardFailedClient(currentPool, client);
|
||||
}
|
||||
} catch (RuntimeException cleanupFailure) {
|
||||
if (readinessFailure == null) {
|
||||
throw cleanupFailure;
|
||||
}
|
||||
readinessFailure.addSuppressed(cleanupFailure);
|
||||
}
|
||||
if (readinessFailure != null) {
|
||||
throw readinessFailure;
|
||||
}
|
||||
}
|
||||
|
||||
<T> T withRequestContext(
|
||||
Context.CancellableContext context,
|
||||
Callable<T> operation
|
||||
) throws Exception {
|
||||
ensureAcceptingOperations();
|
||||
activeContexts.add(context);
|
||||
if (!acceptingOperations) {
|
||||
activeContexts.remove(context);
|
||||
context.cancel(new CancellationException(
|
||||
"Milvus client pool is unavailable"));
|
||||
ensureAcceptingOperations();
|
||||
}
|
||||
try {
|
||||
return context.call(operation);
|
||||
} finally {
|
||||
activeContexts.remove(context);
|
||||
}
|
||||
}
|
||||
|
||||
private void discardFailedClient(
|
||||
ManagedMilvusClientV2Pool currentPool,
|
||||
MilvusClientV2 client
|
||||
) {
|
||||
RuntimeException cleanupFailure = null;
|
||||
try {
|
||||
currentPool.invalidateClient(POOL_KEY, client);
|
||||
} catch (RuntimeException exception) {
|
||||
cleanupFailure = exception;
|
||||
}
|
||||
if (cleanupFailure != null) {
|
||||
throw cleanupFailure;
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Rebuilds the pool when connection or pool settings change.
|
||||
* Active operations are cancelled before the old pool is closed.
|
||||
*
|
||||
* @return true when a new pool was installed
|
||||
*/
|
||||
public synchronized boolean reconfigureIfNeeded(MilvusVectorStoreConfig config) {
|
||||
PoolSettings nextSettings = PoolSettings.from(config);
|
||||
String nextFingerprint = fingerprint(nextSettings);
|
||||
if (nextFingerprint.equals(poolFingerprint)) {
|
||||
return false;
|
||||
}
|
||||
acceptingOperations = false;
|
||||
cancelActiveContexts("Milvus client pool is reconfiguring");
|
||||
interruptActiveOperations();
|
||||
lifecycleLock.writeLock().lock();
|
||||
try {
|
||||
if (closed) {
|
||||
throw new IllegalStateException("Milvus client pool is closed");
|
||||
}
|
||||
if (nextFingerprint.equals(poolFingerprint)) {
|
||||
return false;
|
||||
}
|
||||
ManagedMilvusClientV2Pool replacement = createPool(nextSettings);
|
||||
ManagedMilvusClientV2Pool previous = pool;
|
||||
pool = replacement;
|
||||
poolFingerprint = nextFingerprint;
|
||||
poolGeneration++;
|
||||
initializedCollections.clear();
|
||||
loadedCollections.clear();
|
||||
failCollectionLoads("Milvus client pool was reconfigured");
|
||||
if (previous != null) {
|
||||
previous.close();
|
||||
}
|
||||
return true;
|
||||
} finally {
|
||||
lifecycleLock.writeLock().unlock();
|
||||
if (!closed) {
|
||||
acceptingOperations = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
boolean isCollectionInitialized(String collectionName) {
|
||||
return initializedCollections.contains(collectionName);
|
||||
}
|
||||
|
||||
Object initializedCollectionsLock() {
|
||||
return initializedCollections;
|
||||
}
|
||||
|
||||
void markCollectionInitialized(String collectionName) {
|
||||
initializedCollections.add(collectionName);
|
||||
}
|
||||
|
||||
boolean isCollectionLoaded(String collectionName) {
|
||||
return loadedCollections.contains(collectionName);
|
||||
}
|
||||
|
||||
void markCollectionLoaded(String collectionName) {
|
||||
loadedCollections.add(collectionName);
|
||||
}
|
||||
|
||||
void markCollectionUnloaded(String collectionName) {
|
||||
loadedCollections.remove(collectionName);
|
||||
}
|
||||
|
||||
CollectionLoadTicket beginCollectionLoad(String collectionName) {
|
||||
ensureAcceptingOperations();
|
||||
lifecycleLock.readLock().lock();
|
||||
try {
|
||||
requireOpenPool();
|
||||
CollectionLoadTicket candidate = new CollectionLoadTicket(
|
||||
collectionName,
|
||||
poolGeneration,
|
||||
new CompletableFuture<Void>(),
|
||||
true
|
||||
);
|
||||
CollectionLoadTicket existing = collectionLoads.putIfAbsent(
|
||||
collectionName, candidate);
|
||||
return existing == null ? candidate : existing.asFollower();
|
||||
} finally {
|
||||
lifecycleLock.readLock().unlock();
|
||||
}
|
||||
}
|
||||
|
||||
void completeCollectionLoad(CollectionLoadTicket ticket) {
|
||||
lifecycleLock.readLock().lock();
|
||||
try {
|
||||
requireOpenPool();
|
||||
if (ticket.generation != poolGeneration) {
|
||||
throw new IllegalStateException(
|
||||
"Milvus client pool changed while loading collection: "
|
||||
+ ticket.collectionName
|
||||
);
|
||||
}
|
||||
loadedCollections.add(ticket.collectionName);
|
||||
ticket.completion.complete(null);
|
||||
} finally {
|
||||
lifecycleLock.readLock().unlock();
|
||||
}
|
||||
}
|
||||
|
||||
void failCollectionLoad(
|
||||
CollectionLoadTicket ticket,
|
||||
Throwable failure,
|
||||
boolean retryableForFollowers
|
||||
) {
|
||||
if (retryableForFollowers && ticket.leader) {
|
||||
collectionLoads.remove(ticket.collectionName, ticket);
|
||||
}
|
||||
Throwable sharedFailure = retryableForFollowers
|
||||
? new RetryableCollectionLoadException(failure)
|
||||
: failure;
|
||||
ticket.completion.completeExceptionally(sharedFailure);
|
||||
}
|
||||
|
||||
void endCollectionLoad(CollectionLoadTicket ticket) {
|
||||
if (ticket.leader) {
|
||||
collectionLoads.remove(ticket.collectionName, ticket);
|
||||
}
|
||||
}
|
||||
|
||||
public int getActiveClientCount() {
|
||||
lifecycleLock.readLock().lock();
|
||||
try {
|
||||
return requireOpenPool().getTotalActiveClientNumber();
|
||||
} finally {
|
||||
lifecycleLock.readLock().unlock();
|
||||
}
|
||||
}
|
||||
|
||||
public int getIdleClientCount() {
|
||||
lifecycleLock.readLock().lock();
|
||||
try {
|
||||
return requireOpenPool().getTotalIdleClientNumber();
|
||||
} finally {
|
||||
lifecycleLock.readLock().unlock();
|
||||
}
|
||||
}
|
||||
|
||||
@Override
|
||||
public synchronized void close() {
|
||||
if (closed) {
|
||||
return;
|
||||
}
|
||||
acceptingOperations = false;
|
||||
closed = true;
|
||||
cancelActiveContexts("Milvus client pool is closing");
|
||||
interruptActiveOperations();
|
||||
lifecycleLock.writeLock().lock();
|
||||
try {
|
||||
initializedCollections.clear();
|
||||
loadedCollections.clear();
|
||||
poolGeneration++;
|
||||
failCollectionLoads("Milvus client pool was closed");
|
||||
ManagedMilvusClientV2Pool currentPool = pool;
|
||||
pool = null;
|
||||
poolFingerprint = null;
|
||||
if (currentPool != null) {
|
||||
currentPool.close();
|
||||
}
|
||||
} finally {
|
||||
lifecycleLock.writeLock().unlock();
|
||||
}
|
||||
}
|
||||
|
||||
private ManagedMilvusClientV2Pool requireOpenPool() {
|
||||
ManagedMilvusClientV2Pool currentPool = pool;
|
||||
if (closed || currentPool == null) {
|
||||
throw new IllegalStateException("Milvus client pool is closed");
|
||||
}
|
||||
return currentPool;
|
||||
}
|
||||
|
||||
boolean isClosed() {
|
||||
return closed;
|
||||
}
|
||||
|
||||
private void ensureAcceptingOperations() {
|
||||
if (!acceptingOperations) {
|
||||
throw new IllegalStateException(closed
|
||||
? "Milvus client pool is closed"
|
||||
: "Milvus client pool is reconfiguring");
|
||||
}
|
||||
}
|
||||
|
||||
private void failCollectionLoads(String message) {
|
||||
IllegalStateException failure = new IllegalStateException(message);
|
||||
for (CollectionLoadTicket ticket : collectionLoads.values()) {
|
||||
ticket.completion.completeExceptionally(failure);
|
||||
}
|
||||
collectionLoads.clear();
|
||||
}
|
||||
|
||||
private void cancelActiveContexts(String message) {
|
||||
for (Context.CancellableContext context : activeContexts) {
|
||||
context.cancel(new CancellationException(message));
|
||||
}
|
||||
}
|
||||
|
||||
private void interruptActiveOperations() {
|
||||
Thread currentThread = Thread.currentThread();
|
||||
for (Thread operationThread : activeOperations.keySet()) {
|
||||
if (operationThread != currentThread) {
|
||||
operationThread.interrupt();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private static String fingerprint(PoolSettings settings) {
|
||||
String value = String.join("\u0000",
|
||||
String.valueOf(settings.uri()),
|
||||
String.valueOf(settings.databaseName()),
|
||||
String.valueOf(settings.token()),
|
||||
String.valueOf(settings.username()),
|
||||
String.valueOf(settings.password()),
|
||||
String.valueOf(settings.poolMaxTotal()),
|
||||
String.valueOf(settings.poolMaxTotalPerKey()),
|
||||
String.valueOf(settings.poolMaxIdlePerKey()),
|
||||
String.valueOf(settings.poolMinIdlePerKey()),
|
||||
String.valueOf(settings.poolMaxWaitMillis()),
|
||||
String.valueOf(settings.poolEvictionIntervalMillis()),
|
||||
String.valueOf(settings.poolMinEvictableIdleMillis())
|
||||
);
|
||||
try {
|
||||
byte[] digest = MessageDigest.getInstance("SHA-256")
|
||||
.digest(value.getBytes(StandardCharsets.UTF_8));
|
||||
StringBuilder result = new StringBuilder(digest.length * 2);
|
||||
for (byte item : digest) {
|
||||
result.append(String.format("%02x", item & 0xff));
|
||||
}
|
||||
return result.toString();
|
||||
} catch (NoSuchAlgorithmException exception) {
|
||||
throw new IllegalStateException("SHA-256 is unavailable", exception);
|
||||
}
|
||||
}
|
||||
|
||||
static ConnectConfig buildConnectConfig(MilvusVectorStoreConfig config) {
|
||||
return buildConnectConfig(PoolSettings.from(config));
|
||||
}
|
||||
|
||||
private static ConnectConfig buildConnectConfig(PoolSettings settings) {
|
||||
String uri = normalizeAndValidateUri(settings.uri());
|
||||
String databaseName = StringUtil.hasText(settings.databaseName())
|
||||
? settings.databaseName().trim()
|
||||
: "default";
|
||||
ConnectConfig.ConnectConfigBuilder<?, ?> builder = ConnectConfig.builder()
|
||||
.uri(uri)
|
||||
.dbName(databaseName);
|
||||
if (StringUtil.hasText(settings.token())) {
|
||||
builder.token(settings.token().trim());
|
||||
}
|
||||
if (StringUtil.hasText(settings.username()) && StringUtil.hasText(settings.password())) {
|
||||
builder.username(settings.username().trim());
|
||||
builder.password(settings.password().trim());
|
||||
}
|
||||
return builder.build();
|
||||
}
|
||||
|
||||
private record PoolSettings(
|
||||
String uri,
|
||||
String databaseName,
|
||||
String token,
|
||||
String username,
|
||||
String password,
|
||||
int poolMaxTotal,
|
||||
int poolMaxTotalPerKey,
|
||||
int poolMaxIdlePerKey,
|
||||
int poolMinIdlePerKey,
|
||||
long poolMaxWaitMillis,
|
||||
long poolEvictionIntervalMillis,
|
||||
long poolMinEvictableIdleMillis
|
||||
) {
|
||||
private static PoolSettings from(MilvusVectorStoreConfig config) {
|
||||
return new PoolSettings(
|
||||
config.getUri(),
|
||||
config.getDatabaseName(),
|
||||
config.getToken(),
|
||||
config.getUsername(),
|
||||
config.getPassword(),
|
||||
config.getPoolMaxTotal(),
|
||||
config.getPoolMaxTotalPerKey(),
|
||||
config.getPoolMaxIdlePerKey(),
|
||||
config.getPoolMinIdlePerKey(),
|
||||
config.getPoolMaxWaitMillis(),
|
||||
config.getPoolEvictionIntervalMillis(),
|
||||
config.getPoolMinEvictableIdleMillis()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
private static final class ManagedMilvusClientV2Pool extends MilvusClientV2Pool {
|
||||
|
||||
private ManagedMilvusClientV2Pool(
|
||||
PoolConfig poolConfig,
|
||||
ConnectConfig connectConfig
|
||||
) throws ClassNotFoundException, NoSuchMethodException {
|
||||
super(poolConfig, connectConfig);
|
||||
}
|
||||
|
||||
private void invalidateClient(String key, MilvusClientV2 client) {
|
||||
try {
|
||||
clientPool.invalidateObject(key, client);
|
||||
} catch (Exception exception) {
|
||||
throw new IllegalStateException("Unable to invalidate Milvus client", exception);
|
||||
}
|
||||
}
|
||||
|
||||
private MilvusClientV2 getClient(String key, Duration maxWait) {
|
||||
if (maxWait == null || maxWait.isZero() || maxWait.isNegative()) {
|
||||
throw new IllegalArgumentException("maxWait must be greater than zero");
|
||||
}
|
||||
try {
|
||||
long waitMillis = Math.max(1L, maxWait.toMillis());
|
||||
return clientPool.borrowObject(key, waitMillis);
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
throw new IllegalStateException(
|
||||
"Interrupted while waiting for a Milvus client", exception);
|
||||
} catch (Exception exception) {
|
||||
throw new IllegalStateException(
|
||||
"Unable to borrow a Milvus client", exception);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
static final class CollectionLoadTicket {
|
||||
|
||||
private final String collectionName;
|
||||
private final long generation;
|
||||
private final CompletableFuture<Void> completion;
|
||||
private final boolean leader;
|
||||
|
||||
private CollectionLoadTicket(
|
||||
String collectionName,
|
||||
long generation,
|
||||
CompletableFuture<Void> completion,
|
||||
boolean leader
|
||||
) {
|
||||
this.collectionName = collectionName;
|
||||
this.generation = generation;
|
||||
this.completion = completion;
|
||||
this.leader = leader;
|
||||
}
|
||||
|
||||
boolean isLeader() {
|
||||
return leader;
|
||||
}
|
||||
|
||||
CompletableFuture<Void> completion() {
|
||||
return completion;
|
||||
}
|
||||
|
||||
private CollectionLoadTicket asFollower() {
|
||||
return new CollectionLoadTicket(
|
||||
collectionName, generation, completion, false);
|
||||
}
|
||||
}
|
||||
|
||||
static final class RetryableCollectionLoadException
|
||||
extends RuntimeException {
|
||||
|
||||
private RetryableCollectionLoadException(Throwable cause) {
|
||||
super("The collection load leader exhausted its local budget", cause);
|
||||
}
|
||||
}
|
||||
|
||||
static String normalizeAndValidateUri(String uri) {
|
||||
if (StringUtil.noText(uri)) {
|
||||
throw new IllegalArgumentException(
|
||||
"Milvus uri is required. Example: http://127.0.0.1:19530"
|
||||
);
|
||||
}
|
||||
String normalized = uri.trim();
|
||||
if (!normalized.contains("://")) {
|
||||
normalized = "http://" + normalized;
|
||||
}
|
||||
try {
|
||||
URI parsed = URI.create(normalized);
|
||||
if (StringUtil.noText(parsed.getHost()) || parsed.getPort() <= 0) {
|
||||
throw new IllegalArgumentException("Invalid Milvus uri: " + uri);
|
||||
}
|
||||
} catch (IllegalArgumentException exception) {
|
||||
throw new IllegalArgumentException(
|
||||
"Invalid Milvus uri: " + uri + ". Example: http://127.0.0.1:19530",
|
||||
exception
|
||||
);
|
||||
}
|
||||
return normalized;
|
||||
}
|
||||
}
|
||||
@@ -15,34 +15,45 @@
|
||||
*/
|
||||
package com.easyagents.store.milvus;
|
||||
|
||||
import com.alibaba.fastjson.JSON;
|
||||
import com.alibaba.fastjson.JSONObject;
|
||||
import com.google.gson.Gson;
|
||||
import com.google.gson.JsonObject;
|
||||
import com.easyagents.core.document.Document;
|
||||
import com.easyagents.core.store.DocumentStore;
|
||||
import com.easyagents.core.store.SearchWrapper;
|
||||
import com.easyagents.core.store.StoreOptions;
|
||||
import com.easyagents.core.store.StoreResult;
|
||||
import com.easyagents.core.store.StoreTimeoutException;
|
||||
import com.easyagents.core.util.CollectionUtil;
|
||||
import com.easyagents.core.util.Maps;
|
||||
import com.easyagents.core.util.StringUtil;
|
||||
import io.milvus.v2.client.ConnectConfig;
|
||||
import io.grpc.Context;
|
||||
import io.grpc.Status;
|
||||
import io.milvus.v2.client.MilvusClientV2;
|
||||
import io.milvus.v2.common.ConsistencyLevel;
|
||||
import io.milvus.v2.common.DataType;
|
||||
import io.milvus.v2.common.IndexParam;
|
||||
import io.milvus.v2.exception.MilvusClientException;
|
||||
import io.milvus.v2.service.collection.request.CreateCollectionReq;
|
||||
import io.milvus.v2.service.collection.request.GetLoadStateReq;
|
||||
import io.milvus.v2.service.collection.request.HasCollectionReq;
|
||||
import io.milvus.v2.service.collection.request.LoadCollectionReq;
|
||||
import io.milvus.v2.service.vector.request.*;
|
||||
import io.milvus.v2.service.vector.request.data.FloatVec;
|
||||
import io.milvus.v2.service.vector.response.QueryResp;
|
||||
import io.milvus.v2.service.vector.response.SearchResp;
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
|
||||
import java.net.URI;
|
||||
import java.time.Duration;
|
||||
import java.util.*;
|
||||
import java.util.concurrent.CancellationException;
|
||||
import java.util.concurrent.CompletableFuture;
|
||||
import java.util.concurrent.ExecutionException;
|
||||
import java.util.concurrent.Executors;
|
||||
import java.util.concurrent.ScheduledExecutorService;
|
||||
import java.util.concurrent.ThreadFactory;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.concurrent.TimeoutException;
|
||||
import java.util.function.Function;
|
||||
|
||||
/**
|
||||
* Milvus vector store based on Milvus Java SDK v2.
|
||||
@@ -52,63 +63,74 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
private static final Logger LOG = LoggerFactory.getLogger(MilvusVectorStore.class);
|
||||
private static final long LOAD_TIMEOUT_MS = 30_000L;
|
||||
private static final long LOAD_POLL_INTERVAL_MS = 200L;
|
||||
private static final long DEADLINE_SAFETY_MARGIN_MS = 200L;
|
||||
private static final ScheduledExecutorService DEADLINE_SCHEDULER =
|
||||
Executors.newSingleThreadScheduledExecutor(new ThreadFactory() {
|
||||
@Override
|
||||
public Thread newThread(Runnable runnable) {
|
||||
Thread thread = new Thread(runnable, "milvus-deadline");
|
||||
thread.setDaemon(true);
|
||||
return thread;
|
||||
}
|
||||
});
|
||||
|
||||
private static final String FIELD_ID = "id";
|
||||
private static final String FIELD_CONTENT = "content";
|
||||
private static final String FIELD_METADATA = "metadata";
|
||||
private static final String FIELD_VECTOR = "vector";
|
||||
|
||||
private final MilvusClientV2 client;
|
||||
private static final Gson GSON = new Gson();
|
||||
|
||||
private final MilvusClientManager clientManager;
|
||||
private final MilvusVectorStoreConfig config;
|
||||
private final String defaultCollectionName;
|
||||
private final Set<String> initializedCollections = Collections.synchronizedSet(new HashSet<String>());
|
||||
private final Set<String> loadedCollections = Collections.synchronizedSet(new HashSet<String>());
|
||||
private final boolean ownsClientManager;
|
||||
private volatile MilvusClientV2 compatibilityClient;
|
||||
private volatile boolean closed;
|
||||
|
||||
public MilvusVectorStore(MilvusVectorStoreConfig config) {
|
||||
this.config = config;
|
||||
this.defaultCollectionName = config.getDefaultCollectionName();
|
||||
String uri = normalizeAndValidateUri(config.getUri());
|
||||
String dbName = StringUtil.hasText(config.getDatabaseName()) ? config.getDatabaseName().trim() : "default";
|
||||
|
||||
ConnectConfig.ConnectConfigBuilder<?, ?> builder = ConnectConfig.builder()
|
||||
.uri(uri)
|
||||
.dbName(dbName);
|
||||
|
||||
if (StringUtil.hasText(config.getToken())) {
|
||||
builder.token(config.getToken().trim());
|
||||
}
|
||||
|
||||
if (StringUtil.hasText(config.getUsername()) && StringUtil.hasText(config.getPassword())) {
|
||||
builder.username(config.getUsername().trim());
|
||||
builder.password(config.getPassword().trim());
|
||||
}
|
||||
|
||||
ConnectConfig connectConfig = builder.build();
|
||||
this.client = new MilvusClientV2(connectConfig);
|
||||
this(config, createOwnedClientManager(config), true);
|
||||
}
|
||||
|
||||
private String normalizeAndValidateUri(String uri) {
|
||||
if (StringUtil.noText(uri)) {
|
||||
throw new IllegalArgumentException("Milvus uri is required. Example: http://127.0.0.1:19530");
|
||||
}
|
||||
public MilvusVectorStore(
|
||||
MilvusVectorStoreConfig config,
|
||||
MilvusClientManager clientManager
|
||||
) {
|
||||
this(config, clientManager, false);
|
||||
}
|
||||
|
||||
String normalized = uri.trim();
|
||||
if (!normalized.contains("://")) {
|
||||
normalized = "http://" + normalized;
|
||||
}
|
||||
private MilvusVectorStore(
|
||||
MilvusVectorStoreConfig config,
|
||||
MilvusClientManager clientManager,
|
||||
boolean ownsClientManager
|
||||
) {
|
||||
validateConfig(config);
|
||||
this.config = config;
|
||||
this.defaultCollectionName = config.getDefaultCollectionName();
|
||||
this.clientManager = Objects.requireNonNull(clientManager, "clientManager");
|
||||
this.ownsClientManager = ownsClientManager;
|
||||
}
|
||||
|
||||
URI parsed;
|
||||
try {
|
||||
parsed = URI.create(normalized);
|
||||
} catch (Exception e) {
|
||||
throw new IllegalArgumentException("Invalid Milvus uri: " + uri + ". Example: http://127.0.0.1:19530", e);
|
||||
}
|
||||
private static MilvusClientManager createOwnedClientManager(
|
||||
MilvusVectorStoreConfig config
|
||||
) {
|
||||
validateConfig(config);
|
||||
return new MilvusClientManager(config);
|
||||
}
|
||||
|
||||
if (StringUtil.noText(parsed.getHost()) || parsed.getPort() <= 0) {
|
||||
throw new IllegalArgumentException("Invalid Milvus uri: " + uri + ". Example: http://127.0.0.1:19530");
|
||||
private static void validateConfig(MilvusVectorStoreConfig config) {
|
||||
Objects.requireNonNull(config, "config");
|
||||
if (config.getSearchTimeoutMillis() <= DEADLINE_SAFETY_MARGIN_MS) {
|
||||
throw new IllegalArgumentException(
|
||||
"Milvus searchTimeoutMillis must be greater than "
|
||||
+ DEADLINE_SAFETY_MARGIN_MS
|
||||
);
|
||||
}
|
||||
if (config.getPoolMaxWaitMillis() <= 0L) {
|
||||
throw new IllegalArgumentException(
|
||||
"Milvus poolMaxWaitMillis must be greater than zero"
|
||||
);
|
||||
}
|
||||
|
||||
return normalized;
|
||||
}
|
||||
|
||||
@Override
|
||||
@@ -121,22 +143,25 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
throw new IllegalStateException("CollectionName is null or blank. please config the \"defaultCollectionName\" or store with designative collectionName.");
|
||||
}
|
||||
|
||||
int dimension = getDimension(documents);
|
||||
ensureCollectionExists(collectionName, dimension);
|
||||
|
||||
try {
|
||||
InsertReq.InsertReqBuilder<?, ?> builder = InsertReq.builder();
|
||||
if (StringUtil.hasText(options.getPartitionName())) {
|
||||
builder.partitionName(options.getPartitionName());
|
||||
}
|
||||
InsertReq insertReq = builder
|
||||
.collectionName(collectionName)
|
||||
.data(toMilvusDocuments(documents))
|
||||
.build();
|
||||
client.insert(insertReq);
|
||||
int dimension = getDimension(documents);
|
||||
clientManager.withClient(client -> {
|
||||
ensureCollectionExists(client, collectionName, dimension);
|
||||
InsertReq.InsertReqBuilder<?, ?> builder = InsertReq.builder();
|
||||
if (StringUtil.hasText(options.getPartitionName())) {
|
||||
builder.partitionName(options.getPartitionName());
|
||||
}
|
||||
client.insert(builder
|
||||
.collectionName(collectionName)
|
||||
.data(toMilvusDocuments(documents))
|
||||
.build());
|
||||
return null;
|
||||
});
|
||||
return StoreResult.successWithIds(documents);
|
||||
} catch (MilvusClientException e) {
|
||||
return StoreResult.fail();
|
||||
} catch (RuntimeException e) {
|
||||
LOG.error("Milvus insert failed. collection={}, message={}",
|
||||
collectionName, e.getMessage(), e);
|
||||
return StoreResult.fail(e.getMessage());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -159,7 +184,10 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
.collectionName(collectionName)
|
||||
.ids(MilvusPrimaryKeySupport.normalize(ids))
|
||||
.build();
|
||||
client.delete(deleteReq);
|
||||
clientManager.withClient(client -> {
|
||||
client.delete(deleteReq);
|
||||
return null;
|
||||
});
|
||||
return StoreResult.success();
|
||||
} catch (Exception e) {
|
||||
LOG.error("Milvus delete failed. collection={}, message={}",
|
||||
@@ -178,19 +206,22 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
throw new IllegalStateException("CollectionName is null or blank. please config the \"defaultCollectionName\" or store with designative collectionName.");
|
||||
}
|
||||
|
||||
int dimension = getDimension(documents);
|
||||
ensureCollectionExists(collectionName, dimension);
|
||||
|
||||
try {
|
||||
UpsertReq upsertReq = UpsertReq.builder()
|
||||
.collectionName(collectionName)
|
||||
.partitionName(options.getPartitionName())
|
||||
.data(toMilvusDocuments(documents))
|
||||
.build();
|
||||
client.upsert(upsertReq);
|
||||
int dimension = getDimension(documents);
|
||||
clientManager.withClient(client -> {
|
||||
ensureCollectionExists(client, collectionName, dimension);
|
||||
client.upsert(UpsertReq.builder()
|
||||
.collectionName(collectionName)
|
||||
.partitionName(options.getPartitionName())
|
||||
.data(toMilvusDocuments(documents))
|
||||
.build());
|
||||
return null;
|
||||
});
|
||||
return StoreResult.successWithIds(documents);
|
||||
} catch (Exception e) {
|
||||
return StoreResult.fail();
|
||||
LOG.error("Milvus upsert failed. collection={}, message={}",
|
||||
collectionName, e.getMessage(), e);
|
||||
return StoreResult.fail(e.getMessage());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -200,56 +231,104 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
if (StringUtil.noText(collectionName)) {
|
||||
throw new IllegalStateException("CollectionName is null or blank. please config the \"defaultCollectionName\" or store with designative collectionName.");
|
||||
}
|
||||
ensureCollectionLoaded(collectionName);
|
||||
long timeoutMillis = resolveSearchTimeoutMillis(options);
|
||||
long rpcBudgetMillis = timeoutMillis - DEADLINE_SAFETY_MARGIN_MS;
|
||||
long deadlineNanos = deadlineAfterMillis(rpcBudgetMillis);
|
||||
Context.CancellableContext context = Context.current().withDeadlineAfter(
|
||||
rpcBudgetMillis, TimeUnit.MILLISECONDS, DEADLINE_SCHEDULER);
|
||||
try {
|
||||
return clientManager.withRequestContext(context, () ->
|
||||
searchWithinDeadline(
|
||||
wrapper, options, collectionName, deadlineNanos));
|
||||
} catch (RuntimeException exception) {
|
||||
if (!(exception instanceof StoreTimeoutException)
|
||||
&& deadlineExpired(deadlineNanos, exception)) {
|
||||
throw timeoutException(collectionName, exception);
|
||||
}
|
||||
throw exception;
|
||||
} catch (Exception exception) {
|
||||
throw new IllegalStateException("Milvus search failed", exception);
|
||||
} finally {
|
||||
context.cancel(null);
|
||||
}
|
||||
}
|
||||
|
||||
private List<Document> searchWithinDeadline(
|
||||
SearchWrapper wrapper,
|
||||
StoreOptions options,
|
||||
String collectionName,
|
||||
long deadlineNanos
|
||||
) {
|
||||
String operation = wrapper.getVector() == null
|
||||
|| wrapper.getVector().length == 0
|
||||
? "query"
|
||||
: "search";
|
||||
ensureCollectionLoaded(collectionName, deadlineNanos);
|
||||
try {
|
||||
return searchOnce(wrapper, options, collectionName, deadlineNanos);
|
||||
} catch (RuntimeException exception) {
|
||||
if (!isCollectionNotLoaded(exception)) {
|
||||
throw propagateSearchFailure(
|
||||
operation, collectionName, exception);
|
||||
}
|
||||
clientManager.markCollectionUnloaded(collectionName);
|
||||
try {
|
||||
ensureCollectionLoaded(collectionName, deadlineNanos);
|
||||
return searchOnce(
|
||||
wrapper, options, collectionName, deadlineNanos);
|
||||
} catch (RuntimeException retryException) {
|
||||
retryException.addSuppressed(exception);
|
||||
throw propagateSearchFailure(
|
||||
operation, collectionName, retryException);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private List<Document> searchOnce(
|
||||
SearchWrapper wrapper,
|
||||
StoreOptions options,
|
||||
String collectionName,
|
||||
long deadlineNanos
|
||||
) {
|
||||
if (wrapper.getVector() == null || wrapper.getVector().length == 0) {
|
||||
return queryByCondition(wrapper, options, collectionName);
|
||||
return queryByCondition(
|
||||
wrapper, options, collectionName, deadlineNanos);
|
||||
}
|
||||
return searchByVector(wrapper, options, collectionName);
|
||||
return searchByVector(wrapper, options, collectionName, deadlineNanos);
|
||||
}
|
||||
|
||||
private List<Document> searchByVector(SearchWrapper wrapper, StoreOptions options, String collectionName) {
|
||||
private List<Document> searchByVector(
|
||||
SearchWrapper wrapper,
|
||||
StoreOptions options,
|
||||
String collectionName,
|
||||
long deadlineNanos
|
||||
) {
|
||||
SearchReq searchReq = buildSearchReq(wrapper, options, collectionName);
|
||||
try {
|
||||
SearchResp resp = client.search(searchReq);
|
||||
return parseSearchResults(resp, wrapper.getMinScore());
|
||||
} catch (Exception e) {
|
||||
if (isCollectionNotLoaded(e)) {
|
||||
loadedCollections.remove(collectionName);
|
||||
try {
|
||||
ensureCollectionLoaded(collectionName);
|
||||
SearchResp retryResp = client.search(searchReq);
|
||||
return parseSearchResults(retryResp, wrapper.getMinScore());
|
||||
} catch (Exception retryException) {
|
||||
LOG.warn("Milvus search retry failed after load. collection={}, message={}", collectionName, retryException.getMessage());
|
||||
return Collections.emptyList();
|
||||
}
|
||||
}
|
||||
LOG.warn("Milvus search failed. collection={}, message={}", collectionName, e.getMessage());
|
||||
return Collections.emptyList();
|
||||
}
|
||||
SearchResp resp = withClientBeforeDeadline(
|
||||
deadlineNanos, client -> client.search(searchReq));
|
||||
return parseSearchResults(resp, wrapper.getMinScore());
|
||||
}
|
||||
|
||||
private List<Document> queryByCondition(SearchWrapper wrapper, StoreOptions options, String collectionName) {
|
||||
private List<Document> queryByCondition(
|
||||
SearchWrapper wrapper,
|
||||
StoreOptions options,
|
||||
String collectionName,
|
||||
long deadlineNanos
|
||||
) {
|
||||
QueryReq queryReq = buildQueryReq(wrapper, options, collectionName);
|
||||
try {
|
||||
QueryResp resp = client.query(queryReq);
|
||||
return parseQueryResults(resp);
|
||||
} catch (Exception e) {
|
||||
if (isCollectionNotLoaded(e)) {
|
||||
loadedCollections.remove(collectionName);
|
||||
try {
|
||||
ensureCollectionLoaded(collectionName);
|
||||
QueryResp retryResp = client.query(queryReq);
|
||||
return parseQueryResults(retryResp);
|
||||
} catch (Exception retryException) {
|
||||
LOG.warn("Milvus query retry failed after load. collection={}, message={}", collectionName, retryException.getMessage());
|
||||
return Collections.emptyList();
|
||||
}
|
||||
}
|
||||
LOG.warn("Milvus query failed. collection={}, message={}", collectionName, e.getMessage());
|
||||
return Collections.emptyList();
|
||||
}
|
||||
QueryResp resp = withClientBeforeDeadline(
|
||||
deadlineNanos, client -> client.query(queryReq));
|
||||
return parseQueryResults(resp);
|
||||
}
|
||||
|
||||
private RuntimeException propagateSearchFailure(
|
||||
String operation,
|
||||
String collectionName,
|
||||
RuntimeException exception
|
||||
) {
|
||||
LOG.error("Milvus {} failed. collection={}, message={}",
|
||||
operation, collectionName, exception.getMessage(), exception);
|
||||
return exception;
|
||||
}
|
||||
|
||||
private SearchReq buildSearchReq(SearchWrapper wrapper, StoreOptions options, String collectionName) {
|
||||
@@ -259,7 +338,7 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
.outputFields(getOutputFields(wrapper))
|
||||
.topK(wrapper.getMaxResults())
|
||||
.annsField(FIELD_VECTOR)
|
||||
.data(Collections.singletonList(toFloatList(wrapper.getVector())))
|
||||
.data(Collections.singletonList(new FloatVec(wrapper.getVector())))
|
||||
.searchParams(Maps.of("ef", 64));
|
||||
|
||||
if (CollectionUtil.hasItems(options.getPartitionNamesOrEmpty())) {
|
||||
@@ -305,11 +384,7 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
continue;
|
||||
}
|
||||
document.setId(result.getId());
|
||||
Float distance = result.getDistance();
|
||||
if (distance != null) {
|
||||
double score = (distance + 1.0d) / 2.0d;
|
||||
document.setScore(score);
|
||||
}
|
||||
document.setScore(normalizeScore(result.getScore()));
|
||||
if (minScore == null || document.getScore() == null || document.getScore() >= minScore) {
|
||||
documents.add(document);
|
||||
}
|
||||
@@ -318,6 +393,10 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
return documents;
|
||||
}
|
||||
|
||||
static Double normalizeScore(Float rawScore) {
|
||||
return rawScore == null ? null : (rawScore + 1.0d) / 2.0d;
|
||||
}
|
||||
|
||||
private List<Document> parseQueryResults(QueryResp resp) {
|
||||
List<QueryResp.QueryResult> results = resp.getQueryResults();
|
||||
if (CollectionUtil.noItems(results)) {
|
||||
@@ -360,22 +439,28 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
document.addMetadata(metadata);
|
||||
} else if (metadataObj != null) {
|
||||
@SuppressWarnings("unchecked")
|
||||
Map<String, Object> metadata = JSON.parseObject(JSON.toJSONString(metadataObj), Map.class);
|
||||
Map<String, Object> metadata = GSON.fromJson(
|
||||
GSON.toJsonTree(metadataObj),
|
||||
Map.class
|
||||
);
|
||||
document.addMetadata(metadata);
|
||||
}
|
||||
|
||||
return document;
|
||||
}
|
||||
|
||||
private List<JSONObject> toMilvusDocuments(List<Document> documents) {
|
||||
List<JSONObject> rows = new ArrayList<JSONObject>(documents.size());
|
||||
List<JsonObject> toMilvusDocuments(List<Document> documents) {
|
||||
List<JsonObject> rows = new ArrayList<JsonObject>(documents.size());
|
||||
for (Document doc : documents) {
|
||||
JSONObject row = new JSONObject();
|
||||
row.put(FIELD_ID, String.valueOf(doc.getId()));
|
||||
row.put(FIELD_CONTENT, doc.getContent());
|
||||
row.put(FIELD_VECTOR, toFloatList(doc.getVector()));
|
||||
JsonObject row = new JsonObject();
|
||||
row.addProperty(FIELD_ID, String.valueOf(doc.getId()));
|
||||
row.addProperty(FIELD_CONTENT, doc.getContent());
|
||||
row.add(FIELD_VECTOR, GSON.toJsonTree(toFloatList(doc.getVector())));
|
||||
Map<String, Object> metadatas = doc.getMetadataMap();
|
||||
row.put(FIELD_METADATA, metadatas == null ? new JSONObject() : new JSONObject(metadatas));
|
||||
row.add(
|
||||
FIELD_METADATA,
|
||||
metadatas == null ? new JsonObject() : GSON.toJsonTree(metadatas)
|
||||
);
|
||||
rows.add(row);
|
||||
}
|
||||
return rows;
|
||||
@@ -413,67 +498,175 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
throw new IllegalStateException("Unable to determine vector dimension for Milvus collection.");
|
||||
}
|
||||
|
||||
private void ensureCollectionExists(String collectionName, int dimension) {
|
||||
if (initializedCollections.contains(collectionName)) {
|
||||
private void ensureCollectionExists(
|
||||
MilvusClientV2 client,
|
||||
String collectionName,
|
||||
int dimension
|
||||
) {
|
||||
if (clientManager.isCollectionInitialized(collectionName)) {
|
||||
return;
|
||||
}
|
||||
synchronized (initializedCollections) {
|
||||
if (initializedCollections.contains(collectionName)) {
|
||||
synchronized (clientManager.initializedCollectionsLock()) {
|
||||
if (clientManager.isCollectionInitialized(collectionName)) {
|
||||
return;
|
||||
}
|
||||
Boolean exists = client.hasCollection(HasCollectionReq.builder().collectionName(collectionName).build());
|
||||
if (Boolean.TRUE.equals(exists)) {
|
||||
initializedCollections.add(collectionName);
|
||||
clientManager.markCollectionInitialized(collectionName);
|
||||
return;
|
||||
}
|
||||
if (!config.isAutoCreateCollection()) {
|
||||
throw new IllegalStateException("Milvus collection not found and autoCreateCollection is disabled: " + collectionName);
|
||||
}
|
||||
createCollection(collectionName, dimension);
|
||||
initializedCollections.add(collectionName);
|
||||
createCollection(client, collectionName, dimension);
|
||||
clientManager.markCollectionInitialized(collectionName);
|
||||
}
|
||||
}
|
||||
|
||||
private void ensureCollectionLoaded(String collectionName) {
|
||||
if (loadedCollections.contains(collectionName)) {
|
||||
return;
|
||||
}
|
||||
synchronized (loadedCollections) {
|
||||
if (loadedCollections.contains(collectionName)) {
|
||||
return;
|
||||
private void ensureCollectionLoaded(
|
||||
String collectionName,
|
||||
long deadlineNanos
|
||||
) {
|
||||
while (!clientManager.isCollectionLoaded(collectionName)) {
|
||||
MilvusClientManager.CollectionLoadTicket ticket =
|
||||
clientManager.beginCollectionLoad(collectionName);
|
||||
if (!ticket.isLeader()) {
|
||||
if (awaitCollectionLoad(
|
||||
ticket, collectionName, deadlineNanos)) {
|
||||
return;
|
||||
}
|
||||
continue;
|
||||
}
|
||||
boolean loaded = false;
|
||||
try {
|
||||
loaded = Boolean.TRUE.equals(client.getLoadState(GetLoadStateReq.builder().collectionName(collectionName).build()));
|
||||
} catch (Exception e) {
|
||||
LOG.warn("Milvus getLoadState failed. collection={}, message={}", collectionName, e.getMessage());
|
||||
if (clientManager.isCollectionLoaded(collectionName)) {
|
||||
clientManager.completeCollectionLoad(ticket);
|
||||
return;
|
||||
}
|
||||
withClientBeforeDeadline(deadlineNanos, client -> {
|
||||
boolean loaded = Boolean.TRUE.equals(client.getLoadState(
|
||||
GetLoadStateReq.builder()
|
||||
.collectionName(collectionName)
|
||||
.build()
|
||||
));
|
||||
if (!loaded) {
|
||||
client.loadCollection(LoadCollectionReq.builder()
|
||||
.collectionName(collectionName)
|
||||
.async(false)
|
||||
.build());
|
||||
waitForCollectionLoaded(
|
||||
client, collectionName, deadlineNanos);
|
||||
}
|
||||
return null;
|
||||
});
|
||||
clientManager.completeCollectionLoad(ticket);
|
||||
return;
|
||||
} catch (RuntimeException | Error failure) {
|
||||
clientManager.failCollectionLoad(
|
||||
ticket, failure, isLeaderLocalAbort(failure));
|
||||
throw failure;
|
||||
} finally {
|
||||
clientManager.endCollectionLoad(ticket);
|
||||
}
|
||||
|
||||
if (!loaded) {
|
||||
client.loadCollection(LoadCollectionReq.builder().collectionName(collectionName).build());
|
||||
waitForCollectionLoaded(collectionName);
|
||||
}
|
||||
loadedCollections.add(collectionName);
|
||||
}
|
||||
}
|
||||
|
||||
private void waitForCollectionLoaded(String collectionName) {
|
||||
long deadline = System.currentTimeMillis() + LOAD_TIMEOUT_MS;
|
||||
while (System.currentTimeMillis() < deadline) {
|
||||
private boolean isLeaderLocalAbort(Throwable failure) {
|
||||
if (clientManager.isClosed()) {
|
||||
return false;
|
||||
}
|
||||
if (failure instanceof StoreTimeoutException
|
||||
|| Thread.currentThread().isInterrupted()
|
||||
|| Context.current().isCancelled()) {
|
||||
return true;
|
||||
}
|
||||
Throwable current = failure;
|
||||
while (current != null) {
|
||||
if (current instanceof InterruptedException) {
|
||||
return true;
|
||||
}
|
||||
current = current.getCause();
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
private boolean awaitCollectionLoad(
|
||||
MilvusClientManager.CollectionLoadTicket ticket,
|
||||
String collectionName,
|
||||
long deadlineNanos
|
||||
) {
|
||||
Context currentContext = Context.current();
|
||||
CompletableFuture<Void> cancelled = new CompletableFuture<Void>();
|
||||
Context.CancellationListener cancellationListener = context ->
|
||||
cancelled.completeExceptionally(new CancellationException(
|
||||
"Milvus search was cancelled"));
|
||||
currentContext.addListener(cancellationListener, Runnable::run);
|
||||
try {
|
||||
CompletableFuture.anyOf(ticket.completion(), cancelled).get(
|
||||
remainingNanos(deadlineNanos, collectionName),
|
||||
TimeUnit.NANOSECONDS
|
||||
);
|
||||
return true;
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
throw new IllegalStateException(
|
||||
"Interrupted while loading Milvus collection: "
|
||||
+ collectionName,
|
||||
exception
|
||||
);
|
||||
} catch (TimeoutException exception) {
|
||||
throw timeoutException(collectionName, exception);
|
||||
} catch (ExecutionException exception) {
|
||||
Throwable cause = exception.getCause();
|
||||
if (cause instanceof MilvusClientManager
|
||||
.RetryableCollectionLoadException) {
|
||||
remainingNanos(deadlineNanos, collectionName);
|
||||
return false;
|
||||
}
|
||||
if (cause instanceof RuntimeException runtimeException) {
|
||||
throw runtimeException;
|
||||
}
|
||||
if (cause instanceof Error error) {
|
||||
throw error;
|
||||
}
|
||||
throw new IllegalStateException(
|
||||
"Unable to load Milvus collection: " + collectionName,
|
||||
cause
|
||||
);
|
||||
} catch (CancellationException exception) {
|
||||
throw new IllegalStateException(
|
||||
"Milvus collection load was cancelled: " + collectionName,
|
||||
exception
|
||||
);
|
||||
} finally {
|
||||
currentContext.removeListener(cancellationListener);
|
||||
}
|
||||
}
|
||||
|
||||
private void waitForCollectionLoaded(
|
||||
MilvusClientV2 client,
|
||||
String collectionName,
|
||||
long deadlineNanos
|
||||
) {
|
||||
while (true) {
|
||||
long remainingNanos = remainingNanos(
|
||||
deadlineNanos, collectionName);
|
||||
if (Boolean.TRUE.equals(client.getLoadState(GetLoadStateReq.builder().collectionName(collectionName).build()))) {
|
||||
return;
|
||||
}
|
||||
try {
|
||||
Thread.sleep(LOAD_POLL_INTERVAL_MS);
|
||||
long sleepMillis = Math.min(
|
||||
LOAD_POLL_INTERVAL_MS,
|
||||
Math.max(1L, TimeUnit.NANOSECONDS.toMillis(remainingNanos))
|
||||
);
|
||||
Thread.sleep(sleepMillis);
|
||||
} catch (InterruptedException e) {
|
||||
Thread.currentThread().interrupt();
|
||||
throw new IllegalStateException("Interrupted while loading Milvus collection: " + collectionName, e);
|
||||
}
|
||||
}
|
||||
throw new IllegalStateException("Timeout waiting for Milvus collection loaded: " + collectionName);
|
||||
}
|
||||
|
||||
private boolean isCollectionNotLoaded(Exception e) {
|
||||
private boolean isCollectionNotLoaded(Throwable e) {
|
||||
Throwable current = e;
|
||||
while (current != null) {
|
||||
String message = current.getMessage();
|
||||
@@ -485,7 +678,91 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
return false;
|
||||
}
|
||||
|
||||
private void createCollection(String collectionName, int dimension) {
|
||||
private <T> T withClientBeforeDeadline(
|
||||
long deadlineNanos,
|
||||
Function<MilvusClientV2, T> operation
|
||||
) {
|
||||
long remainingNanos = remainingNanos(deadlineNanos, null);
|
||||
long poolWaitNanos = TimeUnit.MILLISECONDS.toNanos(
|
||||
config.getPoolMaxWaitMillis());
|
||||
try {
|
||||
return clientManager.withClient(
|
||||
Duration.ofNanos(Math.min(remainingNanos, poolWaitNanos)),
|
||||
operation
|
||||
);
|
||||
} catch (RuntimeException exception) {
|
||||
if (deadlineExpired(deadlineNanos, exception)) {
|
||||
throw timeoutException(null, exception);
|
||||
}
|
||||
throw exception;
|
||||
}
|
||||
}
|
||||
|
||||
private long resolveSearchTimeoutMillis(StoreOptions options) {
|
||||
long timeoutMillis = config.getSearchTimeoutMillis();
|
||||
Long requestedTimeoutMillis = options.getTimeoutMillis();
|
||||
if (requestedTimeoutMillis != null) {
|
||||
timeoutMillis = Math.min(timeoutMillis, requestedTimeoutMillis);
|
||||
}
|
||||
if (timeoutMillis <= DEADLINE_SAFETY_MARGIN_MS) {
|
||||
throw new StoreTimeoutException(
|
||||
"Insufficient time remaining for Milvus search"
|
||||
);
|
||||
}
|
||||
return timeoutMillis;
|
||||
}
|
||||
|
||||
private static long deadlineAfterMillis(long timeoutMillis) {
|
||||
long now = System.nanoTime();
|
||||
long timeoutNanos = TimeUnit.MILLISECONDS.toNanos(timeoutMillis);
|
||||
if (now > Long.MAX_VALUE - timeoutNanos) {
|
||||
return Long.MAX_VALUE;
|
||||
}
|
||||
return now + timeoutNanos;
|
||||
}
|
||||
|
||||
private static long remainingNanos(
|
||||
long deadlineNanos,
|
||||
String collectionName
|
||||
) {
|
||||
if (deadlineNanos == Long.MAX_VALUE) {
|
||||
return Long.MAX_VALUE;
|
||||
}
|
||||
long remaining = deadlineNanos - System.nanoTime();
|
||||
if (remaining <= 0L) {
|
||||
throw timeoutException(collectionName, null);
|
||||
}
|
||||
return remaining;
|
||||
}
|
||||
|
||||
private static boolean deadlineExpired(
|
||||
long deadlineNanos,
|
||||
Throwable failure
|
||||
) {
|
||||
if (deadlineNanos != Long.MAX_VALUE
|
||||
&& System.nanoTime() >= deadlineNanos) {
|
||||
return true;
|
||||
}
|
||||
Throwable cancellationCause = Context.current().cancellationCause();
|
||||
return cancellationCause instanceof TimeoutException
|
||||
|| Status.fromThrowable(failure).getCode()
|
||||
== Status.Code.DEADLINE_EXCEEDED;
|
||||
}
|
||||
|
||||
private static StoreTimeoutException timeoutException(
|
||||
String collectionName,
|
||||
Throwable cause
|
||||
) {
|
||||
String suffix = StringUtil.hasText(collectionName)
|
||||
? ": " + collectionName
|
||||
: "";
|
||||
return new StoreTimeoutException(
|
||||
"Timeout waiting for Milvus search" + suffix,
|
||||
cause
|
||||
);
|
||||
}
|
||||
|
||||
private void createCollection(MilvusClientV2 client, String collectionName, int dimension) {
|
||||
List<CreateCollectionReq.FieldSchema> fieldSchemaList = new ArrayList<CreateCollectionReq.FieldSchema>();
|
||||
fieldSchemaList.add(CreateCollectionReq.FieldSchema.builder()
|
||||
.name(FIELD_ID)
|
||||
@@ -531,31 +808,83 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
.indexParams(indexParams)
|
||||
.build();
|
||||
client.createCollection(createCollectionReq);
|
||||
ensureCollectionLoaded(collectionName);
|
||||
ensureCollectionLoadedForWrite(client, collectionName);
|
||||
}
|
||||
|
||||
public MilvusClientV2 getClient() {
|
||||
return client;
|
||||
private void ensureCollectionLoadedForWrite(
|
||||
MilvusClientV2 client,
|
||||
String collectionName
|
||||
) {
|
||||
if (clientManager.isCollectionLoaded(collectionName)) {
|
||||
return;
|
||||
}
|
||||
boolean loaded = Boolean.TRUE.equals(client.getLoadState(
|
||||
GetLoadStateReq.builder().collectionName(collectionName).build()));
|
||||
if (!loaded) {
|
||||
client.loadCollection(LoadCollectionReq.builder()
|
||||
.collectionName(collectionName)
|
||||
.async(false)
|
||||
.build());
|
||||
waitForCollectionLoaded(
|
||||
client,
|
||||
collectionName,
|
||||
deadlineAfterMillis(LOAD_TIMEOUT_MS)
|
||||
);
|
||||
}
|
||||
clientManager.markCollectionLoaded(collectionName);
|
||||
}
|
||||
|
||||
public boolean checkAvailable() {
|
||||
try {
|
||||
return client.hasCollection(HasCollectionReq.builder()
|
||||
.collectionName("__milvus_boot_probe__")
|
||||
.build()) != null;
|
||||
return clientManager.withClient(client -> client.hasCollection(
|
||||
HasCollectionReq.builder()
|
||||
.collectionName("__milvus_boot_probe__")
|
||||
.build()
|
||||
)) != null;
|
||||
} catch (Exception e) {
|
||||
LOG.warn("Milvus availability check failed. message={}", e.getMessage());
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Returns a compatibility client for integrations that used the pre-pool API.
|
||||
* Prefer store operations so pooled lifecycle management remains automatic.
|
||||
*/
|
||||
@Deprecated
|
||||
public synchronized MilvusClientV2 getClient() {
|
||||
if (closed) {
|
||||
throw new IllegalStateException("Milvus vector store is closed");
|
||||
}
|
||||
if (compatibilityClient == null) {
|
||||
compatibilityClient = new MilvusClientV2(
|
||||
MilvusClientManager.buildConnectConfig(config)
|
||||
);
|
||||
}
|
||||
return compatibilityClient;
|
||||
}
|
||||
|
||||
@Override
|
||||
public void close() {
|
||||
try {
|
||||
client.close(1L);
|
||||
} catch (InterruptedException e) {
|
||||
Thread.currentThread().interrupt();
|
||||
LOG.warn("Interrupted while closing Milvus client. uri={}", config.getUri(), e);
|
||||
MilvusClientV2 legacyClient;
|
||||
synchronized (this) {
|
||||
if (closed) {
|
||||
return;
|
||||
}
|
||||
closed = true;
|
||||
legacyClient = compatibilityClient;
|
||||
compatibilityClient = null;
|
||||
}
|
||||
if (legacyClient != null) {
|
||||
try {
|
||||
legacyClient.close(1L);
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
LOG.warn("Interrupted while closing compatibility Milvus client", exception);
|
||||
}
|
||||
}
|
||||
if (ownsClientManager) {
|
||||
clientManager.close();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -30,6 +30,14 @@ public class MilvusVectorStoreConfig implements DocumentStoreConfig {
|
||||
private String password;
|
||||
private String defaultCollectionName;
|
||||
private boolean autoCreateCollection = true;
|
||||
private int poolMaxTotal = 8;
|
||||
private int poolMaxTotalPerKey = 8;
|
||||
private int poolMaxIdlePerKey = 4;
|
||||
private int poolMinIdlePerKey = 1;
|
||||
private long poolMaxWaitMillis = 3_000L;
|
||||
private long poolEvictionIntervalMillis = 60_000L;
|
||||
private long poolMinEvictableIdleMillis = 300_000L;
|
||||
private long searchTimeoutMillis = 10_000L;
|
||||
|
||||
public String getUri() {
|
||||
return uri;
|
||||
@@ -87,6 +95,70 @@ public class MilvusVectorStoreConfig implements DocumentStoreConfig {
|
||||
this.autoCreateCollection = autoCreateCollection;
|
||||
}
|
||||
|
||||
public int getPoolMaxTotal() {
|
||||
return poolMaxTotal;
|
||||
}
|
||||
|
||||
public void setPoolMaxTotal(int poolMaxTotal) {
|
||||
this.poolMaxTotal = poolMaxTotal;
|
||||
}
|
||||
|
||||
public int getPoolMaxTotalPerKey() {
|
||||
return poolMaxTotalPerKey;
|
||||
}
|
||||
|
||||
public void setPoolMaxTotalPerKey(int poolMaxTotalPerKey) {
|
||||
this.poolMaxTotalPerKey = poolMaxTotalPerKey;
|
||||
}
|
||||
|
||||
public int getPoolMaxIdlePerKey() {
|
||||
return poolMaxIdlePerKey;
|
||||
}
|
||||
|
||||
public void setPoolMaxIdlePerKey(int poolMaxIdlePerKey) {
|
||||
this.poolMaxIdlePerKey = poolMaxIdlePerKey;
|
||||
}
|
||||
|
||||
public int getPoolMinIdlePerKey() {
|
||||
return poolMinIdlePerKey;
|
||||
}
|
||||
|
||||
public void setPoolMinIdlePerKey(int poolMinIdlePerKey) {
|
||||
this.poolMinIdlePerKey = poolMinIdlePerKey;
|
||||
}
|
||||
|
||||
public long getPoolMaxWaitMillis() {
|
||||
return poolMaxWaitMillis;
|
||||
}
|
||||
|
||||
public void setPoolMaxWaitMillis(long poolMaxWaitMillis) {
|
||||
this.poolMaxWaitMillis = poolMaxWaitMillis;
|
||||
}
|
||||
|
||||
public long getPoolEvictionIntervalMillis() {
|
||||
return poolEvictionIntervalMillis;
|
||||
}
|
||||
|
||||
public void setPoolEvictionIntervalMillis(long poolEvictionIntervalMillis) {
|
||||
this.poolEvictionIntervalMillis = poolEvictionIntervalMillis;
|
||||
}
|
||||
|
||||
public long getPoolMinEvictableIdleMillis() {
|
||||
return poolMinEvictableIdleMillis;
|
||||
}
|
||||
|
||||
public void setPoolMinEvictableIdleMillis(long poolMinEvictableIdleMillis) {
|
||||
this.poolMinEvictableIdleMillis = poolMinEvictableIdleMillis;
|
||||
}
|
||||
|
||||
public long getSearchTimeoutMillis() {
|
||||
return searchTimeoutMillis;
|
||||
}
|
||||
|
||||
public void setSearchTimeoutMillis(long searchTimeoutMillis) {
|
||||
this.searchTimeoutMillis = searchTimeoutMillis;
|
||||
}
|
||||
|
||||
@Override
|
||||
public boolean checkAvailable() {
|
||||
return StringUtil.hasText(this.uri);
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
package com.easyagents.store.milvus;
|
||||
|
||||
import com.easyagents.core.document.Document;
|
||||
import com.google.gson.JsonObject;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
/**
|
||||
* Milvus SDK 2.3.11 数据适配回归测试。
|
||||
*/
|
||||
public class MilvusVectorStoreCompatibilityTest {
|
||||
|
||||
@Test
|
||||
public void shouldConvertRowsToGsonWithoutLosingMetadata() {
|
||||
MilvusVectorStoreConfig config = new MilvusVectorStoreConfig();
|
||||
config.setUri("http://127.0.0.1:19530");
|
||||
config.setDefaultCollectionName("test");
|
||||
MilvusVectorStore store = new MilvusVectorStore(config);
|
||||
try {
|
||||
Document document = Document.of("正文");
|
||||
document.setId("chunk-1");
|
||||
document.setVector(new float[] { 0.25F, 0.75F });
|
||||
document.addMetadata(Map.of("knowledgeId", "knowledge-1"));
|
||||
|
||||
List<JsonObject> rows = store.toMilvusDocuments(List.of(document));
|
||||
|
||||
Assert.assertEquals(1, rows.size());
|
||||
Assert.assertEquals("chunk-1", rows.get(0).get("id").getAsString());
|
||||
Assert.assertEquals("正文", rows.get(0).get("content").getAsString());
|
||||
Assert.assertEquals(2, rows.get(0).getAsJsonArray("vector").size());
|
||||
Assert.assertEquals(
|
||||
"knowledge-1",
|
||||
rows.get(0).getAsJsonObject("metadata").get("knowledgeId").getAsString()
|
||||
);
|
||||
} finally {
|
||||
store.close();
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldNormalizeUriWithoutExposingCredentialsInPoolKey() {
|
||||
Assert.assertEquals(
|
||||
"http://127.0.0.1:19530",
|
||||
MilvusClientManager.normalizeAndValidateUri("127.0.0.1:19530")
|
||||
);
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRebuildPoolOnlyWhenConnectionSettingsChange() {
|
||||
MilvusVectorStoreConfig config = new MilvusVectorStoreConfig();
|
||||
config.setUri("http://127.0.0.1:19530");
|
||||
MilvusClientManager manager = new MilvusClientManager(config);
|
||||
try {
|
||||
Assert.assertFalse(manager.reconfigureIfNeeded(config));
|
||||
|
||||
config.setPoolMaxTotal(9);
|
||||
|
||||
Assert.assertTrue(manager.reconfigureIfNeeded(config));
|
||||
Assert.assertFalse(manager.reconfigureIfNeeded(config));
|
||||
} finally {
|
||||
manager.close();
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldPreserveCosineScoreNormalization() {
|
||||
Assert.assertEquals(Double.valueOf(1.0D), MilvusVectorStore.normalizeScore(1.0F));
|
||||
Assert.assertEquals(Double.valueOf(0.5D), MilvusVectorStore.normalizeScore(0.0F));
|
||||
Assert.assertEquals(Double.valueOf(0.0D), MilvusVectorStore.normalizeScore(-1.0F));
|
||||
Assert.assertNull(MilvusVectorStore.normalizeScore(null));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRejectCompatibilityClientAfterStoreCloses() {
|
||||
MilvusVectorStoreConfig config = new MilvusVectorStoreConfig();
|
||||
config.setUri("http://127.0.0.1:19530");
|
||||
MilvusVectorStore store = new MilvusVectorStore(config);
|
||||
|
||||
store.close();
|
||||
|
||||
try {
|
||||
store.getClient();
|
||||
Assert.fail("A closed store must not recreate a compatibility client");
|
||||
} catch (IllegalStateException expected) {
|
||||
Assert.assertEquals(
|
||||
"Milvus vector store is closed", expected.getMessage());
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -34,4 +34,24 @@ public class MilvusVectorStoreConfigTest {
|
||||
config.setPassword("Milvus");
|
||||
Assert.assertTrue(config.checkAvailable());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testPoolDefaultsAreBounded() {
|
||||
MilvusVectorStoreConfig config = new MilvusVectorStoreConfig();
|
||||
Assert.assertEquals(8, config.getPoolMaxTotal());
|
||||
Assert.assertEquals(8, config.getPoolMaxTotalPerKey());
|
||||
Assert.assertEquals(4, config.getPoolMaxIdlePerKey());
|
||||
Assert.assertEquals(1, config.getPoolMinIdlePerKey());
|
||||
Assert.assertEquals(3_000L, config.getPoolMaxWaitMillis());
|
||||
Assert.assertEquals(300_000L, config.getPoolMinEvictableIdleMillis());
|
||||
Assert.assertEquals(10_000L, config.getSearchTimeoutMillis());
|
||||
}
|
||||
|
||||
@Test(expected = IllegalArgumentException.class)
|
||||
public void testSearchTimeoutMustLeaveCleanupMargin() {
|
||||
MilvusVectorStoreConfig config = new MilvusVectorStoreConfig();
|
||||
config.setUri("http://127.0.0.1:19530");
|
||||
config.setSearchTimeoutMillis(200L);
|
||||
new MilvusVectorStore(config);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,935 @@
|
||||
package com.easyagents.store.milvus;
|
||||
|
||||
import com.easyagents.core.document.Document;
|
||||
import com.easyagents.core.store.SearchWrapper;
|
||||
import com.easyagents.core.store.StoreOptions;
|
||||
import com.easyagents.core.store.StoreTimeoutException;
|
||||
import io.grpc.Context;
|
||||
import io.grpc.Server;
|
||||
import io.grpc.netty.shaded.io.grpc.netty.NettyServerBuilder;
|
||||
import io.grpc.stub.ServerCallStreamObserver;
|
||||
import io.grpc.stub.StreamObserver;
|
||||
import io.milvus.grpc.CheckHealthRequest;
|
||||
import io.milvus.grpc.CheckHealthResponse;
|
||||
import io.milvus.grpc.ConnectRequest;
|
||||
import io.milvus.grpc.ConnectResponse;
|
||||
import io.milvus.grpc.CollectionSchema;
|
||||
import io.milvus.grpc.DataType;
|
||||
import io.milvus.grpc.DescribeCollectionRequest;
|
||||
import io.milvus.grpc.DescribeCollectionResponse;
|
||||
import io.milvus.grpc.ErrorCode;
|
||||
import io.milvus.grpc.FieldSchema;
|
||||
import io.milvus.grpc.GetLoadStateRequest;
|
||||
import io.milvus.grpc.GetLoadStateResponse;
|
||||
import io.milvus.grpc.ListDatabasesRequest;
|
||||
import io.milvus.grpc.ListDatabasesResponse;
|
||||
import io.milvus.grpc.LoadCollectionRequest;
|
||||
import io.milvus.grpc.LoadState;
|
||||
import io.milvus.grpc.MilvusServiceGrpc;
|
||||
import io.milvus.grpc.QueryRequest;
|
||||
import io.milvus.grpc.QueryResults;
|
||||
import io.milvus.v2.client.MilvusClientV2;
|
||||
import io.milvus.v2.service.vector.request.QueryReq;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
|
||||
import java.io.IOException;
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.Set;
|
||||
import java.util.concurrent.ConcurrentHashMap;
|
||||
import java.util.concurrent.CountDownLatch;
|
||||
import java.util.concurrent.ExecutionException;
|
||||
import java.util.concurrent.ExecutorService;
|
||||
import java.util.concurrent.Executors;
|
||||
import java.util.concurrent.Future;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.concurrent.TimeoutException;
|
||||
import java.util.concurrent.atomic.AtomicBoolean;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
import java.util.concurrent.atomic.AtomicReference;
|
||||
import java.util.function.BooleanSupplier;
|
||||
|
||||
public class MilvusVectorStoreGrpcTest {
|
||||
|
||||
private static final long AWAIT_SECONDS = 5L;
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldUseOneSdkAttemptAndKeepClientReusable() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 2_000L)) {
|
||||
fixture.manager.markCollectionLoaded("docs");
|
||||
Assert.assertEquals(
|
||||
0L,
|
||||
MilvusClientManager.buildConnectConfig(fixture.config).getRpcDeadlineMs()
|
||||
);
|
||||
|
||||
server.failQueriesWithUnavailable();
|
||||
assertSearchFails(fixture.store, "docs");
|
||||
|
||||
Assert.assertEquals(1, server.queryCalls.get());
|
||||
Assert.assertEquals(0, fixture.manager.getActiveClientCount());
|
||||
Assert.assertEquals(1, fixture.manager.getIdleClientCount());
|
||||
|
||||
server.succeedQueries();
|
||||
Assert.assertTrue(search(fixture.store, "docs").isEmpty());
|
||||
Assert.assertEquals(2, server.queryCalls.get());
|
||||
Assert.assertEquals(1, server.connectCalls.get());
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldReleaseAndReuseClientAfterContextDeadline() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 1_000L)) {
|
||||
fixture.manager.markCollectionLoaded("docs");
|
||||
fixture.manager.withClient(client -> client);
|
||||
QueryBlock block = server.blockQueries();
|
||||
ExecutorService executor = Executors.newSingleThreadExecutor();
|
||||
try {
|
||||
Future<List<Document>> search = executor.submit(() ->
|
||||
search(fixture.store, "docs"));
|
||||
|
||||
block.awaitEntered();
|
||||
block.awaitCancelled();
|
||||
Throwable failure = futureFailure(search);
|
||||
Assert.assertTrue(failure.toString(),
|
||||
failure instanceof StoreTimeoutException);
|
||||
|
||||
Assert.assertTrue(server.querySawDeadline.get());
|
||||
Assert.assertEquals(0, fixture.manager.getActiveClientCount());
|
||||
Assert.assertEquals(1, fixture.manager.getIdleClientCount());
|
||||
|
||||
server.succeedQueries();
|
||||
Assert.assertTrue(search(fixture.store, "docs").isEmpty());
|
||||
Assert.assertEquals(1, server.connectCalls.get());
|
||||
} finally {
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldReleaseAndReuseClientAfterThreadInterrupt() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 5_000L)) {
|
||||
fixture.manager.markCollectionLoaded("docs");
|
||||
fixture.manager.withClient(client -> client);
|
||||
QueryBlock block = server.blockQueries();
|
||||
CountDownLatch finished = new CountDownLatch(1);
|
||||
AtomicReference<Throwable> failure = new AtomicReference<>();
|
||||
AtomicBoolean interrupted = new AtomicBoolean();
|
||||
Thread searchThread = new Thread(() -> {
|
||||
try {
|
||||
search(fixture.store, "docs");
|
||||
} catch (Throwable exception) {
|
||||
failure.set(exception);
|
||||
} finally {
|
||||
interrupted.set(Thread.currentThread().isInterrupted());
|
||||
finished.countDown();
|
||||
}
|
||||
}, "milvus-interrupt-test");
|
||||
searchThread.start();
|
||||
|
||||
block.awaitEntered();
|
||||
searchThread.interrupt();
|
||||
Assert.assertTrue(finished.await(AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
block.awaitCancelled();
|
||||
|
||||
Assert.assertNotNull(failure.get());
|
||||
Assert.assertTrue(interrupted.get());
|
||||
Assert.assertEquals(0, fixture.manager.getActiveClientCount());
|
||||
Assert.assertEquals(1, fixture.manager.getIdleClientCount());
|
||||
|
||||
server.succeedQueries();
|
||||
Assert.assertTrue(search(fixture.store, "docs").isEmpty());
|
||||
Assert.assertEquals(1, server.connectCalls.get());
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldKeepPoolAfterOrdinaryBusinessFailure() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 2_000L)) {
|
||||
MilvusClientV2 first = fixture.manager.withClient(client -> client);
|
||||
|
||||
try {
|
||||
fixture.manager.withClient(client -> {
|
||||
throw new IllegalArgumentException("synthetic business failure");
|
||||
});
|
||||
Assert.fail("The business failure must be propagated");
|
||||
} catch (IllegalArgumentException expected) {
|
||||
Assert.assertEquals("synthetic business failure", expected.getMessage());
|
||||
}
|
||||
|
||||
MilvusClientV2 second = fixture.manager.withClient(client -> client);
|
||||
Assert.assertSame(first, second);
|
||||
Assert.assertEquals(1, server.connectCalls.get());
|
||||
Assert.assertEquals(0, fixture.manager.getActiveClientCount());
|
||||
Assert.assertEquals(1, fixture.manager.getIdleClientCount());
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldAttachContextToEveryClientOperation() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 2_000L)) {
|
||||
Context callerContext = Context.current();
|
||||
Context operationContext = fixture.manager.withClient(client ->
|
||||
Context.current());
|
||||
|
||||
Assert.assertNotSame(callerContext, operationContext);
|
||||
Assert.assertTrue(operationContext.isCancelled());
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldCapPoolWaitByRemainingSearchDeadline() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 5_000L)) {
|
||||
fixture.manager.markCollectionLoaded("docs");
|
||||
ExecutorService executor = Executors.newSingleThreadExecutor();
|
||||
CountDownLatch borrowed = new CountDownLatch(1);
|
||||
CountDownLatch release = new CountDownLatch(1);
|
||||
try {
|
||||
Future<?> holder = executor.submit(() ->
|
||||
fixture.manager.withClient(client -> {
|
||||
borrowed.countDown();
|
||||
try {
|
||||
Assert.assertTrue(release.await(
|
||||
AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
throw new IllegalStateException(exception);
|
||||
}
|
||||
return null;
|
||||
}));
|
||||
Assert.assertTrue(borrowed.await(
|
||||
AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
|
||||
StoreOptions options = StoreOptions.ofCollectionName("docs");
|
||||
options.setTimeoutMillis(500L);
|
||||
long startedAt = System.nanoTime();
|
||||
try {
|
||||
search(fixture.store, options);
|
||||
Assert.fail("Pool wait must respect the remaining deadline");
|
||||
} catch (RuntimeException expected) {
|
||||
Assert.assertTrue(expected.toString(),
|
||||
expected instanceof StoreTimeoutException);
|
||||
long elapsedMillis = TimeUnit.NANOSECONDS.toMillis(
|
||||
System.nanoTime() - startedAt);
|
||||
Assert.assertTrue("elapsedMillis=" + elapsedMillis,
|
||||
elapsedMillis < 800L);
|
||||
}
|
||||
|
||||
release.countDown();
|
||||
holder.get(AWAIT_SECONDS, TimeUnit.SECONDS);
|
||||
Assert.assertEquals(0, fixture.manager.getActiveClientCount());
|
||||
} finally {
|
||||
release.countDown();
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldInvalidateOnlyClosedClient() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 2, 2_000L)) {
|
||||
ExecutorService executor = Executors.newSingleThreadExecutor();
|
||||
CountDownLatch firstBorrowed = new CountDownLatch(1);
|
||||
CountDownLatch releaseFirst = new CountDownLatch(1);
|
||||
try {
|
||||
Future<?> holder = executor.submit(() ->
|
||||
fixture.manager.withClient(client -> {
|
||||
firstBorrowed.countDown();
|
||||
try {
|
||||
Assert.assertTrue(releaseFirst.await(
|
||||
AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
throw new IllegalStateException(exception);
|
||||
}
|
||||
return null;
|
||||
}));
|
||||
Assert.assertTrue(firstBorrowed.await(
|
||||
AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
fixture.manager.withClient(client -> client);
|
||||
releaseFirst.countDown();
|
||||
holder.get(AWAIT_SECONDS, TimeUnit.SECONDS);
|
||||
Assert.assertEquals(2, fixture.manager.getIdleClientCount());
|
||||
|
||||
AtomicReference<MilvusClientV2> closed = new AtomicReference<>();
|
||||
try {
|
||||
fixture.manager.withClient(client -> {
|
||||
closed.set(client);
|
||||
client.close();
|
||||
throw new IllegalStateException("synthetic closed client");
|
||||
});
|
||||
Assert.fail("The closed-client failure must be propagated");
|
||||
} catch (IllegalStateException expected) {
|
||||
Assert.assertEquals(
|
||||
"synthetic closed client", expected.getMessage());
|
||||
}
|
||||
|
||||
Assert.assertEquals(1, fixture.manager.getIdleClientCount());
|
||||
MilvusClientV2 remaining = fixture.manager.withClient(
|
||||
client -> client);
|
||||
Assert.assertNotSame(closed.get(), remaining);
|
||||
Assert.assertEquals(2, server.connectCalls.get());
|
||||
} finally {
|
||||
releaseFirst.countDown();
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldLoadSameCollectionOnlyOnce() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 2, 3_000L)) {
|
||||
LoadGate gate = server.blockLoad("shared");
|
||||
ExecutorService executor = Executors.newFixedThreadPool(2);
|
||||
try {
|
||||
Future<List<Document>> first = executor.submit(() ->
|
||||
search(fixture.store, "shared"));
|
||||
gate.awaitEntered();
|
||||
|
||||
Future<List<Document>> second = executor.submit(() ->
|
||||
search(fixture.store, "shared"));
|
||||
MilvusClientManager.CollectionLoadTicket follower =
|
||||
fixture.manager.beginCollectionLoad("shared");
|
||||
Assert.assertFalse(follower.isLeader());
|
||||
awaitCondition(() -> follower.completion().getNumberOfDependents() > 0);
|
||||
|
||||
Assert.assertEquals(1, server.loadCalls("shared"));
|
||||
Assert.assertEquals(1, fixture.manager.getActiveClientCount());
|
||||
gate.release();
|
||||
|
||||
Assert.assertTrue(first.get(AWAIT_SECONDS, TimeUnit.SECONDS).isEmpty());
|
||||
Assert.assertTrue(second.get(AWAIT_SECONDS, TimeUnit.SECONDS).isEmpty());
|
||||
Assert.assertEquals(1, server.loadCalls("shared"));
|
||||
Assert.assertTrue(fixture.manager.isCollectionLoaded("shared"));
|
||||
} finally {
|
||||
gate.release();
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldRemoveLeaderTicketAfterLateCacheHit() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer()) {
|
||||
MilvusVectorStoreConfig config = new MilvusVectorStoreConfig();
|
||||
config.setUri(server.uri());
|
||||
config.setDefaultCollectionName("docs");
|
||||
config.setPoolMinIdlePerKey(0);
|
||||
config.setSearchTimeoutMillis(2_000L);
|
||||
RacingMilvusClientManager manager =
|
||||
new RacingMilvusClientManager(config, "late-hit");
|
||||
MilvusVectorStore store = new MilvusVectorStore(config, manager);
|
||||
try {
|
||||
Assert.assertTrue(search(store, "late-hit").isEmpty());
|
||||
manager.markCollectionUnloaded("late-hit");
|
||||
|
||||
Assert.assertTrue(search(store, "late-hit").isEmpty());
|
||||
|
||||
Assert.assertEquals(1, server.loadCalls("late-hit"));
|
||||
Assert.assertTrue(manager.isCollectionLoaded("late-hit"));
|
||||
} finally {
|
||||
store.close();
|
||||
manager.close();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldLoadDifferentCollectionsInParallel() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 2, 3_000L)) {
|
||||
LoadGate firstGate = server.blockLoad("first");
|
||||
LoadGate secondGate = server.blockLoad("second");
|
||||
ExecutorService executor = Executors.newFixedThreadPool(2);
|
||||
try {
|
||||
Future<List<Document>> first = executor.submit(() ->
|
||||
search(fixture.store, "first"));
|
||||
Future<List<Document>> second = executor.submit(() ->
|
||||
search(fixture.store, "second"));
|
||||
|
||||
firstGate.awaitEntered();
|
||||
secondGate.awaitEntered();
|
||||
Assert.assertEquals(2, server.activeLoads.get());
|
||||
Assert.assertEquals(2, server.maxConcurrentLoads.get());
|
||||
Assert.assertEquals(2, fixture.manager.getActiveClientCount());
|
||||
|
||||
firstGate.release();
|
||||
secondGate.release();
|
||||
Assert.assertTrue(first.get(AWAIT_SECONDS, TimeUnit.SECONDS).isEmpty());
|
||||
Assert.assertTrue(second.get(AWAIT_SECONDS, TimeUnit.SECONDS).isEmpty());
|
||||
} finally {
|
||||
firstGate.release();
|
||||
secondGate.release();
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldReelectFollowerAfterLoadLeaderTimesOut() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 3_000L)) {
|
||||
LoadGate gate = server.blockLoad("reelect");
|
||||
ExecutorService executor = Executors.newFixedThreadPool(2);
|
||||
try {
|
||||
StoreOptions shortBudget =
|
||||
StoreOptions.ofCollectionName("reelect");
|
||||
shortBudget.setTimeoutMillis(500L);
|
||||
Future<List<Document>> first = executor.submit(() ->
|
||||
search(fixture.store, shortBudget));
|
||||
gate.awaitEntered();
|
||||
|
||||
Future<List<Document>> second = executor.submit(() ->
|
||||
search(fixture.store, "reelect"));
|
||||
awaitCondition(() -> server.loadCalls("reelect") == 2);
|
||||
gate.release();
|
||||
|
||||
Throwable firstFailure = futureFailure(first);
|
||||
Assert.assertTrue(firstFailure.toString(),
|
||||
firstFailure instanceof StoreTimeoutException);
|
||||
Assert.assertTrue(second.get(
|
||||
AWAIT_SECONDS, TimeUnit.SECONDS).isEmpty());
|
||||
Assert.assertEquals(2, server.loadCalls("reelect"));
|
||||
Assert.assertTrue(
|
||||
fixture.manager.isCollectionLoaded("reelect"));
|
||||
} finally {
|
||||
gate.release();
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldCancelWaitingFollowerWithoutBorrowingClient() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 5_000L)) {
|
||||
MilvusClientManager.CollectionLoadTicket leader =
|
||||
fixture.manager.beginCollectionLoad("waiting");
|
||||
CountDownLatch finished = new CountDownLatch(1);
|
||||
AtomicReference<Throwable> failure = new AtomicReference<>();
|
||||
Thread follower = new Thread(() -> {
|
||||
try {
|
||||
search(fixture.store, "waiting");
|
||||
} catch (Throwable exception) {
|
||||
failure.set(exception);
|
||||
} finally {
|
||||
finished.countDown();
|
||||
}
|
||||
}, "milvus-load-follower-test");
|
||||
try {
|
||||
follower.start();
|
||||
awaitCondition(() -> leader.completion().getNumberOfDependents() > 0);
|
||||
|
||||
Assert.assertEquals(0, fixture.manager.getActiveClientCount());
|
||||
Assert.assertEquals(0, server.connectCalls.get());
|
||||
follower.interrupt();
|
||||
Assert.assertTrue(finished.await(AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
|
||||
Assert.assertNotNull(failure.get());
|
||||
Assert.assertFalse(leader.completion().isDone());
|
||||
Assert.assertEquals(0, fixture.manager.getActiveClientCount());
|
||||
Assert.assertEquals(0, server.connectCalls.get());
|
||||
} finally {
|
||||
follower.interrupt();
|
||||
fixture.manager.failCollectionLoad(
|
||||
leader, new IllegalStateException("test cleanup"), false);
|
||||
fixture.manager.endCollectionLoad(leader);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldRecoverAfterCollectionLoadFailure() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 2_000L)) {
|
||||
server.failNextLoad("recoverable");
|
||||
LoadGate gate = server.blockLoad("recoverable");
|
||||
ExecutorService executor = Executors.newFixedThreadPool(2);
|
||||
try {
|
||||
Future<List<Document>> first = executor.submit(() ->
|
||||
search(fixture.store, "recoverable"));
|
||||
gate.awaitEntered();
|
||||
Future<List<Document>> second = executor.submit(() ->
|
||||
search(fixture.store, "recoverable"));
|
||||
MilvusClientManager.CollectionLoadTicket follower =
|
||||
fixture.manager.beginCollectionLoad("recoverable");
|
||||
Assert.assertFalse(follower.isLeader());
|
||||
awaitCondition(() ->
|
||||
follower.completion().getNumberOfDependents() > 0);
|
||||
gate.release();
|
||||
|
||||
assertFutureFails(first);
|
||||
assertFutureFails(second);
|
||||
} finally {
|
||||
gate.release();
|
||||
executor.shutdownNow();
|
||||
}
|
||||
|
||||
Assert.assertEquals(1, server.loadCalls("recoverable"));
|
||||
Assert.assertFalse(fixture.manager.isCollectionLoaded("recoverable"));
|
||||
Assert.assertEquals(0, fixture.manager.getActiveClientCount());
|
||||
Assert.assertEquals(1, fixture.manager.getIdleClientCount());
|
||||
|
||||
Assert.assertTrue(search(fixture.store, "recoverable").isEmpty());
|
||||
Assert.assertEquals(2, server.loadCalls("recoverable"));
|
||||
Assert.assertTrue(fixture.manager.isCollectionLoaded("recoverable"));
|
||||
Assert.assertEquals(1, server.connectCalls.get());
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldCancelActiveRpcWhenManagerCloses() throws Exception {
|
||||
FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 5_000L);
|
||||
ExecutorService executor = Executors.newFixedThreadPool(2);
|
||||
try {
|
||||
fixture.manager.markCollectionLoaded("docs");
|
||||
QueryBlock block = server.blockQueries();
|
||||
Future<?> operation = executor.submit(() ->
|
||||
query(fixture.manager));
|
||||
block.awaitEntered();
|
||||
|
||||
Future<?> close = executor.submit(fixture.manager::close);
|
||||
block.awaitCancelled();
|
||||
close.get(AWAIT_SECONDS, TimeUnit.SECONDS);
|
||||
assertFutureFails(operation);
|
||||
|
||||
try {
|
||||
fixture.manager.withClient(client -> null);
|
||||
Assert.fail("A closed manager must reject client borrows");
|
||||
} catch (IllegalStateException expected) {
|
||||
Assert.assertEquals("Milvus client pool is closed", expected.getMessage());
|
||||
}
|
||||
} finally {
|
||||
executor.shutdownNow();
|
||||
fixture.close();
|
||||
server.close();
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldCancelActiveRpcWhenManagerReconfigures() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 5_000L)) {
|
||||
ExecutorService executor = Executors.newFixedThreadPool(2);
|
||||
try {
|
||||
QueryBlock block = server.blockQueries();
|
||||
Future<?> operation = executor.submit(() ->
|
||||
query(fixture.manager));
|
||||
block.awaitEntered();
|
||||
|
||||
fixture.config.setPoolMaxTotal(2);
|
||||
Future<Boolean> reconfigure = executor.submit(() ->
|
||||
fixture.manager.reconfigureIfNeeded(fixture.config));
|
||||
block.awaitCancelled();
|
||||
|
||||
Assert.assertTrue(reconfigure.get(
|
||||
AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
assertFutureFails(operation);
|
||||
|
||||
server.succeedQueries();
|
||||
query(fixture.manager);
|
||||
} finally {
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private static Object query(MilvusClientManager manager) {
|
||||
return manager.withClient(client -> client.query(
|
||||
QueryReq.builder()
|
||||
.collectionName("docs")
|
||||
.filter("id == \"synthetic-id\"")
|
||||
.outputFields(List.of("id"))
|
||||
.build()
|
||||
));
|
||||
}
|
||||
|
||||
private static List<Document> search(MilvusVectorStore store, String collection) {
|
||||
return search(store, StoreOptions.ofCollectionName(collection));
|
||||
}
|
||||
|
||||
private static List<Document> search(
|
||||
MilvusVectorStore store,
|
||||
StoreOptions options
|
||||
) {
|
||||
SearchWrapper wrapper = new SearchWrapper();
|
||||
wrapper.setWithVector(false);
|
||||
wrapper.eq("id", "synthetic-id");
|
||||
return store.search(wrapper, options);
|
||||
}
|
||||
|
||||
private static void assertSearchFails(MilvusVectorStore store, String collection) {
|
||||
try {
|
||||
search(store, collection);
|
||||
Assert.fail("The synthetic Milvus failure must be propagated");
|
||||
} catch (RuntimeException expected) {
|
||||
Assert.assertNotNull(expected);
|
||||
}
|
||||
}
|
||||
|
||||
private static void assertFutureFails(Future<?> future)
|
||||
throws InterruptedException, TimeoutException {
|
||||
futureFailure(future);
|
||||
}
|
||||
|
||||
private static Throwable futureFailure(Future<?> future)
|
||||
throws InterruptedException, TimeoutException {
|
||||
try {
|
||||
future.get(AWAIT_SECONDS, TimeUnit.SECONDS);
|
||||
Assert.fail("The synthetic Milvus failure must be propagated");
|
||||
} catch (ExecutionException expected) {
|
||||
Assert.assertNotNull(expected.getCause());
|
||||
return expected.getCause();
|
||||
}
|
||||
throw new AssertionError("Expected future to fail");
|
||||
}
|
||||
|
||||
private static void awaitCondition(BooleanSupplier condition)
|
||||
throws InterruptedException, TimeoutException {
|
||||
long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(AWAIT_SECONDS);
|
||||
while (!condition.getAsBoolean()) {
|
||||
if (System.nanoTime() >= deadline) {
|
||||
throw new TimeoutException("Timed out waiting for test condition");
|
||||
}
|
||||
if (Thread.interrupted()) {
|
||||
throw new InterruptedException();
|
||||
}
|
||||
Thread.onSpinWait();
|
||||
}
|
||||
}
|
||||
|
||||
private static final class Fixture implements AutoCloseable {
|
||||
|
||||
private final MilvusVectorStoreConfig config;
|
||||
private final MilvusClientManager manager;
|
||||
private final MilvusVectorStore store;
|
||||
|
||||
private Fixture(FakeMilvusServer server, int poolSize, long searchTimeoutMillis) {
|
||||
config = new MilvusVectorStoreConfig();
|
||||
config.setUri(server.uri());
|
||||
config.setDefaultCollectionName("docs");
|
||||
config.setPoolMaxTotal(poolSize);
|
||||
config.setPoolMaxTotalPerKey(poolSize);
|
||||
config.setPoolMaxIdlePerKey(poolSize);
|
||||
config.setPoolMinIdlePerKey(0);
|
||||
config.setPoolMaxWaitMillis(1_000L);
|
||||
config.setSearchTimeoutMillis(searchTimeoutMillis);
|
||||
manager = new MilvusClientManager(config);
|
||||
store = new MilvusVectorStore(config, manager);
|
||||
}
|
||||
|
||||
@Override
|
||||
public void close() {
|
||||
store.close();
|
||||
manager.close();
|
||||
}
|
||||
}
|
||||
|
||||
private static final class RacingMilvusClientManager
|
||||
extends MilvusClientManager {
|
||||
|
||||
private final String collectionName;
|
||||
private final AtomicInteger observations = new AtomicInteger();
|
||||
|
||||
private RacingMilvusClientManager(
|
||||
MilvusVectorStoreConfig config,
|
||||
String collectionName
|
||||
) {
|
||||
super(config);
|
||||
this.collectionName = collectionName;
|
||||
}
|
||||
|
||||
@Override
|
||||
boolean isCollectionLoaded(String requestedCollectionName) {
|
||||
if (collectionName.equals(requestedCollectionName)) {
|
||||
int observation = observations.getAndIncrement();
|
||||
if (observation == 0) {
|
||||
return false;
|
||||
}
|
||||
if (observation == 1) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
return super.isCollectionLoaded(requestedCollectionName);
|
||||
}
|
||||
}
|
||||
|
||||
private static final class QueryBlock {
|
||||
|
||||
private final CountDownLatch entered = new CountDownLatch(1);
|
||||
private final CountDownLatch cancelled = new CountDownLatch(1);
|
||||
|
||||
private void awaitEntered() throws InterruptedException {
|
||||
Assert.assertTrue(entered.await(AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
}
|
||||
|
||||
private void awaitCancelled() throws InterruptedException {
|
||||
Assert.assertTrue(cancelled.await(AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
}
|
||||
}
|
||||
|
||||
private static final class LoadGate {
|
||||
|
||||
private final CountDownLatch entered = new CountDownLatch(1);
|
||||
private final CountDownLatch release = new CountDownLatch(1);
|
||||
|
||||
private void awaitEntered() throws InterruptedException {
|
||||
Assert.assertTrue(entered.await(AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
}
|
||||
|
||||
private void release() {
|
||||
release.countDown();
|
||||
}
|
||||
}
|
||||
|
||||
private static final class FakeMilvusServer implements AutoCloseable {
|
||||
|
||||
private static final io.milvus.grpc.Status SUCCESS =
|
||||
io.milvus.grpc.Status.newBuilder()
|
||||
.setErrorCode(ErrorCode.Success)
|
||||
.setCode(0)
|
||||
.build();
|
||||
|
||||
private final AtomicInteger connectCalls = new AtomicInteger();
|
||||
private final AtomicInteger queryCalls = new AtomicInteger();
|
||||
private final AtomicInteger activeLoads = new AtomicInteger();
|
||||
private final AtomicInteger maxConcurrentLoads = new AtomicInteger();
|
||||
private final AtomicBoolean querySawDeadline = new AtomicBoolean();
|
||||
private final ConcurrentHashMap<String, AtomicInteger> loadCalls =
|
||||
new ConcurrentHashMap<>();
|
||||
private final ConcurrentHashMap<String, LoadGate> loadGates =
|
||||
new ConcurrentHashMap<>();
|
||||
private final Set<String> loadedCollections = ConcurrentHashMap.newKeySet();
|
||||
private final Set<String> failNextLoads = ConcurrentHashMap.newKeySet();
|
||||
private final ExecutorService rpcExecutor = Executors.newCachedThreadPool();
|
||||
private final Server server;
|
||||
private volatile QueryAction queryAction = QueryAction.SUCCESS;
|
||||
private volatile QueryBlock queryBlock;
|
||||
|
||||
private FakeMilvusServer() throws IOException {
|
||||
server = NettyServerBuilder.forPort(0)
|
||||
.executor(rpcExecutor)
|
||||
.addService(new Service())
|
||||
.build()
|
||||
.start();
|
||||
}
|
||||
|
||||
private String uri() {
|
||||
return "http://127.0.0.1:" + server.getPort();
|
||||
}
|
||||
|
||||
private void succeedQueries() {
|
||||
queryAction = QueryAction.SUCCESS;
|
||||
queryBlock = null;
|
||||
}
|
||||
|
||||
private void failQueriesWithUnavailable() {
|
||||
queryAction = QueryAction.UNAVAILABLE;
|
||||
queryBlock = null;
|
||||
}
|
||||
|
||||
private QueryBlock blockQueries() {
|
||||
QueryBlock block = new QueryBlock();
|
||||
queryBlock = block;
|
||||
queryAction = QueryAction.BLOCK;
|
||||
return block;
|
||||
}
|
||||
|
||||
private LoadGate blockLoad(String collectionName) {
|
||||
LoadGate gate = new LoadGate();
|
||||
loadGates.put(collectionName, gate);
|
||||
return gate;
|
||||
}
|
||||
|
||||
private void failNextLoad(String collectionName) {
|
||||
failNextLoads.add(collectionName);
|
||||
}
|
||||
|
||||
private int loadCalls(String collectionName) {
|
||||
AtomicInteger calls = loadCalls.get(collectionName);
|
||||
return calls == null ? 0 : calls.get();
|
||||
}
|
||||
|
||||
@Override
|
||||
public void close() {
|
||||
for (LoadGate gate : new ArrayList<>(loadGates.values())) {
|
||||
gate.release();
|
||||
}
|
||||
server.shutdownNow();
|
||||
try {
|
||||
server.awaitTermination(AWAIT_SECONDS, TimeUnit.SECONDS);
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
} finally {
|
||||
rpcExecutor.shutdownNow();
|
||||
}
|
||||
}
|
||||
|
||||
private enum QueryAction {
|
||||
SUCCESS,
|
||||
UNAVAILABLE,
|
||||
BLOCK
|
||||
}
|
||||
|
||||
private final class Service extends MilvusServiceGrpc.MilvusServiceImplBase {
|
||||
|
||||
@Override
|
||||
public void connect(
|
||||
ConnectRequest request,
|
||||
StreamObserver<ConnectResponse> observer
|
||||
) {
|
||||
connectCalls.incrementAndGet();
|
||||
observer.onNext(ConnectResponse.newBuilder()
|
||||
.setStatus(SUCCESS)
|
||||
.setIdentifier(1L)
|
||||
.build());
|
||||
observer.onCompleted();
|
||||
}
|
||||
|
||||
@Override
|
||||
public void listDatabases(
|
||||
ListDatabasesRequest request,
|
||||
StreamObserver<ListDatabasesResponse> observer
|
||||
) {
|
||||
observer.onNext(ListDatabasesResponse.newBuilder()
|
||||
.setStatus(SUCCESS)
|
||||
.addDbNames("default")
|
||||
.build());
|
||||
observer.onCompleted();
|
||||
}
|
||||
|
||||
@Override
|
||||
public void checkHealth(
|
||||
CheckHealthRequest request,
|
||||
StreamObserver<CheckHealthResponse> observer
|
||||
) {
|
||||
observer.onNext(CheckHealthResponse.newBuilder()
|
||||
.setStatus(SUCCESS)
|
||||
.setIsHealthy(true)
|
||||
.build());
|
||||
observer.onCompleted();
|
||||
}
|
||||
|
||||
@Override
|
||||
public void getLoadState(
|
||||
GetLoadStateRequest request,
|
||||
StreamObserver<GetLoadStateResponse> observer
|
||||
) {
|
||||
LoadState state = loadedCollections.contains(request.getCollectionName())
|
||||
? LoadState.LoadStateLoaded
|
||||
: LoadState.LoadStateNotLoad;
|
||||
observer.onNext(GetLoadStateResponse.newBuilder()
|
||||
.setStatus(SUCCESS)
|
||||
.setState(state)
|
||||
.build());
|
||||
observer.onCompleted();
|
||||
}
|
||||
|
||||
@Override
|
||||
public void describeCollection(
|
||||
DescribeCollectionRequest request,
|
||||
StreamObserver<DescribeCollectionResponse> observer
|
||||
) {
|
||||
CollectionSchema schema = CollectionSchema.newBuilder()
|
||||
.setName(request.getCollectionName())
|
||||
.addFields(FieldSchema.newBuilder()
|
||||
.setName("id")
|
||||
.setIsPrimaryKey(true)
|
||||
.setDataType(DataType.VarChar)
|
||||
.build())
|
||||
.build();
|
||||
observer.onNext(DescribeCollectionResponse.newBuilder()
|
||||
.setStatus(SUCCESS)
|
||||
.setCollectionName(request.getCollectionName())
|
||||
.setSchema(schema)
|
||||
.build());
|
||||
observer.onCompleted();
|
||||
}
|
||||
|
||||
@Override
|
||||
public void loadCollection(
|
||||
LoadCollectionRequest request,
|
||||
StreamObserver<io.milvus.grpc.Status> observer
|
||||
) {
|
||||
String collectionName = request.getCollectionName();
|
||||
loadCalls.computeIfAbsent(
|
||||
collectionName, ignored -> new AtomicInteger()).incrementAndGet();
|
||||
LoadGate gate = loadGates.get(collectionName);
|
||||
if (gate != null) {
|
||||
int active = activeLoads.incrementAndGet();
|
||||
maxConcurrentLoads.accumulateAndGet(active, Math::max);
|
||||
gate.entered.countDown();
|
||||
try {
|
||||
if (!gate.release.await(AWAIT_SECONDS, TimeUnit.SECONDS)) {
|
||||
observer.onError(io.grpc.Status.DEADLINE_EXCEEDED
|
||||
.withDescription("test load gate timed out")
|
||||
.asRuntimeException());
|
||||
return;
|
||||
}
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
observer.onError(io.grpc.Status.CANCELLED
|
||||
.withCause(exception)
|
||||
.asRuntimeException());
|
||||
return;
|
||||
} finally {
|
||||
activeLoads.decrementAndGet();
|
||||
}
|
||||
}
|
||||
if (failNextLoads.remove(collectionName)) {
|
||||
observer.onNext(io.milvus.grpc.Status.newBuilder()
|
||||
.setErrorCode(ErrorCode.UnexpectedError)
|
||||
.setCode(1)
|
||||
.setReason("synthetic load failure")
|
||||
.build());
|
||||
observer.onCompleted();
|
||||
return;
|
||||
}
|
||||
loadedCollections.add(collectionName);
|
||||
observer.onNext(SUCCESS);
|
||||
observer.onCompleted();
|
||||
}
|
||||
|
||||
@Override
|
||||
public void query(
|
||||
QueryRequest request,
|
||||
StreamObserver<QueryResults> observer
|
||||
) {
|
||||
queryCalls.incrementAndGet();
|
||||
querySawDeadline.compareAndSet(
|
||||
false, Context.current().getDeadline() != null);
|
||||
QueryAction action = queryAction;
|
||||
if (action == QueryAction.UNAVAILABLE) {
|
||||
observer.onError(io.grpc.Status.UNAVAILABLE
|
||||
.withDescription("synthetic query failure")
|
||||
.asRuntimeException());
|
||||
return;
|
||||
}
|
||||
if (action == QueryAction.BLOCK) {
|
||||
QueryBlock block = queryBlock;
|
||||
@SuppressWarnings("unchecked")
|
||||
ServerCallStreamObserver<QueryResults> serverObserver =
|
||||
(ServerCallStreamObserver<QueryResults>) observer;
|
||||
serverObserver.setOnCancelHandler(block.cancelled::countDown);
|
||||
block.entered.countDown();
|
||||
return;
|
||||
}
|
||||
observer.onNext(QueryResults.newBuilder()
|
||||
.setStatus(SUCCESS)
|
||||
.setCollectionName(request.getCollectionName())
|
||||
.build());
|
||||
observer.onCompleted();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,190 @@
|
||||
package com.easyagents.store.milvus;
|
||||
|
||||
import com.easyagents.core.document.Document;
|
||||
import com.easyagents.core.store.SearchWrapper;
|
||||
import com.easyagents.core.store.StoreOptions;
|
||||
import io.milvus.v2.service.collection.request.DropCollectionReq;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Assume;
|
||||
import org.junit.Test;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.UUID;
|
||||
import java.util.concurrent.CountDownLatch;
|
||||
import java.util.concurrent.ExecutionException;
|
||||
import java.util.concurrent.ExecutorService;
|
||||
import java.util.concurrent.Executors;
|
||||
import java.util.concurrent.Future;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.concurrent.atomic.AtomicReference;
|
||||
|
||||
/**
|
||||
* Opt-in compatibility smoke tests for a real Milvus instance.
|
||||
*/
|
||||
public class MilvusVectorStoreIntegrationTest {
|
||||
|
||||
@Test
|
||||
public void shouldCrudAgainstRealMilvusAndRecoverPooledClients() throws Exception {
|
||||
String uri = System.getenv("MILVUS_TEST_URI");
|
||||
Assume.assumeTrue("MILVUS_TEST_URI is not configured", uri != null && !uri.isBlank());
|
||||
String collectionName = "easy_agents_sdk_2311_" + UUID.randomUUID().toString().replace("-", "");
|
||||
MilvusVectorStoreConfig config = new MilvusVectorStoreConfig();
|
||||
config.setUri(uri);
|
||||
config.setDefaultCollectionName(collectionName);
|
||||
config.setPoolMaxTotal(1);
|
||||
config.setPoolMaxTotalPerKey(1);
|
||||
config.setPoolMaxIdlePerKey(1);
|
||||
config.setPoolMinIdlePerKey(0);
|
||||
config.setPoolMaxWaitMillis(250L);
|
||||
MilvusClientManager manager = new MilvusClientManager(config);
|
||||
MilvusVectorStore store = new MilvusVectorStore(config, manager);
|
||||
StoreOptions options = StoreOptions.ofCollectionName(collectionName);
|
||||
try {
|
||||
Document first = document("chunk-1", "first", 1.0F, 0.0F);
|
||||
Document second = document("chunk-2", "second", 0.0F, 1.0F);
|
||||
Assert.assertTrue(store.store(List.of(first, second), options).isSuccess());
|
||||
|
||||
SearchWrapper nearest = new SearchWrapper();
|
||||
nearest.setVector(new float[] { 1.0F, 0.0F });
|
||||
nearest.setMaxResults(1);
|
||||
List<Document> initial = store.search(nearest, options);
|
||||
Assert.assertEquals(1, initial.size());
|
||||
Assert.assertEquals("chunk-1", String.valueOf(initial.get(0).getId()));
|
||||
|
||||
Document updated = document("chunk-1", "updated", 1.0F, 0.0F);
|
||||
Assert.assertTrue(store.update(List.of(updated), options).isSuccess());
|
||||
Assert.assertEquals("updated", store.search(nearest, options).get(0).getContent());
|
||||
|
||||
Assert.assertTrue(store.delete(List.of("chunk-2"), options).isSuccess());
|
||||
SearchWrapper deleted = new SearchWrapper();
|
||||
deleted.setWithVector(false);
|
||||
deleted.eq("id", "chunk-2");
|
||||
Assert.assertTrue(store.search(deleted, options).isEmpty());
|
||||
|
||||
assertQueryFailureIsNotReportedAsEmpty(store, options);
|
||||
assertPoolExhaustionIsBounded(manager);
|
||||
assertBusinessFailurePreservesClient(manager);
|
||||
Assert.assertTrue(store.checkAvailable());
|
||||
} finally {
|
||||
try {
|
||||
manager.withClient(client -> {
|
||||
client.dropCollection(DropCollectionReq.builder()
|
||||
.collectionName(collectionName)
|
||||
.build());
|
||||
return null;
|
||||
});
|
||||
} finally {
|
||||
store.close();
|
||||
manager.close();
|
||||
}
|
||||
}
|
||||
try {
|
||||
manager.getActiveClientCount();
|
||||
Assert.fail("A closed pool must reject further use");
|
||||
} catch (IllegalStateException expected) {
|
||||
Assert.assertEquals("Milvus client pool is closed", expected.getMessage());
|
||||
}
|
||||
}
|
||||
|
||||
private static void assertQueryFailureIsNotReportedAsEmpty(
|
||||
MilvusVectorStore store,
|
||||
StoreOptions options
|
||||
) {
|
||||
SearchWrapper invalid = new SearchWrapper();
|
||||
invalid.setWithVector(false);
|
||||
invalid.eq("id", "chunk-1");
|
||||
StoreOptions invalidOptions = StoreOptions.ofCollectionName(
|
||||
options.getCollectionName()
|
||||
).partitionName("__missing_partition__");
|
||||
try {
|
||||
store.search(invalid, invalidOptions);
|
||||
Assert.fail("A Milvus query failure must not be reported as an empty result");
|
||||
} catch (RuntimeException expected) {
|
||||
Assert.assertNotNull(expected.getMessage());
|
||||
}
|
||||
}
|
||||
|
||||
private static void assertBusinessFailurePreservesClient(MilvusClientManager manager)
|
||||
throws InterruptedException, ExecutionException {
|
||||
ExecutorService executor = Executors.newFixedThreadPool(2);
|
||||
CountDownLatch borrowed = new CountDownLatch(1);
|
||||
CountDownLatch waiterStarted = new CountDownLatch(1);
|
||||
CountDownLatch fail = new CountDownLatch(1);
|
||||
AtomicReference<Object> failedClient = new AtomicReference<>();
|
||||
try {
|
||||
Future<?> failing = executor.submit(() -> {
|
||||
try {
|
||||
manager.withClient(client -> {
|
||||
failedClient.set(client);
|
||||
borrowed.countDown();
|
||||
try {
|
||||
if (!fail.await(2, TimeUnit.SECONDS)) {
|
||||
throw new IllegalStateException("Timed out waiting to fail client");
|
||||
}
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
throw new IllegalStateException(exception);
|
||||
}
|
||||
throw new IllegalStateException("synthetic RPC failure");
|
||||
});
|
||||
Assert.fail("The synthetic client failure must be propagated");
|
||||
} catch (IllegalStateException expected) {
|
||||
Assert.assertEquals("synthetic RPC failure", expected.getMessage());
|
||||
}
|
||||
});
|
||||
Assert.assertTrue(borrowed.await(2, TimeUnit.SECONDS));
|
||||
Future<Object> waiting = executor.submit(() -> {
|
||||
waiterStarted.countDown();
|
||||
return manager.withClient(client -> client);
|
||||
});
|
||||
Assert.assertTrue(waiterStarted.await(2, TimeUnit.SECONDS));
|
||||
fail.countDown();
|
||||
failing.get();
|
||||
Assert.assertSame(failedClient.get(), waiting.get());
|
||||
} finally {
|
||||
fail.countDown();
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
|
||||
private static void assertPoolExhaustionIsBounded(MilvusClientManager manager)
|
||||
throws InterruptedException, ExecutionException {
|
||||
ExecutorService executor = Executors.newFixedThreadPool(2);
|
||||
CountDownLatch borrowed = new CountDownLatch(1);
|
||||
CountDownLatch release = new CountDownLatch(1);
|
||||
try {
|
||||
Future<?> holder = executor.submit(() -> manager.withClient(client -> {
|
||||
borrowed.countDown();
|
||||
try {
|
||||
if (!release.await(2, TimeUnit.SECONDS)) {
|
||||
throw new IllegalStateException("Timed out waiting to release pooled client");
|
||||
}
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
throw new IllegalStateException(exception);
|
||||
}
|
||||
return null;
|
||||
}));
|
||||
Assert.assertTrue(borrowed.await(2, TimeUnit.SECONDS));
|
||||
Future<?> waiter = executor.submit(() -> manager.withClient(client -> null));
|
||||
try {
|
||||
waiter.get();
|
||||
Assert.fail("Pool exhaustion must fail after the configured wait");
|
||||
} catch (ExecutionException expected) {
|
||||
Assert.assertNotNull(expected.getCause());
|
||||
}
|
||||
release.countDown();
|
||||
holder.get();
|
||||
} finally {
|
||||
release.countDown();
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
|
||||
private static Document document(String id, String content, float first, float second) {
|
||||
Document document = Document.of(content);
|
||||
document.setId(id);
|
||||
document.setVector(new float[] { first, second });
|
||||
return document;
|
||||
}
|
||||
}
|
||||
7
pom.xml
7
pom.xml
@@ -54,6 +54,7 @@
|
||||
<calcite.version>1.42.0</calcite.version>
|
||||
<h2.version>2.3.232</h2.version>
|
||||
<quartz.version>2.5.2</quartz.version>
|
||||
<milvus.version>2.3.11</milvus.version>
|
||||
</properties>
|
||||
|
||||
|
||||
@@ -152,6 +153,12 @@
|
||||
<version>${quartz.version}</version>
|
||||
</dependency>
|
||||
|
||||
<dependency>
|
||||
<groupId>io.milvus</groupId>
|
||||
<artifactId>milvus-sdk-java</artifactId>
|
||||
<version>${milvus.version}</version>
|
||||
</dependency>
|
||||
|
||||
<!--easy-agents dependency management-->
|
||||
<dependency>
|
||||
<groupId>com.easyagents</groupId>
|
||||
|
||||
Reference in New Issue
Block a user