|
|
@@ -5,24 +5,27 @@ import com.aliyun.fc.runtime.PojoRequestHandler;
|
|
|
import com.aliyun.ha3engine.vector.models.*;
|
|
|
import com.fasterxml.jackson.databind.JsonNode;
|
|
|
import lombok.extern.slf4j.Slf4j;
|
|
|
+import net.yaoyi.algorithm.domain.model.TaskImage;
|
|
|
+import net.yaoyi.algorithm.domain.model.TaskImageDupDetail;
|
|
|
+import net.yaoyi.algorithm.domain.rabbitmq.TaskImageDupMsg;
|
|
|
+import net.yaoyi.algorithm.domain.rabbitmq.TaskImageMsg;
|
|
|
import net.yaoyi.pipeline.common.JacksonUtils;
|
|
|
import net.yaoyi.pipeline.model.index.TaskImageIndex;
|
|
|
-import net.yaoyi.pipeline.model.mq.ImageMultiMsg;
|
|
|
-import net.yaoyi.pipeline.model.mq.TaskImageMsg;
|
|
|
+import net.yaoyi.pipeline.model.rabbitmq.TaskImageResultMsg;
|
|
|
|
|
|
+import java.math.BigDecimal;
|
|
|
import java.util.*;
|
|
|
import java.util.function.Function;
|
|
|
import java.util.stream.Collectors;
|
|
|
|
|
|
@Slf4j
|
|
|
-public class OpenSearchHandler implements PojoRequestHandler<TaskImageMsg, List<ImageMultiMsg>> {
|
|
|
- private static final String TABLE_NAME = "task_image_cos_2048";
|
|
|
+public class OpenSearchHandler implements PojoRequestHandler<TaskImageMsg, TaskImageResultMsg> {
|
|
|
|
|
|
private static final int TOP_K = 10;
|
|
|
private static final float SCORE_THRESHOLD = 0.7f;
|
|
|
|
|
|
@Override
|
|
|
- public List<ImageMultiMsg> handleRequest(TaskImageMsg taskImageMsg, Context context) {
|
|
|
+ public TaskImageResultMsg handleRequest(TaskImageMsg taskImageMsg, Context context) {
|
|
|
log.info("收到请求消息:{}", JacksonUtils.toJson(taskImageMsg));
|
|
|
Long taskId = taskImageMsg.getTaskId();
|
|
|
String namespace = taskImageMsg.getNamespace();
|
|
|
@@ -30,6 +33,10 @@ public class OpenSearchHandler implements PojoRequestHandler<TaskImageMsg, List<
|
|
|
// 1. 从数据到索引:将 image 数据转换为向量并写入 OpenSearch
|
|
|
saveByTask(taskImageMsg);
|
|
|
|
|
|
+ if (taskImageMsg.getStatus() < 2) {
|
|
|
+ return new TaskImageResultMsg();
|
|
|
+ }
|
|
|
+
|
|
|
try {
|
|
|
Thread.sleep(2000);
|
|
|
} catch (Exception e) {
|
|
|
@@ -37,71 +44,75 @@ public class OpenSearchHandler implements PojoRequestHandler<TaskImageMsg, List<
|
|
|
}
|
|
|
|
|
|
// 2. 以 namespace + taskId 维度来完成数据更新,image相同的数据保持id不变
|
|
|
- List<TaskImageIndex> list = fetchByTaskId(namespace, taskId);
|
|
|
+ List<TaskImageIndex> list = fetchByTaskId(taskImageMsg.getIndexTableName(), namespace, taskId);
|
|
|
for (TaskImageIndex taskImageIndex : list) {
|
|
|
if (taskImageIndex.getVector() == null) {
|
|
|
throw new RuntimeException("尚未更新索引信息,id:" + taskImageIndex.getId());
|
|
|
}
|
|
|
|
|
|
- if (taskImageIndex.getStatus() == 0) {
|
|
|
- taskImageIndex.setMultiResult(null);
|
|
|
- } else {
|
|
|
- // 维度一:namespace 维度
|
|
|
- String nsFilter = String.format("sourceSendPackageDeptId = %d AND taskId != %d AND status = 1", taskImageIndex.getSourceSendPackageDeptId(), taskImageIndex.getTaskId());
|
|
|
- List<TaskImageIndex.MultiResult> nsResult = knnSearch(namespace, taskImageIndex.getVector(), nsFilter, "namespace");
|
|
|
- List<TaskImageIndex.MultiResult> multiResultList = new ArrayList<>(nsResult);
|
|
|
-
|
|
|
- // 维度二:namespace + taskId 维度
|
|
|
- String nsTaskFilter = String.format("sourceSendPackageDeptId = %d AND taskId = %d AND id != \"%s\" AND status = 1", taskImageIndex.getSourceSendPackageDeptId(), taskImageIndex.getTaskId(), taskImageIndex.getId());
|
|
|
- List<TaskImageIndex.MultiResult> nsTaskResult = knnSearch(namespace, taskImageIndex.getVector(), nsTaskFilter, "task");
|
|
|
- multiResultList.addAll(nsTaskResult);
|
|
|
- taskImageIndex.setMultiResult(JacksonUtils.toJson(multiResultList));
|
|
|
- }
|
|
|
+ // 维度一:namespace 维度
|
|
|
+ String nsFilter = String.format("sourceSendPackageDeptId = %d AND taskId != %d AND status > 0", taskImageIndex.getSourceSendPackageDeptId(), taskImageIndex.getTaskId());
|
|
|
+ List<TaskImageIndex.MultiResult> nsResult = knnSearch(taskImageMsg.getIndexTableName(), namespace, taskImageIndex.getVector(), nsFilter, "namespace");
|
|
|
+ List<TaskImageIndex.MultiResult> multiResultList = new ArrayList<>(nsResult);
|
|
|
+
|
|
|
+ // 维度二:namespace + taskId 维度
|
|
|
+ String nsTaskFilter = String.format("sourceSendPackageDeptId = %d AND taskId = %d AND id != \"%s\" AND status > 0", taskImageIndex.getSourceSendPackageDeptId(), taskImageIndex.getTaskId(), taskImageIndex.getId());
|
|
|
+ List<TaskImageIndex.MultiResult> nsTaskResult = knnSearch(taskImageMsg.getIndexTableName(), namespace, taskImageIndex.getVector(), nsTaskFilter, "task");
|
|
|
+ multiResultList.addAll(nsTaskResult);
|
|
|
+ taskImageIndex.setMultiResult(JacksonUtils.toJson(multiResultList));
|
|
|
}
|
|
|
|
|
|
//先落库
|
|
|
- pushDocument(list, Collections.emptyList());
|
|
|
+ pushDocument(taskImageMsg.getIndexTableName(), list, Collections.emptyList());
|
|
|
//构建消息体
|
|
|
- List<ImageMultiMsg> msgList = new ArrayList<>();
|
|
|
+ List<TaskImageDupMsg> msgList = new ArrayList<>();
|
|
|
for (TaskImageIndex taskImageIndex : list) {
|
|
|
msgList.add(buildImageMultiMsg(taskImageMsg, taskImageIndex));
|
|
|
}
|
|
|
- return msgList;
|
|
|
+
|
|
|
+ TaskImageResultMsg taskImageResultMsg = new TaskImageResultMsg();
|
|
|
+ taskImageResultMsg.setQueueName(taskImageMsg.getResultQueueName());
|
|
|
+ taskImageResultMsg.setMsgList(msgList);
|
|
|
+ return taskImageResultMsg;
|
|
|
}
|
|
|
|
|
|
private void saveByTask(TaskImageMsg taskImageMsg) {
|
|
|
- List<TaskImageMsg.ImageInfo> images = taskImageMsg.getImages();
|
|
|
-
|
|
|
- List<TaskImageIndex> taskImageList = fetchByTaskId(taskImageMsg.getNamespace(), taskImageMsg.getTaskId());
|
|
|
- Map<String, TaskImageIndex> map = taskImageList.stream().collect(Collectors.toMap(TaskImageIndex::getImage, Function.identity()));
|
|
|
-
|
|
|
List<TaskImageIndex> addList = new ArrayList<>();
|
|
|
List<TaskImageIndex> deleteList = new ArrayList<>();
|
|
|
- for (TaskImageMsg.ImageInfo imageInfo : images) {
|
|
|
- TaskImageIndex exists = map.remove(imageInfo.getImage());
|
|
|
- if (exists != null) {
|
|
|
- TaskImageIndex index = buildTaskImageIndex(exists.getId(), taskImageMsg, imageInfo);
|
|
|
- index.setVector(exists.getVector());
|
|
|
- addList.add(index);
|
|
|
- } else {
|
|
|
- String id = UUID.randomUUID().toString().replaceAll("-", "");
|
|
|
- addList.add(buildTaskImageIndex(id, taskImageMsg, imageInfo));
|
|
|
+
|
|
|
+ List<TaskImageIndex> taskImageList = fetchByTaskId(taskImageMsg.getIndexTableName(), taskImageMsg.getNamespace(), taskImageMsg.getTaskId());
|
|
|
+ if (taskImageMsg.getStatus() == 0) {
|
|
|
+ deleteList = taskImageList;
|
|
|
+ } else {
|
|
|
+ Map<String, TaskImageIndex> map = taskImageList.stream().collect(Collectors.toMap(TaskImageIndex::getImage, Function.identity()));
|
|
|
+ List<TaskImage> images = taskImageMsg.getImages();
|
|
|
+ for (TaskImage imageInfo : images) {
|
|
|
+ TaskImageIndex exists = map.remove(imageInfo.getImage());
|
|
|
+ if (exists != null) {
|
|
|
+ TaskImageIndex index = buildTaskImageIndex(exists.getId(), taskImageMsg, imageInfo);
|
|
|
+ index.setVector(exists.getVector());
|
|
|
+ index.setMultiResult(exists.getMultiResult());
|
|
|
+ addList.add(index);
|
|
|
+ } else {
|
|
|
+ String id = UUID.randomUUID().toString().replaceAll("-", "");
|
|
|
+ addList.add(buildTaskImageIndex(id, taskImageMsg, imageInfo));
|
|
|
+ }
|
|
|
}
|
|
|
- }
|
|
|
|
|
|
- if (!map.isEmpty()) {
|
|
|
- deleteList = new ArrayList<>(map.values());
|
|
|
+ if (!map.isEmpty()) {
|
|
|
+ deleteList = new ArrayList<>(map.values());
|
|
|
+ }
|
|
|
}
|
|
|
|
|
|
- pushDocument(addList, deleteList);
|
|
|
+ pushDocument(taskImageMsg.getIndexTableName(), addList, deleteList);
|
|
|
}
|
|
|
|
|
|
- private void pushDocument(List<TaskImageIndex> addList, List<TaskImageIndex> deleteList) {
|
|
|
+ private void pushDocument(String tableName, List<TaskImageIndex> addList, List<TaskImageIndex> deleteList) {
|
|
|
if (addList.isEmpty() && deleteList.isEmpty()) {
|
|
|
return;
|
|
|
}
|
|
|
|
|
|
- String fullTableName = Contents.INSTANCE_ID + "_" + TABLE_NAME;
|
|
|
+ String fullTableName = Contents.INSTANCE_ID + "_" + tableName;
|
|
|
|
|
|
ArrayList<Map<String, ?>> documents = new ArrayList<>();
|
|
|
for (TaskImageIndex taskImageIndex : addList) {
|
|
|
@@ -129,10 +140,10 @@ public class OpenSearchHandler implements PojoRequestHandler<TaskImageMsg, List<
|
|
|
}
|
|
|
|
|
|
|
|
|
- private List<TaskImageIndex> fetchByTaskId(String namespace, Long taskId) {
|
|
|
+ private List<TaskImageIndex> fetchByTaskId(String tableName, String namespace, Long taskId) {
|
|
|
try {
|
|
|
FetchRequest request = new FetchRequest();
|
|
|
- request.setTableName(TABLE_NAME);
|
|
|
+ request.setTableName(tableName);
|
|
|
request.setFilter(String.format("taskId=%d AND namespace=\"%s\"", taskId, namespace));
|
|
|
request.setIncludeVector(true);
|
|
|
|
|
|
@@ -206,8 +217,7 @@ public class OpenSearchHandler implements PojoRequestHandler<TaskImageMsg, List<
|
|
|
return index;
|
|
|
}
|
|
|
|
|
|
-
|
|
|
- private List<TaskImageIndex.MultiResult> knnSearch(String namespace, List<Float> vector, String filter, String type) {
|
|
|
+ private List<TaskImageIndex.MultiResult> knnSearch(String tableName, String namespace, List<Float> vector, String filter, String type) {
|
|
|
try {
|
|
|
QueryRequest knn = new QueryRequest();
|
|
|
knn.setNamespace(namespace);
|
|
|
@@ -217,7 +227,7 @@ public class OpenSearchHandler implements PojoRequestHandler<TaskImageMsg, List<
|
|
|
knn.setScoreThreshold(SCORE_THRESHOLD);
|
|
|
|
|
|
SearchRequest request = new SearchRequest();
|
|
|
- request.setTableName(TABLE_NAME);
|
|
|
+ request.setTableName(tableName);
|
|
|
request.setKnn(knn);
|
|
|
request.setOutputFields(List.of("taskId", "image", "taskFiledValue"));
|
|
|
|
|
|
@@ -239,8 +249,7 @@ public class OpenSearchHandler implements PojoRequestHandler<TaskImageMsg, List<
|
|
|
|
|
|
private List<TaskImageIndex.MultiResult> parseToMultiResultList(JsonNode bodyNode, String type) {
|
|
|
JsonNode resultArray;
|
|
|
- if (bodyNode == null || !bodyNode.has("result")
|
|
|
- || !(resultArray = bodyNode.get("result")).isArray() || resultArray.isEmpty()) {
|
|
|
+ if (bodyNode == null || !bodyNode.has("result") || !(resultArray = bodyNode.get("result")).isArray() || resultArray.isEmpty()) {
|
|
|
return Collections.emptyList();
|
|
|
}
|
|
|
|
|
|
@@ -264,7 +273,7 @@ public class OpenSearchHandler implements PojoRequestHandler<TaskImageMsg, List<
|
|
|
return list;
|
|
|
}
|
|
|
|
|
|
- private TaskImageIndex buildTaskImageIndex(String id, TaskImageMsg msg, TaskImageMsg.ImageInfo imageInfo) {
|
|
|
+ private TaskImageIndex buildTaskImageIndex(String id, TaskImageMsg msg, TaskImage imageInfo) {
|
|
|
TaskImageIndex index = new TaskImageIndex();
|
|
|
index.setId(id);
|
|
|
index.setNamespace(msg.getNamespace());
|
|
|
@@ -273,7 +282,7 @@ public class OpenSearchHandler implements PojoRequestHandler<TaskImageMsg, List<
|
|
|
index.setTaskTypeId(msg.getTaskTypeId());
|
|
|
index.setUserId(msg.getUserId());
|
|
|
index.setStatus(msg.getStatus());
|
|
|
- index.setTaskTime(msg.getTaskTime());
|
|
|
+ index.setTaskTime(msg.getTaskTime().toString());
|
|
|
index.setImage(imageInfo.getImage());
|
|
|
index.setTaskFiledValue(imageInfo.getTaskFiledValue());
|
|
|
return index;
|
|
|
@@ -290,23 +299,22 @@ public class OpenSearchHandler implements PojoRequestHandler<TaskImageMsg, List<
|
|
|
return vector;
|
|
|
}
|
|
|
|
|
|
- private ImageMultiMsg buildImageMultiMsg(TaskImageMsg taskImageMsg, TaskImageIndex taskImageIndex) {
|
|
|
- ImageMultiMsg msg = new ImageMultiMsg();
|
|
|
+ private TaskImageDupMsg buildImageMultiMsg(TaskImageMsg taskImageMsg, TaskImageIndex taskImageIndex) {
|
|
|
+ TaskImageDupMsg msg = new TaskImageDupMsg();
|
|
|
msg.setAlgoConfigId(taskImageMsg.getAlgoConfigId());
|
|
|
msg.setTaskId(taskImageIndex.getTaskId());
|
|
|
msg.setImage(taskImageIndex.getImage());
|
|
|
- msg.setTaskFiledValue(taskImageIndex.getTaskFiledValue());
|
|
|
|
|
|
if (taskImageIndex.getMultiResult() != null) {
|
|
|
TaskImageIndex.MultiResult[] results = JacksonUtils.parseObject(taskImageIndex.getMultiResult(), TaskImageIndex.MultiResult[].class);
|
|
|
if (results != null) {
|
|
|
- List<ImageMultiMsg.MultiInfo> multiInfoList = new ArrayList<>();
|
|
|
+ List<TaskImageDupDetail> multiInfoList = new ArrayList<>();
|
|
|
for (TaskImageIndex.MultiResult r : results) {
|
|
|
- ImageMultiMsg.MultiInfo info = new ImageMultiMsg.MultiInfo();
|
|
|
+ TaskImageDupDetail info = new TaskImageDupDetail();
|
|
|
info.setTaskId(r.getTaskId());
|
|
|
info.setImage(r.getImage());
|
|
|
info.setTaskFiledValue(r.getTaskFiledValue());
|
|
|
- info.setScore(r.getScore());
|
|
|
+ info.setScore(BigDecimal.valueOf(r.getScore()));
|
|
|
multiInfoList.add(info);
|
|
|
}
|
|
|
msg.setMultiResult(multiInfoList);
|