package org.example.service; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.stereotype.Service; import java.util.ArrayList; import java.util.Collections; import java.util.List; /** * 向量嵌入服务 * 使用阿里云 DashScope Text Embedding API */ @Service public class VectorEmbeddingService { private static final Logger logger = LoggerFactory.getLogger(VectorEmbeddingService.class); @Autowired private EmbeddingModel embeddingModel; /** * 生成向量嵌入 * 调用阿里云 DashScope Text Embedding API * * @param content 文本内容 * @return 向量嵌入(浮点数列表) */ public List generateEmbedding(String content) { try { if (content == null || content.trim().isEmpty()) { logger.warn("内容为空,无法生成向量"); throw new IllegalArgumentException("内容不能为空"); } logger.debug("开始生成向量嵌入, 内容长度: {} 字符", content.length()); float[] embedding = embeddingModel.embed(content); List floatEmbedding = new ArrayList<>(embedding.length); for (float v : embedding) { floatEmbedding.add(v); } logger.info("成功生成向量嵌入, 内容长度: {} 字符, 向量维度: {}", content.length(), floatEmbedding.size()); return floatEmbedding; } catch (Exception e) { logger.error("生成向量嵌入失败, 内容长度: {}", content != null ? content.length() : 0, e); throw new RuntimeException("生成向量嵌入失败: " + e.getMessage(), e); } } public List> generateEmbeddings(List contents) { try { if (contents == null || contents.isEmpty()) { logger.warn("内容列表为空,无法生成向量"); return Collections.emptyList(); } logger.info("开始批量生成向量嵌入, 数量: {}", contents.size()); List embeddings = embeddingModel.embed(contents); List> result = new ArrayList<>(); for (float[] embedding : embeddings) { List floatEmbedding = new ArrayList<>(embedding.length); for (float v : embedding) { floatEmbedding.add(v); } result.add(floatEmbedding); } logger.info("成功批量生成向量嵌入, 数量: {}, 维度: {}", result.size(), result.isEmpty() ? 0 : result.get(0).size()); return result; } catch (Exception e) { logger.error("批量生成向量嵌入失败", e); throw new RuntimeException("批量生成向量嵌入失败: " + e.getMessage(), e); } } /** * 生成查询向量 * * @param query 查询文本 * @return 向量嵌入 */ public List generateQueryVector(String query) { return generateEmbedding(query); } /** * 计算两个向量的余弦相似度 * * @param vector1 向量1 * @param vector2 向量2 * @return 余弦相似度 [-1, 1] */ public float calculateCosineSimilarity(List vector1, List vector2) { if (vector1.size() != vector2.size()) { throw new IllegalArgumentException("向量维度不匹配"); } float dotProduct = 0.0f; float norm1 = 0.0f; float norm2 = 0.0f; for (int i = 0; i < vector1.size(); i++) { dotProduct += vector1.get(i) * vector2.get(i); norm1 += vector1.get(i) * vector1.get(i); norm2 += vector2.get(i) * vector2.get(i); } return dotProduct / (float) (Math.sqrt(norm1) * Math.sqrt(norm2)); } }