OpenSearchHandler.java 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297
  1. package net.yaoyi.pipeline.opensearch;
  2. import com.aliyun.fc.runtime.Context;
  3. import com.aliyun.fc.runtime.PojoRequestHandler;
  4. import com.aliyun.ha3engine.vector.models.*;
  5. import com.fasterxml.jackson.databind.JsonNode;
  6. import lombok.extern.slf4j.Slf4j;
  7. import net.yaoyi.pipeline.common.JacksonUtils;
  8. import net.yaoyi.pipeline.model.index.TaskImageIndex;
  9. import net.yaoyi.pipeline.model.mq.ImageMultiMsg;
  10. import net.yaoyi.pipeline.model.mq.TaskImageMsg;
  11. import java.util.*;
  12. import java.util.function.Function;
  13. import java.util.stream.Collectors;
  14. @Slf4j
  15. public class OpenSearchHandler implements PojoRequestHandler<TaskImageMsg, List<ImageMultiMsg>> {
  16. private static final String TABLE_NAME = "task_image_cos_2048";
  17. private static final int TOP_K = 10;
  18. private static final float SCORE_THRESHOLD = 0.7f;
  19. @Override
  20. public List<ImageMultiMsg> handleRequest(TaskImageMsg taskImageMsg, Context context) {
  21. Long taskId = taskImageMsg.getTaskId();
  22. String namespace = taskImageMsg.getNamespace();
  23. // 1. 从数据到索引:将 image 数据转换为向量并写入 OpenSearch
  24. saveByTask(taskImageMsg);
  25. try {
  26. Thread.sleep(2000);
  27. } catch (Exception e) {
  28. log.warn("等待被中断:" + e.getMessage());
  29. }
  30. // 2. 以 namespace + taskId 维度来完成数据更新,image相同的数据保持id不变
  31. List<TaskImageIndex> list = fetchByTaskId(namespace, taskId);
  32. for (TaskImageIndex taskImageIndex : list) {
  33. if (taskImageIndex.getVector() == null) {
  34. throw new RuntimeException("尚未更新索引信息,id:" + taskImageIndex.getId());
  35. }
  36. // 维度一:namespace 维度
  37. String nsFilter = String.format("sourceSendPackageDeptId = %d and taskId != %d", taskImageIndex.getSourceSendPackageDeptId(), taskImageIndex.getTaskId());
  38. List<TaskImageIndex.MultiResult> nsResult = knnSearch(namespace, taskImageIndex.getVector(), nsFilter, "namespace");
  39. List<TaskImageIndex.MultiResult> multiResultList = new ArrayList<>(nsResult);
  40. // 维度二:namespace + taskId 维度
  41. String nsTaskFilter = String.format("sourceSendPackageDeptId = %d and taskId = %d and id != '%s'", taskImageIndex.getSourceSendPackageDeptId(), taskImageIndex.getTaskId(), taskImageIndex.getId());
  42. List<TaskImageIndex.MultiResult> nsTaskResult = knnSearch(namespace, taskImageIndex.getVector(), nsTaskFilter, "task");
  43. multiResultList.addAll(nsTaskResult);
  44. taskImageIndex.setMultiResult(JacksonUtils.toJsonString(multiResultList));
  45. }
  46. //先落库
  47. pushDocument(list, Collections.emptyList());
  48. //构建消息体
  49. List<ImageMultiMsg> msgList = new ArrayList<>();
  50. for (TaskImageIndex taskImageIndex : list) {
  51. msgList.add(buildImageMultiMsg(taskImageMsg, taskImageIndex));
  52. }
  53. return msgList;
  54. }
  55. private void saveByTask(TaskImageMsg taskImageMsg) {
  56. String[] images = taskImageMsg.getImage();
  57. List<TaskImageIndex> taskImageList = fetchByTaskId(taskImageMsg.getNamespace(), taskImageMsg.getTaskId());
  58. Map<String, TaskImageIndex> map = taskImageList.stream().collect(Collectors.toMap(TaskImageIndex::getImage, Function.identity()));
  59. List<TaskImageIndex> addList = new ArrayList<>();
  60. List<TaskImageIndex> deleteList = new ArrayList<>();
  61. for (String image : images) {
  62. TaskImageIndex exists = map.remove(image);
  63. if (exists != null) {
  64. TaskImageIndex index = buildTaskImageIndex(exists.getId(), taskImageMsg, image);
  65. index.setVector(exists.getVector());
  66. addList.add(index);
  67. } else {
  68. String id = UUID.randomUUID().toString().replaceAll("-", "");
  69. addList.add(buildTaskImageIndex(id, taskImageMsg, image));
  70. }
  71. }
  72. if (!map.isEmpty()) {
  73. deleteList = new ArrayList<>(map.values());
  74. }
  75. pushDocument(addList, deleteList);
  76. }
  77. private void pushDocument(List<TaskImageIndex> addList, List<TaskImageIndex> deleteList) {
  78. if (addList.isEmpty() && deleteList.isEmpty()) {
  79. return;
  80. }
  81. String fullTableName = Contents.INSTANCE_ID + "_" + TABLE_NAME;
  82. ArrayList<Map<String, ?>> documents = new ArrayList<>();
  83. for (TaskImageIndex taskImageIndex : addList) {
  84. documents.add(Map.of("cmd", "add", "fields", JacksonUtils.toMap(taskImageIndex)));
  85. }
  86. for (TaskImageIndex taskImageIndex : deleteList) {
  87. documents.add(Map.of("cmd", "delete", "id", Map.of("id", taskImageIndex.getId())));
  88. }
  89. PushDocumentsRequest request = new PushDocumentsRequest();
  90. request.setBody(documents);
  91. try {
  92. PushDocumentsResponse response = OpenSearchUtils.getClient().pushDocuments(fullTableName, "id", request);
  93. String responseBody = response.getBody();
  94. JsonNode bodyNode = JacksonUtils.toJsonNode(responseBody);
  95. if (!bodyNode.has("status") || !"OK".equals(bodyNode.get("status").asText())) {
  96. log.error("push document with multiResult failed, request:{}, response:{}", JacksonUtils.toJsonString(request), responseBody);
  97. throw new RuntimeException("push document failed");
  98. }
  99. } catch (Exception e) {
  100. log.error("push document with multiResult exception, request:{}", JacksonUtils.toJsonString(request), e);
  101. throw new RuntimeException(e);
  102. }
  103. }
  104. private List<TaskImageIndex> fetchByTaskId(String namespace, Long taskId) {
  105. try {
  106. FetchRequest request = new FetchRequest();
  107. request.setTableName(TABLE_NAME);
  108. request.setFilter(String.format("taskId:%d AND namespace:\"%s\"", taskId, namespace));
  109. request.setIncludeVector(true);
  110. SearchResponse response = OpenSearchUtils.getClient().fetch(request);
  111. JsonNode bodyNode = JacksonUtils.toJsonNode(response.getBody());
  112. if (bodyNode.has("errorCode")) {
  113. throw new RuntimeException(String.format("查询失败,namespace:[%s], taskId:%d ", namespace, taskId));
  114. }
  115. int totalCount = bodyNode.get("totalCount").asInt();
  116. if (totalCount == 0) {
  117. return Collections.emptyList();
  118. }
  119. return parseToModelList(bodyNode.get("result"));
  120. } catch (Exception e) {
  121. throw new RuntimeException(String.format("查询异常,namespace:[%s], taskId:%d ", namespace, taskId), e);
  122. }
  123. }
  124. private List<TaskImageIndex> parseToModelList(JsonNode result) {
  125. List<TaskImageIndex> list = new ArrayList<>();
  126. for (JsonNode jsonNode : result) {
  127. list.add(parseToModel(jsonNode));
  128. }
  129. return list;
  130. }
  131. private TaskImageIndex parseToModel(JsonNode result) {
  132. TaskImageIndex index = new TaskImageIndex();
  133. if (result.has("id")) {
  134. index.setId(result.get("id").asText());
  135. }
  136. if (result.has("namespace")) {
  137. index.setNamespace(result.get("namespace").asText());
  138. }
  139. if (result.has("vector")) {
  140. index.setVector(extractVector(result.get("vector")));
  141. }
  142. JsonNode fields = result.get("fields");
  143. if (fields != null) {
  144. if (fields.has("taskId")) {
  145. index.setTaskId(fields.get("taskId").asLong());
  146. }
  147. if (fields.has("sourceSendPackageDeptId")) {
  148. index.setSourceSendPackageDeptId(fields.get("sourceSendPackageDeptId").asLong());
  149. }
  150. if (fields.has("taskTypeId")) {
  151. index.setTaskTypeId(fields.get("taskTypeId").asInt());
  152. }
  153. if (fields.has("userId")) {
  154. index.setUserId(fields.get("userId").asLong());
  155. }
  156. if (fields.has("status")) {
  157. index.setStatus(fields.get("status").asInt());
  158. }
  159. if (fields.has("image")) {
  160. index.setImage(fields.get("image").asText());
  161. }
  162. if (fields.has("taskTime")) {
  163. index.setTaskTime(fields.get("taskTime").asText());
  164. }
  165. if (fields.has("multiResult")) {
  166. index.setMultiResult(fields.get("multiResult").asText());
  167. }
  168. }
  169. return index;
  170. }
  171. private List<TaskImageIndex.MultiResult> knnSearch(String namespace, List<Float> vector, String filter, String type) {
  172. try {
  173. QueryRequest knn = new QueryRequest();
  174. knn.setNamespace(namespace);
  175. knn.setVector(vector);
  176. knn.setTopK(TOP_K);
  177. knn.setFilter(filter);
  178. knn.setScoreThreshold(SCORE_THRESHOLD);
  179. knn.setOutputFields(List.of("id", "namespace", "fields"));
  180. SearchRequest request = new SearchRequest();
  181. request.setTableName(TABLE_NAME);
  182. request.setKnn(knn);
  183. SearchResponse response = OpenSearchUtils.getClient().search(request);
  184. return parseToMultiResultList(response.getBody(), type);
  185. } catch (Exception e) {
  186. log.error("knnSearch exception, filter:{}", filter, e);
  187. throw new RuntimeException("对比异常,filter:" + filter, e);
  188. }
  189. }
  190. private List<TaskImageIndex.MultiResult> parseToMultiResultList(String body, String type) {
  191. JsonNode bodyNode = JacksonUtils.toJsonNode(body);
  192. JsonNode resultArray;
  193. if (bodyNode == null || !bodyNode.has("result")
  194. || !(resultArray = bodyNode.get("result")).isArray() || resultArray.isEmpty()) {
  195. return Collections.emptyList();
  196. }
  197. List<TaskImageIndex.MultiResult> list = new ArrayList<>();
  198. for (JsonNode jsonNode : resultArray) {
  199. TaskImageIndex.MultiResult multiResult = new TaskImageIndex.MultiResult();
  200. JsonNode fields = jsonNode.get("fields");
  201. multiResult.setTaskId(fields.get("taskId").asLong());
  202. multiResult.setImage(fields.get("image").asText());
  203. if (jsonNode.has("score")) {
  204. multiResult.setScore(jsonNode.get("score").asDouble());
  205. } else if (jsonNode.has("similarityScore")) {
  206. multiResult.setScore(jsonNode.get("similarityScore").asDouble());
  207. }
  208. multiResult.setType(type);
  209. }
  210. return list;
  211. }
  212. private TaskImageIndex buildTaskImageIndex(String id, TaskImageMsg msg, String image) {
  213. TaskImageIndex index = new TaskImageIndex();
  214. index.setId(id);
  215. index.setNamespace(msg.getNamespace());
  216. index.setTaskId(msg.getTaskId());
  217. index.setSourceSendPackageDeptId(msg.getSourceSendPackageDeptId());
  218. index.setTaskTypeId(msg.getTaskTypeId());
  219. index.setUserId(msg.getUserId());
  220. index.setStatus(msg.getStatus());
  221. index.setTaskTime(msg.getTaskTime());
  222. index.setImage(image);
  223. return index;
  224. }
  225. private List<Float> extractVector(JsonNode vectorNode) {
  226. if (vectorNode == null || !vectorNode.isArray() || vectorNode.isEmpty()) {
  227. return null;
  228. }
  229. List<Float> vector = new ArrayList<>(vectorNode.size());
  230. for (JsonNode v : vectorNode) {
  231. vector.add((float) v.asDouble());
  232. }
  233. return vector;
  234. }
  235. private ImageMultiMsg buildImageMultiMsg(TaskImageMsg taskImageMsg, TaskImageIndex taskImageIndex) {
  236. ImageMultiMsg msg = new ImageMultiMsg();
  237. msg.setAlgoConfigId(taskImageMsg.getAlgoConfigId());
  238. msg.setTaskId(taskImageIndex.getTaskId());
  239. msg.setImage(taskImageIndex.getImage());
  240. if (taskImageIndex.getMultiResult() != null) {
  241. TaskImageIndex.MultiResult[] results = JacksonUtils.parseObject(taskImageIndex.getMultiResult(), TaskImageIndex.MultiResult[].class);
  242. if (results != null) {
  243. List<ImageMultiMsg.MultiInfo> multiInfoList = new ArrayList<>();
  244. for (TaskImageIndex.MultiResult r : results) {
  245. ImageMultiMsg.MultiInfo info = new ImageMultiMsg.MultiInfo();
  246. info.setTaskId(r.getTaskId());
  247. info.setImage(r.getImage());
  248. info.setScore(r.getScore());
  249. multiInfoList.add(info);
  250. }
  251. msg.setMultiResult(multiInfoList);
  252. }
  253. }
  254. return msg;
  255. }
  256. }