237 lines
8.4 KiB
Java
237 lines
8.4 KiB
Java
package org.example.service;
|
|
|
|
import io.milvus.client.MilvusServiceClient;
|
|
import io.milvus.grpc.DataType;
|
|
import io.milvus.grpc.FlushResponse;
|
|
import io.milvus.grpc.MutationResult;
|
|
import io.milvus.grpc.SearchResults;
|
|
import io.milvus.grpc.ShowCollectionsResponse;
|
|
import io.milvus.common.clientenum.ConsistencyLevelEnum;
|
|
import io.milvus.param.ConnectParam;
|
|
import io.milvus.param.IndexType;
|
|
import io.milvus.param.MetricType;
|
|
import io.milvus.param.R;
|
|
import io.milvus.param.RpcStatus;
|
|
import io.milvus.param.collection.*;
|
|
import io.milvus.param.dml.InsertParam;
|
|
import io.milvus.param.dml.SearchParam;
|
|
import io.milvus.param.index.CreateIndexParam;
|
|
import io.milvus.response.SearchResultsWrapper;
|
|
import org.junit.jupiter.api.*;
|
|
|
|
import java.util.Arrays;
|
|
import java.util.Collections;
|
|
import java.util.List;
|
|
import java.util.concurrent.TimeUnit;
|
|
|
|
import static org.junit.jupiter.api.Assertions.*;
|
|
|
|
@DisplayName("Milvus 连接验证")
|
|
@TestMethodOrder(MethodOrderer.OrderAnnotation.class)
|
|
class MilvusConnectionTest {
|
|
|
|
private static final String COLLECTION = "conn_test";
|
|
private static final int DIM = 128;
|
|
|
|
private static MilvusServiceClient client;
|
|
|
|
@BeforeAll
|
|
static void connect() {
|
|
String host = envOrDefault("MILVUS_HOST",
|
|
"in03-4a578da0f27ce9d.serverless.aws-eu-central-1.cloud.zilliz.com");
|
|
int port = Integer.parseInt(envOrDefault("MILVUS_PORT", "443"));
|
|
String token = System.getenv("MILVUS_TOKEN");
|
|
|
|
assertNotNull(token, "环境变量 MILVUS_TOKEN 未设置");
|
|
|
|
ConnectParam connectParam = ConnectParam.newBuilder()
|
|
.withHost(host)
|
|
.withPort(port)
|
|
.withToken(token)
|
|
.withSecure(true)
|
|
.withDatabaseName("db_4a578da0f27ce9d")
|
|
.withConnectTimeout(30, TimeUnit.SECONDS)
|
|
.build();
|
|
|
|
client = new MilvusServiceClient(connectParam);
|
|
System.out.println("连接目标: " + host + ":" + port);
|
|
}
|
|
|
|
@AfterAll
|
|
static void disconnect() {
|
|
if (client != null) {
|
|
try {
|
|
client.dropCollection(DropCollectionParam.newBuilder()
|
|
.withCollectionName(COLLECTION).build());
|
|
} catch (Exception ignored) {}
|
|
client.close();
|
|
}
|
|
}
|
|
|
|
private static String safeMsg(R<?> resp) {
|
|
try {
|
|
return resp.getMessage();
|
|
} catch (Exception e) {
|
|
return "(no message)";
|
|
}
|
|
}
|
|
|
|
@Test
|
|
@Order(1)
|
|
@DisplayName("1. 连接成功 - 能列出 collection")
|
|
void listCollections() {
|
|
R<ShowCollectionsResponse> resp = client.showCollections(
|
|
ShowCollectionsParam.newBuilder().build());
|
|
|
|
System.out.println("listCollections status: " + resp.getStatus() + ", msg: " + safeMsg(resp));
|
|
assertEquals(0, resp.getStatus(), "连接失败,status=" + resp.getStatus());
|
|
|
|
List<String> names = resp.getData().getCollectionNamesList();
|
|
System.out.println("现有 collections: " + names);
|
|
}
|
|
|
|
@Test
|
|
@Order(2)
|
|
@DisplayName("2. 创建测试 collection")
|
|
void createCollection() {
|
|
client.dropCollection(DropCollectionParam.newBuilder()
|
|
.withCollectionName(COLLECTION).build());
|
|
|
|
FieldType idField = FieldType.newBuilder()
|
|
.withName("id")
|
|
.withDataType(DataType.Int64)
|
|
.withPrimaryKey(true)
|
|
.withAutoID(true)
|
|
.build();
|
|
|
|
FieldType vectorField = FieldType.newBuilder()
|
|
.withName("vector")
|
|
.withDataType(DataType.FloatVector)
|
|
.withDimension(DIM)
|
|
.build();
|
|
|
|
CollectionSchemaParam schema = CollectionSchemaParam.newBuilder()
|
|
.addFieldType(idField)
|
|
.addFieldType(vectorField)
|
|
.build();
|
|
|
|
R<RpcStatus> resp = client.createCollection(
|
|
CreateCollectionParam.newBuilder()
|
|
.withCollectionName(COLLECTION)
|
|
.withSchema(schema)
|
|
.build());
|
|
|
|
System.out.println("createCollection status: " + resp.getStatus() + ", msg: " + safeMsg(resp));
|
|
assertEquals(0, resp.getStatus(), "创建 collection 失败");
|
|
}
|
|
|
|
@Test
|
|
@Order(3)
|
|
@DisplayName("3. 插入数据 + flush")
|
|
void insertAndFlush() {
|
|
List<Float> vec1 = makeVector(1.0f);
|
|
List<Float> vec2 = makeVector(2.0f);
|
|
List<Float> vec3 = makeVector(3.0f);
|
|
|
|
List<InsertParam.Field> fields = Collections.singletonList(
|
|
new InsertParam.Field("vector", Arrays.asList(vec1, vec2, vec3))
|
|
);
|
|
|
|
R<MutationResult> insertResp = client.insert(
|
|
InsertParam.newBuilder()
|
|
.withCollectionName(COLLECTION)
|
|
.withFields(fields)
|
|
.build());
|
|
|
|
System.out.println("insert status: " + insertResp.getStatus() + ", msg: " + safeMsg(insertResp));
|
|
assertEquals(0, insertResp.getStatus(), "插入失败");
|
|
|
|
// 官方示例要求:insert 后必须 flush,数据才对搜索可见
|
|
R<FlushResponse> flushResp = client.flush(FlushParam.newBuilder()
|
|
.withCollectionNames(Collections.singletonList(COLLECTION))
|
|
.withSyncFlush(true)
|
|
.withSyncFlushWaitingTimeout(30L)
|
|
.build());
|
|
|
|
System.out.println("flush status: " + flushResp.getStatus() + ", msg: " + safeMsg(flushResp));
|
|
assertEquals(0, flushResp.getStatus(), "flush 失败");
|
|
System.out.println("插入 3 条数据并 flush 完成");
|
|
}
|
|
|
|
@Test
|
|
@Order(4)
|
|
@DisplayName("4. 创建索引 + 加载")
|
|
void createIndexAndLoad() {
|
|
R<RpcStatus> indexResp = client.createIndex(
|
|
CreateIndexParam.newBuilder()
|
|
.withCollectionName(COLLECTION)
|
|
.withFieldName("vector")
|
|
.withIndexType(IndexType.AUTOINDEX)
|
|
.withMetricType(MetricType.L2)
|
|
.build());
|
|
|
|
System.out.println("createIndex status: " + indexResp.getStatus() + ", msg: " + safeMsg(indexResp));
|
|
assertEquals(0, indexResp.getStatus(), "创建索引失败");
|
|
|
|
R<RpcStatus> loadResp = client.loadCollection(
|
|
LoadCollectionParam.newBuilder()
|
|
.withCollectionName(COLLECTION)
|
|
.withSyncLoad(true)
|
|
.withSyncLoadWaitingTimeout(30L)
|
|
.build());
|
|
|
|
System.out.println("load status: " + loadResp.getStatus() + ", msg: " + safeMsg(loadResp));
|
|
assertEquals(0, loadResp.getStatus(), "加载失败");
|
|
System.out.println("索引创建 + 加载完成");
|
|
}
|
|
|
|
@Test
|
|
@Order(5)
|
|
@DisplayName("5. 向量搜索")
|
|
void search() throws InterruptedException {
|
|
Thread.sleep(3000);
|
|
|
|
List<Float> queryVec = makeVector(1.1f);
|
|
|
|
R<SearchResults> resp = null;
|
|
for (int retry = 0; retry < 10; retry++) {
|
|
resp = client.search(
|
|
SearchParam.newBuilder()
|
|
.withCollectionName(COLLECTION)
|
|
.withMetricType(MetricType.L2)
|
|
.withTopK(2)
|
|
.withVectors(Collections.singletonList(queryVec))
|
|
.withVectorFieldName("vector")
|
|
.withParams("{}")
|
|
.withConsistencyLevel(ConsistencyLevelEnum.STRONG)
|
|
.build());
|
|
|
|
if (resp.getStatus() == 0) break;
|
|
System.out.println("search retry " + (retry + 1) + ": status=" + resp.getStatus() + ", msg=" + safeMsg(resp));
|
|
Thread.sleep(5000);
|
|
}
|
|
|
|
System.out.println("search status: " + resp.getStatus() + ", msg: " + safeMsg(resp));
|
|
assertEquals(0, resp.getStatus(), "搜索失败");
|
|
|
|
SearchResultsWrapper wrapper = new SearchResultsWrapper(resp.getData().getResults());
|
|
List<SearchResultsWrapper.IDScore> scores = wrapper.getIDScore(0);
|
|
|
|
assertFalse(scores.isEmpty(), "搜索结果不应为空");
|
|
System.out.println("搜索结果 (top " + scores.size() + "):");
|
|
for (SearchResultsWrapper.IDScore idScore : scores) {
|
|
System.out.println(" score=" + idScore.getScore() + ", id=" + idScore.getLongID());
|
|
}
|
|
}
|
|
|
|
private static List<Float> makeVector(float val) {
|
|
Float[] arr = new Float[DIM];
|
|
Arrays.fill(arr, val);
|
|
return Arrays.asList(arr);
|
|
}
|
|
|
|
private static String envOrDefault(String key, String defaultVal) {
|
|
String val = System.getenv(key);
|
|
return (val != null && !val.isEmpty()) ? val : defaultVal;
|
|
}
|
|
} |