|
|
@@ -0,0 +1,293 @@
|
|
|
+package net.yaoyi.pipeline.opensearch;
|
|
|
+
|
|
|
+import com.aliyun.fc.runtime.Context;
|
|
|
+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.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 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 = "";
|
|
|
+
|
|
|
+ private static final int TOP_K = 10;
|
|
|
+ private static final float SCORE_THRESHOLD = 0.7f;
|
|
|
+
|
|
|
+ @Override
|
|
|
+ public List<ImageMultiMsg> handleRequest(TaskImageMsg taskImageMsg, Context context) {
|
|
|
+ Long taskId = taskImageMsg.getTaskId();
|
|
|
+ String namespace = taskImageMsg.getNamespace();
|
|
|
+
|
|
|
+ // 1. 从数据到索引:将 image 数据转换为向量并写入 OpenSearch
|
|
|
+ saveByTask(taskImageMsg);
|
|
|
+
|
|
|
+ try {
|
|
|
+ Thread.sleep(2000);
|
|
|
+ } catch (Exception e) {
|
|
|
+ log.warn("等待被中断:" + e.getMessage());
|
|
|
+ }
|
|
|
+
|
|
|
+ // 2. 以 namespace + taskId 维度来完成数据更新,image相同的数据保持id不变
|
|
|
+ List<TaskImageIndex> list = fetchByTaskId(namespace, taskId);
|
|
|
+ for (TaskImageIndex taskImageIndex : list) {
|
|
|
+ if (taskImageIndex.getVector() == null) {
|
|
|
+ throw new RuntimeException("尚未更新索引信息,id:" + taskImageIndex.getId());
|
|
|
+ }
|
|
|
+
|
|
|
+ // 维度一:namespace 维度
|
|
|
+ String nsFilter = "taskId!=" + taskId;
|
|
|
+ List<TaskImageIndex.MultiResult> nsResult = knnSearch(namespace, taskImageIndex.getVector(), nsFilter, "namespace");
|
|
|
+ List<TaskImageIndex.MultiResult> multiResultList = new ArrayList<>(nsResult);
|
|
|
+
|
|
|
+ // 维度二:namespace + taskId 维度
|
|
|
+ String nsTaskFilter = "taskId=" + taskId + " and id!=\"" + taskImageIndex.getId() + "\"";
|
|
|
+ List<TaskImageIndex.MultiResult> nsTaskResult = knnSearch(namespace, taskImageIndex.getVector(), nsTaskFilter, "task");
|
|
|
+ multiResultList.addAll(nsTaskResult);
|
|
|
+ taskImageIndex.setMultiResult(JacksonUtils.toJsonString(multiResultList));
|
|
|
+ }
|
|
|
+
|
|
|
+ //先落库
|
|
|
+ pushDocument(list, Collections.emptyList());
|
|
|
+ //构建消息体
|
|
|
+ List<ImageMultiMsg> msgList = new ArrayList<>();
|
|
|
+ for (TaskImageIndex taskImageIndex : list) {
|
|
|
+ msgList.add(buildImageMultiMsg(taskImageMsg, taskImageIndex));
|
|
|
+ }
|
|
|
+ return msgList;
|
|
|
+ }
|
|
|
+
|
|
|
+ private void saveByTask(TaskImageMsg taskImageMsg) {
|
|
|
+ String[] images = taskImageMsg.getImage();
|
|
|
+
|
|
|
+ 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 (String image : images) {
|
|
|
+ TaskImageIndex exists = map.remove(image);
|
|
|
+ if (exists != null) {
|
|
|
+ TaskImageIndex index = buildTaskImageIndex(exists.getId(), taskImageMsg, image);
|
|
|
+ index.setVector(exists.getVector());
|
|
|
+ addList.add(index);
|
|
|
+ } else {
|
|
|
+ String id = UUID.randomUUID().toString().replaceAll("-", "");
|
|
|
+ addList.add(buildTaskImageIndex(id, taskImageMsg, image));
|
|
|
+ }
|
|
|
+ }
|
|
|
+
|
|
|
+ if (!map.isEmpty()) {
|
|
|
+ deleteList = new ArrayList<>(map.values());
|
|
|
+ }
|
|
|
+
|
|
|
+ pushDocument(addList, deleteList);
|
|
|
+ }
|
|
|
+
|
|
|
+ private void pushDocument(List<TaskImageIndex> addList, List<TaskImageIndex> deleteList) {
|
|
|
+ if (addList.isEmpty() && deleteList.isEmpty()) {
|
|
|
+ return;
|
|
|
+ }
|
|
|
+
|
|
|
+ String fullTableName = Contents.INSTANCE_ID + "_" + TABLE_NAME;
|
|
|
+
|
|
|
+ ArrayList<Map<String, ?>> documents = new ArrayList<>();
|
|
|
+ for (TaskImageIndex taskImageIndex : addList) {
|
|
|
+ documents.add(Map.of("cmd", "add", "fields", JacksonUtils.toMap(taskImageIndex)));
|
|
|
+ }
|
|
|
+ for (TaskImageIndex taskImageIndex : deleteList) {
|
|
|
+ documents.add(Map.of("cmd", "delete", "id", Map.of("id", taskImageIndex.getId())));
|
|
|
+ }
|
|
|
+
|
|
|
+ PushDocumentsRequest request = new PushDocumentsRequest();
|
|
|
+ request.setBody(documents);
|
|
|
+ try {
|
|
|
+ PushDocumentsResponse response = OpenSearchUtils.getClient().pushDocuments(fullTableName, "id", request);
|
|
|
+ String responseBody = response.getBody();
|
|
|
+ JsonNode bodyNode = JacksonUtils.toJsonNode(responseBody);
|
|
|
+
|
|
|
+ if (!bodyNode.has("status") || !"OK".equals(bodyNode.get("status").asText())) {
|
|
|
+ log.error("push document with multiResult failed, request:{}, response:{}", JacksonUtils.toJsonString(request), responseBody);
|
|
|
+ throw new RuntimeException("push document failed");
|
|
|
+ }
|
|
|
+ } catch (Exception e) {
|
|
|
+ log.error("push document with multiResult exception, request:{}", JacksonUtils.toJsonString(request), e);
|
|
|
+ throw new RuntimeException(e);
|
|
|
+ }
|
|
|
+ }
|
|
|
+
|
|
|
+
|
|
|
+ private List<TaskImageIndex> fetchByTaskId(String namespace, Long taskId) {
|
|
|
+ try {
|
|
|
+ FetchRequest request = new FetchRequest();
|
|
|
+ request.setTableName(TABLE_NAME);
|
|
|
+ request.setFilter(String.format("taskId:%d AND namespace:\"%s\"", taskId, namespace));
|
|
|
+ request.setIncludeVector(true);
|
|
|
+
|
|
|
+ SearchResponse response = OpenSearchUtils.getClient().fetch(request);
|
|
|
+ JsonNode bodyNode = JacksonUtils.toJsonNode(response.getBody());
|
|
|
+ if (bodyNode.has("errorCode")) {
|
|
|
+ throw new RuntimeException(String.format("查询失败,namespace:[%s], taskId:%d ", namespace, taskId));
|
|
|
+ }
|
|
|
+
|
|
|
+ int totalCount = bodyNode.get("totalCount").asInt();
|
|
|
+ if (totalCount == 0) {
|
|
|
+ return Collections.emptyList();
|
|
|
+ }
|
|
|
+ return parseToModelList(bodyNode.get("result"));
|
|
|
+ } catch (Exception e) {
|
|
|
+ throw new RuntimeException(String.format("查询异常,namespace:[%s], taskId:%d ", namespace, taskId), e);
|
|
|
+ }
|
|
|
+ }
|
|
|
+
|
|
|
+ private List<TaskImageIndex> parseToModelList(JsonNode result) {
|
|
|
+ List<TaskImageIndex> list = new ArrayList<>();
|
|
|
+ for (JsonNode jsonNode : result) {
|
|
|
+ list.add(parseToModel(jsonNode));
|
|
|
+ }
|
|
|
+ return list;
|
|
|
+ }
|
|
|
+
|
|
|
+ private TaskImageIndex parseToModel(JsonNode result) {
|
|
|
+ TaskImageIndex index = new TaskImageIndex();
|
|
|
+ if (result.has("id")) {
|
|
|
+ index.setId(result.get("id").asText());
|
|
|
+ }
|
|
|
+ if (result.has("namespace")) {
|
|
|
+ index.setNamespace(result.get("namespace").asText());
|
|
|
+ }
|
|
|
+ if (result.has("vector")) {
|
|
|
+ index.setVector(extractVector(result.get("vector")));
|
|
|
+ }
|
|
|
+
|
|
|
+ JsonNode fields = result.get("fields");
|
|
|
+ if (fields != null) {
|
|
|
+ if (fields.has("taskId")) {
|
|
|
+ index.setTaskId(fields.get("taskId").asLong());
|
|
|
+ }
|
|
|
+ if (fields.has("taskTypeId")) {
|
|
|
+ index.setTaskTypeId(fields.get("taskTypeId").asInt());
|
|
|
+ }
|
|
|
+ if (fields.has("userId")) {
|
|
|
+ index.setUserId(fields.get("userId").asLong());
|
|
|
+ }
|
|
|
+ if (fields.has("status")) {
|
|
|
+ index.setStatus(fields.get("status").asInt());
|
|
|
+ }
|
|
|
+ if (fields.has("image")) {
|
|
|
+ index.setImage(fields.get("image").asText());
|
|
|
+ }
|
|
|
+ if (fields.has("taskTime")) {
|
|
|
+ index.setTaskTime(fields.get("taskTime").asText());
|
|
|
+ }
|
|
|
+ if (fields.has("multiResult")) {
|
|
|
+ index.setMultiResult(fields.get("multiResult").asText());
|
|
|
+ }
|
|
|
+ }
|
|
|
+ return index;
|
|
|
+ }
|
|
|
+
|
|
|
+
|
|
|
+ private List<TaskImageIndex.MultiResult> knnSearch(String namespace, List<Float> vector, String filter, String type) {
|
|
|
+ try {
|
|
|
+ QueryRequest knn = new QueryRequest();
|
|
|
+ knn.setNamespace(namespace);
|
|
|
+ knn.setVector(vector);
|
|
|
+ knn.setTopK(TOP_K);
|
|
|
+ knn.setFilter(filter);
|
|
|
+ knn.setScoreThreshold(SCORE_THRESHOLD);
|
|
|
+ knn.setOutputFields(List.of("id", "namespace", "fields"));
|
|
|
+
|
|
|
+ SearchRequest request = new SearchRequest();
|
|
|
+ request.setTableName(TABLE_NAME);
|
|
|
+ request.setKnn(knn);
|
|
|
+
|
|
|
+ SearchResponse response = OpenSearchUtils.getClient().search(request);
|
|
|
+ return parseToMultiResultList(response.getBody(), type);
|
|
|
+ } catch (Exception e) {
|
|
|
+ log.error("knnSearch exception, filter:{}", filter, e);
|
|
|
+ throw new RuntimeException("对比异常,filter:" + filter, e);
|
|
|
+ }
|
|
|
+ }
|
|
|
+
|
|
|
+ private List<TaskImageIndex.MultiResult> parseToMultiResultList(String body, String type) {
|
|
|
+ JsonNode bodyNode = JacksonUtils.toJsonNode(body);
|
|
|
+
|
|
|
+ JsonNode resultArray;
|
|
|
+ if (bodyNode == null || !bodyNode.has("result")
|
|
|
+ || !(resultArray = bodyNode.get("result")).isArray() || resultArray.isEmpty()) {
|
|
|
+ return Collections.emptyList();
|
|
|
+ }
|
|
|
+
|
|
|
+ List<TaskImageIndex.MultiResult> list = new ArrayList<>();
|
|
|
+ for (JsonNode jsonNode : resultArray) {
|
|
|
+ TaskImageIndex.MultiResult multiResult = new TaskImageIndex.MultiResult();
|
|
|
+ JsonNode fields = jsonNode.get("fields");
|
|
|
+ multiResult.setTaskId(fields.get("taskId").asLong());
|
|
|
+ multiResult.setImage(fields.get("image").asText());
|
|
|
+
|
|
|
+ if (jsonNode.has("score")) {
|
|
|
+ multiResult.setScore(jsonNode.get("score").asDouble());
|
|
|
+ } else if (jsonNode.has("similarityScore")) {
|
|
|
+ multiResult.setScore(jsonNode.get("similarityScore").asDouble());
|
|
|
+ }
|
|
|
+ multiResult.setType(type);
|
|
|
+ }
|
|
|
+ return list;
|
|
|
+ }
|
|
|
+
|
|
|
+ private TaskImageIndex buildTaskImageIndex(String id, TaskImageMsg msg, String image) {
|
|
|
+ TaskImageIndex index = new TaskImageIndex();
|
|
|
+ index.setId(id);
|
|
|
+ index.setNamespace(msg.getNamespace());
|
|
|
+ index.setTaskId(msg.getTaskId());
|
|
|
+ index.setTaskTypeId(msg.getTaskTypeId());
|
|
|
+ index.setUserId(msg.getUserId());
|
|
|
+ index.setStatus(msg.getStatus());
|
|
|
+ index.setTaskTime(msg.getTaskTime());
|
|
|
+ index.setImage(image);
|
|
|
+ return index;
|
|
|
+ }
|
|
|
+
|
|
|
+ private List<Float> extractVector(JsonNode vectorNode) {
|
|
|
+ if (vectorNode == null || !vectorNode.isArray() || vectorNode.isEmpty()) {
|
|
|
+ return null;
|
|
|
+ }
|
|
|
+ List<Float> vector = new ArrayList<>(vectorNode.size());
|
|
|
+ for (JsonNode v : vectorNode) {
|
|
|
+ vector.add((float) v.asDouble());
|
|
|
+ }
|
|
|
+ return vector;
|
|
|
+ }
|
|
|
+
|
|
|
+ private ImageMultiMsg buildImageMultiMsg(TaskImageMsg taskImageMsg, TaskImageIndex taskImageIndex) {
|
|
|
+ ImageMultiMsg msg = new ImageMultiMsg();
|
|
|
+ msg.setAlgoConfigId(taskImageMsg.getAlgoConfigId());
|
|
|
+ msg.setTaskId(taskImageIndex.getTaskId());
|
|
|
+ msg.setImage(taskImageIndex.getImage());
|
|
|
+
|
|
|
+ if (taskImageIndex.getMultiResult() != null) {
|
|
|
+ TaskImageIndex.MultiResult[] results = JacksonUtils.parseObject(taskImageIndex.getMultiResult(), TaskImageIndex.MultiResult[].class);
|
|
|
+ if (results != null) {
|
|
|
+ List<ImageMultiMsg.MultiInfo> multiInfoList = new ArrayList<>();
|
|
|
+ for (TaskImageIndex.MultiResult r : results) {
|
|
|
+ ImageMultiMsg.MultiInfo info = new ImageMultiMsg.MultiInfo();
|
|
|
+ info.setTaskId(r.getTaskId());
|
|
|
+ info.setImage(r.getImage());
|
|
|
+ info.setScore(r.getScore());
|
|
|
+ multiInfoList.add(info);
|
|
|
+ }
|
|
|
+ msg.setMultiResult(multiInfoList);
|
|
|
+ }
|
|
|
+ }
|
|
|
+ return msg;
|
|
|
+ }
|
|
|
+
|
|
|
+}
|