diff --git a/backend/ruoyi-modules/ruoyi-aihr/src/main/java/org/dromara/aihr/personal/service/PersonalRetrievalService.java b/backend/ruoyi-modules/ruoyi-aihr/src/main/java/org/dromara/aihr/personal/service/PersonalRetrievalService.java new file mode 100644 index 00000000..792a56eb --- /dev/null +++ b/backend/ruoyi-modules/ruoyi-aihr/src/main/java/org/dromara/aihr/personal/service/PersonalRetrievalService.java @@ -0,0 +1,240 @@ +package org.dromara.aihr.personal.service; + +import org.dromara.aihr.personal.domain.PersonalAssistantDto.PersonalSearchRequest; +import org.dromara.aihr.personal.domain.PersonalAssistantDto.SearchHitResponse; +import org.dromara.aihr.personal.domain.PersonalAssistantDto.SearchScope; +import org.dromara.aihr.personal.support.PersonalKnowledgeProperties; +import org.dromara.aihr.personal.support.PersonalOwner; +import org.springframework.beans.factory.ObjectProvider; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.jdbc.core.JdbcTemplate; +import org.springframework.jdbc.core.RowMapper; +import org.springframework.stereotype.Service; + +import java.time.LocalDate; +import java.time.LocalDateTime; +import java.util.ArrayList; +import java.util.Comparator; +import java.util.HashMap; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.Optional; + +@Service +public class PersonalRetrievalService { + + private static final int MAX_LIMIT = 50; + private static final int RRF_K = 60; + + private final JdbcTemplate jdbcTemplate; + private final PersonalVectorStore vectorStore; + private final QueryEmbeddingProvider embeddingProvider; + private final PersonalKnowledgeProperties properties; + + @Autowired + public PersonalRetrievalService(JdbcTemplate jdbcTemplate, PersonalVectorStore vectorStore, + ObjectProvider embeddingProviders, + PersonalKnowledgeProperties properties) { + this(jdbcTemplate, vectorStore, embeddingProviders.getIfAvailable(() -> query -> Optional.empty()), properties); + } + + public PersonalRetrievalService(JdbcTemplate jdbcTemplate, PersonalVectorStore vectorStore, + QueryEmbeddingProvider embeddingProvider, PersonalKnowledgeProperties properties) { + this.jdbcTemplate = jdbcTemplate; + this.vectorStore = vectorStore; + this.embeddingProvider = embeddingProvider == null ? query -> Optional.empty() : embeddingProvider; + this.properties = properties; + } + + public List search(PersonalOwner owner, PersonalSearchRequest request) { + requireOwner(owner); + ValidatedRequest validated = validate(request); + if (!validated.personalScope()) { + return List.of(); + } + + List fulltext = fulltext(owner, validated); + Optional vectorJson; + try { + vectorJson = embeddingProvider.embed(validated.query()); + } catch (RuntimeException ex) { + vectorJson = Optional.empty(); + } + if (vectorJson.isEmpty() || vectorJson.get().isBlank()) { + return fulltext; + } + try { + List vectorMatches = vectorStore.query(owner, vectorJson.get(), validated.limit()); + if (vectorMatches.isEmpty()) { + return fulltext; + } + List hydrated = hydrate(owner, vectorMatches, validated); + return mergeRrf(fulltext, hydrated, validated.limit()); + } catch (RuntimeException ex) { + return fulltext; + } + } + + private List fulltext(PersonalOwner owner, ValidatedRequest request) { + StringBuilder sql = new StringBuilder(""" + select f.id as fragment_id, i.title, f.content, i.captured_at, + match(f.content) against (? in natural language mode) as relevance + from aihr_personal_fragment f + join aihr_personal_item i + on i.id = f.item_id and i.tenant_id = f.tenant_id and i.owner_user_id = f.owner_user_id + where f.tenant_id = ? and f.owner_user_id = ? + and i.status = 'READY' + and match(f.content) against (? in natural language mode) + """); + List args = new ArrayList<>(); + args.add(request.query()); + args.add(owner.tenantId()); + args.add(owner.userId()); + args.add(request.query()); + if (request.dateFrom() != null) { + sql.append(" and i.captured_at >= ?"); + args.add(request.dateFrom().atStartOfDay()); + } + if (request.dateTo() != null) { + sql.append(" and i.captured_at < ?"); + args.add(request.dateTo().plusDays(1).atStartOfDay()); + } + appendItemFilter(sql, args, request.itemIds(), "i.id"); + sql.append(" order by relevance desc, f.id asc limit ?"); + args.add(request.limit()); + return List.copyOf(jdbcTemplate.query(sql.toString(), hitMapper(), args.toArray())); + } + + private List hydrate(PersonalOwner owner, List matches, + ValidatedRequest request) { + List fragmentIds = matches.stream().map(PersonalVectorStore.VectorMatch::fragmentId).distinct().toList(); + if (fragmentIds.isEmpty()) { + return List.of(); + } + StringBuilder sql = new StringBuilder(""" + select f.id as fragment_id, i.title, f.content, i.captured_at, 0 as relevance + from aihr_personal_fragment f + join aihr_personal_item i + on i.id = f.item_id and i.tenant_id = f.tenant_id and i.owner_user_id = f.owner_user_id + where f.tenant_id = ? and f.owner_user_id = ? + and i.status = 'READY' and f.id in ( + """); + sql.append("?,".repeat(fragmentIds.size())); + sql.setLength(sql.length() - 1); + sql.append(")"); + List args = new ArrayList<>(); + args.add(owner.tenantId()); + args.add(owner.userId()); + args.addAll(fragmentIds); + if (request.dateFrom() != null) { + sql.append(" and i.captured_at >= ?"); + args.add(request.dateFrom().atStartOfDay()); + } + if (request.dateTo() != null) { + sql.append(" and i.captured_at < ?"); + args.add(request.dateTo().plusDays(1).atStartOfDay()); + } + appendItemFilter(sql, args, request.itemIds(), "i.id"); + List rows = jdbcTemplate.query(sql.toString(), hitMapper(), args.toArray()); + Map byId = new HashMap<>(); + rows.forEach(hit -> byId.put(hit.sourceId(), hit)); + List ordered = new ArrayList<>(); + for (PersonalVectorStore.VectorMatch match : matches) { + SearchHitResponse hit = byId.get(Long.toString(match.fragmentId())); + if (hit != null) { + ordered.add(new SearchHitResponse(hit.domain(), hit.sourceId(), hit.title(), hit.excerpt(), + hit.capturedAt(), match.score())); + } + } + return ordered; + } + + private RowMapper hitMapper() { + return (rs, rowNum) -> new SearchHitResponse( + "PERSONAL", + Long.toString(rs.getLong("fragment_id")), + rs.getString("title"), + excerpt(rs.getString("content")), + rs.getObject("captured_at", LocalDateTime.class), + rs.getDouble("relevance") + ); + } + + static List mergeRrf(List lexical, List vector, int limit) { + Map hits = new LinkedHashMap<>(); + Map scores = new HashMap<>(); + addRanking(lexical, hits, scores); + addRanking(vector, hits, scores); + return hits.values().stream() + .map(hit -> new SearchHitResponse(hit.domain(), hit.sourceId(), hit.title(), hit.excerpt(), + hit.capturedAt(), scores.getOrDefault(hit.sourceId(), 0D))) + .sorted(Comparator.comparingDouble(SearchHitResponse::score).reversed() + .thenComparing(SearchHitResponse::sourceId)) + .limit(limit) + .toList(); + } + + private static void addRanking(List ranking, Map hits, + Map scores) { + for (int rank = 0; rank < ranking.size(); rank++) { + SearchHitResponse hit = ranking.get(rank); + hits.putIfAbsent(hit.sourceId(), hit); + scores.merge(hit.sourceId(), 1D / (RRF_K + rank + 1), Double::sum); + } + } + + private ValidatedRequest validate(PersonalSearchRequest request) { + if (request == null || request.queryText() == null || request.queryText().isBlank() + || request.queryText().trim().length() > 1000) { + throw new IllegalArgumentException("PERSONAL_SEARCH_QUERY_INVALID"); + } + if (request.dateFrom() != null && request.dateTo() != null && request.dateFrom().isAfter(request.dateTo())) { + throw new IllegalArgumentException("PERSONAL_SEARCH_DATE_RANGE_INVALID"); + } + List itemIds = request.itemIds() == null ? List.of() : request.itemIds().stream().distinct().toList(); + if (itemIds.size() > 100 || itemIds.stream().anyMatch(id -> id == null || id <= 0)) { + throw new IllegalArgumentException("PERSONAL_SEARCH_ITEM_SCOPE_INVALID"); + } + boolean personal = request.scope() == null || request.scope().isEmpty() + || request.scope().contains(SearchScope.PERSONAL); + int configured = properties.getRetrievalLimit() > 0 ? properties.getRetrievalLimit() : 10; + int limit = request.limit() == null ? configured : request.limit(); + limit = Math.max(1, Math.min(MAX_LIMIT, limit)); + return new ValidatedRequest(request.queryText().trim(), personal, request.dateFrom(), request.dateTo(), itemIds, limit); + } + + private static void appendItemFilter(StringBuilder sql, List args, List itemIds, String column) { + if (itemIds.isEmpty()) { + return; + } + sql.append(" and ").append(column).append(" in ("); + sql.append("?,".repeat(itemIds.size())); + sql.setLength(sql.length() - 1); + sql.append(")"); + args.addAll(itemIds); + } + + private static String excerpt(String content) { + if (content == null) { + return ""; + } + String normalized = content.replaceAll("\\s+", " ").trim(); + return normalized.length() <= 240 ? normalized : normalized.substring(0, 240) + "…"; + } + + private static void requireOwner(PersonalOwner owner) { + if (owner == null || owner.tenantId() == null || owner.tenantId().isBlank() || owner.userId() <= 0) { + throw new IllegalStateException("个人知识空间需要有效登录身份"); + } + } + + @FunctionalInterface + public interface QueryEmbeddingProvider { + Optional embed(String queryText); + } + + private record ValidatedRequest(String query, boolean personalScope, LocalDate dateFrom, LocalDate dateTo, + List itemIds, int limit) { + } +} diff --git a/backend/ruoyi-modules/ruoyi-aihr/src/main/java/org/dromara/aihr/personal/service/PersonalVectorStore.java b/backend/ruoyi-modules/ruoyi-aihr/src/main/java/org/dromara/aihr/personal/service/PersonalVectorStore.java new file mode 100644 index 00000000..e1cb5559 --- /dev/null +++ b/backend/ruoyi-modules/ruoyi-aihr/src/main/java/org/dromara/aihr/personal/service/PersonalVectorStore.java @@ -0,0 +1,338 @@ +package org.dromara.aihr.personal.service; + +import com.fasterxml.jackson.core.JsonProcessingException; +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.fasterxml.jackson.databind.node.ArrayNode; +import com.fasterxml.jackson.databind.node.ObjectNode; +import org.dromara.aihr.personal.support.PersonalKnowledgeProperties; +import org.dromara.aihr.personal.support.PersonalOwner; +import org.springframework.stereotype.Service; + +import java.net.URI; +import java.net.http.HttpClient; +import java.net.http.HttpRequest; +import java.net.http.HttpResponse; +import java.time.Duration; +import java.time.LocalDateTime; +import java.util.ArrayList; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.regex.Pattern; + +@Service +public class PersonalVectorStore { + + private static final String UNAVAILABLE = "PERSONAL_VECTOR_STORE_UNAVAILABLE"; + private static final Pattern SAFE_COLLECTION = Pattern.compile("[A-Za-z0-9_-]{1,120}"); + + private final PersonalKnowledgeProperties properties; + private final ObjectMapper objectMapper; + private final HttpTransport transport; + private final String collection; + private final AtomicInteger vectorDimension = new AtomicInteger(); + + public PersonalVectorStore(PersonalKnowledgeProperties properties, ObjectMapper objectMapper) { + this(properties, objectMapper, javaTransport(properties)); + } + + private PersonalVectorStore(PersonalKnowledgeProperties properties, ObjectMapper objectMapper, HttpTransport transport) { + this.properties = properties; + this.objectMapper = objectMapper; + this.transport = transport; + this.collection = validateCollection(properties.getQdrantCollection()); + } + + public static PersonalVectorStore forTest(PersonalKnowledgeProperties properties, ObjectMapper objectMapper, + HttpTransport transport) { + return new PersonalVectorStore(properties, objectMapper, transport); + } + + public void ensureCollection(int dimension) { + validateDimension(dimension); + TransportResponse current = send("GET", collectionPath(), null); + if (success(current.status())) { + int remoteDimension; + try { + JsonNode size = objectMapper.readTree(current.body()).path("result").path("config") + .path("params").path("vectors").path("size"); + if (!size.canConvertToInt() || size.asInt() <= 0) { + throw unavailable(); + } + remoteDimension = size.asInt(); + } catch (Exception ex) { + throw unavailable(); + } + if (remoteDimension != dimension) { + throw new IllegalStateException("PERSONAL_VECTOR_DIMENSION_MISMATCH"); + } + } else if (current.status() == 404) { + ObjectNode vectors = objectMapper.createObjectNode(); + vectors.put("size", dimension); + vectors.put("distance", "Cosine"); + ObjectNode body = objectMapper.createObjectNode(); + body.set("vectors", vectors); + requireMutation(send("PUT", collectionPath(), body)); + } else { + throw unavailable(); + } + setOrValidateDimension(dimension); + ensurePayloadIndex("tenant_id", "keyword"); + ensurePayloadIndex("owner_user_id", "integer"); + ensurePayloadIndex("item_id", "integer"); + } + + /** + * Stores one personal vector. Task 5 deliberately does not fabricate embeddings; a later worker integration must + * call ensureCollection and this method only after a real embedding provider returns a vector. + */ + public void upsert(PersonalOwner owner, VectorPoint point, String vectorJson) { + requireOwner(owner); + if (point == null || point.spaceId() <= 0 || point.itemId() <= 0 || point.fragmentId() <= 0) { + throw new IllegalArgumentException("PERSONAL_VECTOR_POINT_INVALID"); + } + ArrayNode vector = parseVector(vectorJson); + setOrValidateDimension(vector.size()); + + ObjectNode payload = objectMapper.createObjectNode(); + payload.put("tenant_id", owner.tenantId()); + payload.put("owner_user_id", owner.userId()); + payload.put("space_id", point.spaceId()); + payload.put("item_id", point.itemId()); + payload.put("fragment_id", point.fragmentId()); + payload.put("source", point.source() == null ? "" : point.source()); + if (point.capturedAt() != null) { + payload.put("captured_at", point.capturedAt().toString()); + } + ObjectNode qdrantPoint = objectMapper.createObjectNode(); + qdrantPoint.put("id", point.fragmentId()); + qdrantPoint.set("vector", vector); + qdrantPoint.set("payload", payload); + ObjectNode body = objectMapper.createObjectNode(); + body.putArray("points").add(qdrantPoint); + requireMutation(send("PUT", collectionPath() + "/points?wait=true", body)); + } + + public List query(PersonalOwner owner, String vectorJson, int limit) { + requireOwner(owner); + ArrayNode vector = parseVector(vectorJson); + setOrValidateDimension(vector.size()); + ObjectNode body = objectMapper.createObjectNode(); + body.set("query", vector); + body.set("filter", ownerFilter(owner, null)); + body.put("limit", Math.max(1, Math.min(50, limit))); + body.put("with_payload", true); + body.put("with_vector", false); + TransportResponse response; + try { + response = send("POST", collectionPath() + "/points/query", body); + } catch (IllegalStateException ex) { + return List.of(); + } + if (!success(response.status())) { + return List.of(); + } + try { + JsonNode points = objectMapper.readTree(response.body()).path("result").path("points"); + if (!points.isArray()) { + return List.of(); + } + List matches = new ArrayList<>(); + for (JsonNode point : points) { + JsonNode fragmentId = point.path("payload").path("fragment_id"); + JsonNode score = point.path("score"); + if (fragmentId.canConvertToLong() && fragmentId.asLong() > 0 && score.isNumber() + && Double.isFinite(score.asDouble())) { + matches.add(new VectorMatch(fragmentId.asLong(), score.asDouble())); + } + } + return List.copyOf(matches); + } catch (Exception ex) { + return List.of(); + } + } + + public void deleteItem(PersonalOwner owner, long itemId) { + requireOwner(owner); + if (itemId <= 0) { + throw new IllegalArgumentException("PERSONAL_ITEM_ID_INVALID"); + } + ObjectNode body = objectMapper.createObjectNode(); + body.set("filter", ownerFilter(owner, itemId)); + TransportResponse response = send("POST", collectionPath() + "/points/delete?wait=true", body); + if (response.status() != 404) { + requireMutation(response); + } + } + + private void ensurePayloadIndex(String field, String schema) { + ObjectNode body = objectMapper.createObjectNode(); + body.put("field_name", field); + body.put("field_schema", schema); + TransportResponse response = send("PUT", collectionPath() + "/index?wait=true", body); + if (!success(response.status()) && response.status() != 409) { + throw unavailable(); + } + } + + private ObjectNode ownerFilter(PersonalOwner owner, Long itemId) { + ObjectNode filter = objectMapper.createObjectNode(); + ArrayNode must = filter.putArray("must"); + must.add(match("tenant_id", owner.tenantId())); + must.add(match("owner_user_id", owner.userId())); + if (itemId != null) { + must.add(match("item_id", itemId)); + } + return filter; + } + + private ObjectNode match(String key, String value) { + ObjectNode condition = objectMapper.createObjectNode(); + condition.put("key", key); + condition.putObject("match").put("value", value); + return condition; + } + + private ObjectNode match(String key, long value) { + ObjectNode condition = objectMapper.createObjectNode(); + condition.put("key", key); + condition.putObject("match").put("value", value); + return condition; + } + + private ArrayNode parseVector(String json) { + try { + JsonNode parsed = objectMapper.readTree(json == null ? "" : json); + if (!(parsed instanceof ArrayNode array) || array.isEmpty()) { + throw new IllegalArgumentException("PERSONAL_VECTOR_INVALID"); + } + for (JsonNode value : array) { + if (!value.isNumber() || !Double.isFinite(value.asDouble())) { + throw new IllegalArgumentException("PERSONAL_VECTOR_INVALID"); + } + } + return array; + } catch (JsonProcessingException ex) { + throw new IllegalArgumentException("PERSONAL_VECTOR_INVALID"); + } + } + + private void setOrValidateDimension(int dimension) { + validateDimension(dimension); + int current = vectorDimension.get(); + if (current == 0) { + vectorDimension.compareAndSet(0, dimension); + current = vectorDimension.get(); + } + if (current != dimension) { + throw new IllegalArgumentException("PERSONAL_VECTOR_DIMENSION_MISMATCH"); + } + } + + private static void validateDimension(int dimension) { + if (dimension <= 0 || dimension > 65536) { + throw new IllegalArgumentException("PERSONAL_VECTOR_DIMENSION_INVALID"); + } + } + + private void requireOwner(PersonalOwner owner) { + if (owner == null || owner.tenantId() == null || owner.tenantId().isBlank() || owner.userId() <= 0) { + throw new IllegalStateException("个人知识空间需要有效登录身份"); + } + } + + private TransportResponse send(String method, String path, JsonNode body) { + try { + Map headers = new LinkedHashMap<>(); + headers.put("Content-Type", "application/json"); + String apiKey = firstNonBlank(properties.getQdrantApiKey(), System.getProperty("aihr.qdrant.apiKey"), + System.getenv("AIHR_QDRANT_API_KEY")); + if (!apiKey.isBlank()) { + if (apiKey.indexOf('\r') >= 0 || apiKey.indexOf('\n') >= 0) { + throw unavailable(); + } + headers.put("api-key", apiKey); + } + return transport.send(new TransportRequest(method, path, + body == null ? "" : objectMapper.writeValueAsString(body), Map.copyOf(headers))); + } catch (Exception ex) { + throw unavailable(); + } + } + + private void requireMutation(TransportResponse response) { + if (!success(response.status())) { + throw unavailable(); + } + } + + private String collectionPath() { + return "/collections/" + collection; + } + + private static boolean success(int status) { + return status >= 200 && status < 300; + } + + private static IllegalStateException unavailable() { + return new IllegalStateException(UNAVAILABLE); + } + + private static String validateCollection(String configured) { + String value = configured == null || configured.isBlank() ? "aihr_personal_knowledge" : configured.trim(); + if (!SAFE_COLLECTION.matcher(value).matches()) { + throw new IllegalArgumentException("PERSONAL_QDRANT_COLLECTION_INVALID"); + } + return value; + } + + private static HttpTransport javaTransport(PersonalKnowledgeProperties properties) { + String configured = firstNonBlank(properties.getQdrantUrl(), System.getProperty("aihr.qdrant.url"), + System.getenv("AIHR_QDRANT_URL"), "http://127.0.0.1:6333"); + URI base = URI.create(configured.endsWith("/") ? configured.substring(0, configured.length() - 1) : configured); + if (!("http".equalsIgnoreCase(base.getScheme()) || "https".equalsIgnoreCase(base.getScheme())) + || base.getHost() == null || base.getUserInfo() != null || base.getQuery() != null || base.getFragment() != null) { + throw new IllegalArgumentException("PERSONAL_QDRANT_URL_INVALID"); + } + int seconds = Math.max(1, Math.min(30, properties.getQdrantTimeoutSeconds())); + HttpClient client = HttpClient.newBuilder().connectTimeout(Duration.ofSeconds(Math.min(3, seconds))).build(); + return request -> { + HttpRequest.Builder builder = HttpRequest.newBuilder(base.resolve(request.path())) + .timeout(Duration.ofSeconds(seconds)); + request.headers().forEach(builder::header); + builder.method(request.method(), request.body().isEmpty() ? HttpRequest.BodyPublishers.noBody() + : HttpRequest.BodyPublishers.ofString(request.body())); + HttpResponse response = client.send(builder.build(), HttpResponse.BodyHandlers.ofString()); + return new TransportResponse(response.statusCode(), response.body()); + }; + } + + private static String firstNonBlank(String... values) { + for (String value : values) { + if (value != null && !value.isBlank()) { + return value.trim(); + } + } + return ""; + } + + @FunctionalInterface + public interface HttpTransport { + TransportResponse send(TransportRequest request) throws Exception; + } + + public record TransportRequest(String method, String path, String body, Map headers) { + } + + public record TransportResponse(int status, String body) { + } + + public record VectorPoint(long spaceId, long itemId, long fragmentId, String source, + LocalDateTime capturedAt) { + } + + public record VectorMatch(long fragmentId, double score) { + } +} diff --git a/backend/ruoyi-modules/ruoyi-aihr/src/main/java/org/dromara/aihr/personal/support/PersonalKnowledgeProperties.java b/backend/ruoyi-modules/ruoyi-aihr/src/main/java/org/dromara/aihr/personal/support/PersonalKnowledgeProperties.java index db085111..8f74e5f7 100644 --- a/backend/ruoyi-modules/ruoyi-aihr/src/main/java/org/dromara/aihr/personal/support/PersonalKnowledgeProperties.java +++ b/backend/ruoyi-modules/ruoyi-aihr/src/main/java/org/dromara/aihr/personal/support/PersonalKnowledgeProperties.java @@ -15,6 +15,10 @@ public class PersonalKnowledgeProperties { private int maxItems = 1000; private int downloadUrlMinutes = 5; private String qdrantCollection = "aihr_personal_knowledge"; + private String qdrantUrl = ""; + private String qdrantApiKey = ""; + private int qdrantTimeoutSeconds = 3; + private int retrievalLimit = 10; /** Optional sys_oss_config key. Blank selects the system default client. */ private String ossConfigKey = ""; private int chunkSize = 800; diff --git a/backend/ruoyi-modules/ruoyi-aihr/src/test/java/org/dromara/aihr/personal/PersonalRetrievalServiceTest.java b/backend/ruoyi-modules/ruoyi-aihr/src/test/java/org/dromara/aihr/personal/PersonalRetrievalServiceTest.java new file mode 100644 index 00000000..a98967de --- /dev/null +++ b/backend/ruoyi-modules/ruoyi-aihr/src/test/java/org/dromara/aihr/personal/PersonalRetrievalServiceTest.java @@ -0,0 +1,117 @@ +package org.dromara.aihr.personal; + +import org.dromara.aihr.personal.domain.PersonalAssistantDto.PersonalSearchRequest; +import org.dromara.aihr.personal.domain.PersonalAssistantDto.SearchHitResponse; +import org.dromara.aihr.personal.domain.PersonalAssistantDto.SearchScope; +import org.dromara.aihr.personal.service.PersonalRetrievalService; +import org.dromara.aihr.personal.service.PersonalVectorStore; +import org.dromara.aihr.personal.support.PersonalKnowledgeProperties; +import org.dromara.aihr.personal.support.PersonalOwner; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.Tag; +import org.mockito.ArgumentCaptor; +import org.springframework.jdbc.core.JdbcTemplate; +import org.springframework.jdbc.core.RowMapper; + +import java.time.LocalDate; +import java.time.LocalDateTime; +import java.util.List; + +import static org.junit.jupiter.api.Assertions.*; +import static org.mockito.ArgumentMatchers.*; +import static org.mockito.Mockito.*; + +@Tag("dev") +class PersonalRetrievalServiceTest { + + @Test + void fulltextSqlPreservesOwnerJoinFiltersDatesAndPreparedItemIds() { + JdbcTemplate jdbc = mock(JdbcTemplate.class); + when(jdbc.query(anyString(), any(RowMapper.class), any(Object[].class))).thenReturn(List.of()); + PersonalRetrievalService service = service(jdbc, query -> java.util.Optional.empty(), vectorStore(List.of())); + + service.search(new PersonalOwner("tenant-a", 7, null), new PersonalSearchRequest( + "收费标准", List.of(SearchScope.PERSONAL), LocalDate.of(2026, 1, 1), LocalDate.of(2026, 2, 1), List.of(3L, 5L), 200)); + + ArgumentCaptor sql = ArgumentCaptor.forClass(String.class); + ArgumentCaptor args = ArgumentCaptor.forClass(Object[].class); + verify(jdbc).query(sql.capture(), any(RowMapper.class), args.capture()); + String normalized = sql.getValue().replaceAll("\\s+", " "); + assertTrue(normalized.contains("i.id = f.item_id and i.tenant_id = f.tenant_id and i.owner_user_id = f.owner_user_id")); + assertTrue(normalized.contains("f.tenant_id = ? and f.owner_user_id = ?")); + assertTrue(normalized.contains("i.status = 'READY'")); + assertTrue(normalized.contains("match(f.content) against (? in natural language mode)")); + assertTrue(normalized.contains("i.id in (?,?)")); + assertFalse(sql.getValue().contains("3,5")); + assertEquals("tenant-a", args.getValue()[1]); + assertEquals(7L, args.getValue()[2]); + assertEquals(50, args.getValue()[args.getValue().length - 1]); + } + + @Test + void excludesPersonalScopeAndRejectsInvalidRequests() { + JdbcTemplate jdbc = mock(JdbcTemplate.class); + PersonalRetrievalService service = service(jdbc, query -> java.util.Optional.empty(), vectorStore(List.of())); + PersonalOwner owner = new PersonalOwner("t", 1, null); + assertTrue(service.search(owner, new PersonalSearchRequest("q", List.of(SearchScope.ENTERPRISE), null, null, null, 10)).isEmpty()); + assertThrows(IllegalArgumentException.class, () -> service.search(owner, new PersonalSearchRequest(" ", null, null, null, null, 10))); + assertThrows(IllegalArgumentException.class, () -> service.search(owner, new PersonalSearchRequest("q", null, + LocalDate.of(2026, 2, 1), LocalDate.of(2026, 1, 1), null, 10))); + verifyNoInteractions(jdbc); + } + + @Test + void vectorHydrationRechecksOwnerAndReadyAndRrfDedupesDeterministically() { + JdbcTemplate jdbc = mock(JdbcTemplate.class); + SearchHitResponse lexical = hit("10", "Lexical"); + SearchHitResponse vector = hit("20", "Vector"); + when(jdbc.query(anyString(), any(RowMapper.class), any(Object[].class))) + .thenReturn(List.of(lexical), List.of(vector)); + PersonalVectorStore vectors = vectorStore(List.of( + new PersonalVectorStore.VectorMatch(20, .99), new PersonalVectorStore.VectorMatch(10, .8))); + PersonalRetrievalService service = service(jdbc, query -> java.util.Optional.of("[0.1,0.2]"), vectors); + + List hits = service.search(new PersonalOwner("tenant-a", 7, null), + new PersonalSearchRequest("问题", List.of(SearchScope.PERSONAL), LocalDate.of(2026, 1, 1), + LocalDate.of(2026, 1, 31), List.of(55L), 10)); + + assertEquals(List.of("10", "20"), hits.stream().map(SearchHitResponse::sourceId).toList()); + ArgumentCaptor sql = ArgumentCaptor.forClass(String.class); + verify(jdbc, times(2)).query(sql.capture(), any(RowMapper.class), any(Object[].class)); + String hydration = sql.getAllValues().get(1).replaceAll("\\s+", " "); + assertTrue(hydration.contains("f.tenant_id = ? and f.owner_user_id = ?")); + assertTrue(hydration.contains("i.status = 'READY'")); + assertTrue(hydration.contains("f.id in (")); + assertTrue(hydration.contains("i.captured_at >= ?")); + assertTrue(hydration.contains("i.captured_at < ?")); + assertTrue(hydration.contains("i.id in (?)")); + } + + @Test + void missingEmbeddingOrQdrantFailureFallsBackToFulltext() { + JdbcTemplate jdbc = mock(JdbcTemplate.class); + when(jdbc.query(anyString(), any(RowMapper.class), any(Object[].class))).thenReturn(List.of(hit("1", "Only"))); + PersonalVectorStore broken = mock(PersonalVectorStore.class); + when(broken.query(any(), anyString(), anyInt())).thenThrow(new IllegalStateException("down")); + PersonalRetrievalService service = service(jdbc, query -> java.util.Optional.of("[1,2]"), broken); + assertEquals(List.of("1"), service.search(new PersonalOwner("t", 1, null), + new PersonalSearchRequest("q", null, null, null, null, 10)).stream().map(SearchHitResponse::sourceId).toList()); + verify(jdbc, times(1)).query(anyString(), any(RowMapper.class), any(Object[].class)); + } + + private PersonalRetrievalService service(JdbcTemplate jdbc, PersonalRetrievalService.QueryEmbeddingProvider provider, + PersonalVectorStore vectors) { + PersonalKnowledgeProperties properties = new PersonalKnowledgeProperties(); + return new PersonalRetrievalService(jdbc, vectors, provider, properties); + } + + private PersonalVectorStore vectorStore(List matches) { + PersonalVectorStore vectors = mock(PersonalVectorStore.class); + when(vectors.query(any(), anyString(), anyInt())).thenReturn(matches); + return vectors; + } + + private SearchHitResponse hit(String id, String title) { + return new SearchHitResponse("PERSONAL", id, title, title + " excerpt", LocalDateTime.of(2026, 1, 1, 0, 0), 1); + } +} diff --git a/backend/ruoyi-modules/ruoyi-aihr/src/test/java/org/dromara/aihr/personal/PersonalVectorStoreTest.java b/backend/ruoyi-modules/ruoyi-aihr/src/test/java/org/dromara/aihr/personal/PersonalVectorStoreTest.java new file mode 100644 index 00000000..9a34e6b0 --- /dev/null +++ b/backend/ruoyi-modules/ruoyi-aihr/src/test/java/org/dromara/aihr/personal/PersonalVectorStoreTest.java @@ -0,0 +1,145 @@ +package org.dromara.aihr.personal; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import org.dromara.aihr.personal.service.PersonalVectorStore; +import org.dromara.aihr.personal.support.PersonalKnowledgeProperties; +import org.dromara.aihr.personal.support.PersonalOwner; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.Tag; + +import java.time.LocalDateTime; +import java.util.ArrayList; +import java.util.List; +import java.util.Map; + +import static org.junit.jupiter.api.Assertions.*; + +@Tag("dev") +class PersonalVectorStoreTest { + + private final ObjectMapper mapper = new ObjectMapper(); + + @Test + void queryAndDeleteAlwaysCarryTenantAndNumericOwnerFilters() throws Exception { + List seen = new ArrayList<>(); + PersonalVectorStore store = fixture(seen, request -> request.path().endsWith("/points/query") + ? ok("{\"result\":{\"points\":[{\"id\":\"99\",\"score\":0.9,\"payload\":{\"fragment_id\":99,\"owner_user_id\":999}}]}}") + : ok("{}")); + PersonalOwner owner = new PersonalOwner("000001", 42, "ext"); + + assertEquals(99, store.query(owner, "[0.1,0.2]", 100).get(0).fragmentId()); + store.deleteItem(owner, 7); + + assertEquals("/collections/aihr_personal_knowledge/points/query", seen.get(0).path()); + JsonNode query = mapper.readTree(seen.get(0).body()); + assertEquals(50, query.path("limit").asInt()); + assertFilter(query.path("filter"), "000001", 42, null); + JsonNode delete = mapper.readTree(seen.get(1).body()); + assertFilter(delete.path("filter"), "000001", 42, 7L); + assertFalse(seen.get(0).body().contains("999"), "Qdrant payload owner must not influence authorization filter"); + } + + @Test + void upsertUsesPersonalCollectionAndOwnerScopedPayload() throws Exception { + List seen = new ArrayList<>(); + PersonalVectorStore store = fixture(seen, request -> ok("{}")); + PersonalOwner owner = new PersonalOwner("t-1", 8, "ext"); + + store.upsert(owner, new PersonalVectorStore.VectorPoint(11, 12, 13, "file", LocalDateTime.of(2026, 7, 1, 9, 0)), "[1,2]"); + + JsonNode body = mapper.readTree(seen.get(0).body()); + JsonNode point = body.path("points").get(0); + assertEquals(13, point.path("id").asLong()); + assertEquals("t-1", point.path("payload").path("tenant_id").asText()); + assertEquals(8, point.path("payload").path("owner_user_id").asLong()); + assertEquals(11, point.path("payload").path("space_id").asLong()); + assertEquals(12, point.path("payload").path("item_id").asLong()); + assertEquals(13, point.path("payload").path("fragment_id").asLong()); + assertFalse(point.path("payload").has("content")); + } + + @Test + void collectionCreationAndPayloadIndexesAreStable() throws Exception { + List seen = new ArrayList<>(); + PersonalVectorStore store = fixture(seen, request -> request.method().equals("GET") + ? new PersonalVectorStore.TransportResponse(404, "") : ok("{}")); + + store.ensureCollection(2); + + assertEquals("/collections/aihr_personal_knowledge", seen.get(1).path()); + assertEquals(2, mapper.readTree(seen.get(1).body()).path("vectors").path("size").asInt()); + List indexFields = seen.stream().filter(r -> r.path().contains("/index?")) + .map(r -> read(r.body()).path("field_name").asText()).toList(); + assertEquals(List.of("tenant_id", "owner_user_id", "item_id"), indexFields); + } + + @Test + void rejectsMalformedNonFiniteAndDimensionMismatchBeforeTransport() { + List seen = new ArrayList<>(); + PersonalVectorStore store = fixture(seen, request -> ok("{}")); + PersonalOwner owner = new PersonalOwner("t", 1, null); + var point = new PersonalVectorStore.VectorPoint(1, 2, 3, "text", LocalDateTime.now()); + + for (String invalid : List.of("", "{}", "[]", "[1,\"x\"]", "[1e999]", "[NaN]")) { + assertThrows(IllegalArgumentException.class, () -> store.upsert(owner, point, invalid), invalid); + } + store.upsert(owner, point, "[1,2]"); + assertThrows(IllegalArgumentException.class, () -> store.upsert(owner, point, "[1,2,3]")); + assertThrows(IllegalStateException.class, () -> store.query(null, "[1,2]", 5)); + } + + @Test + void validatesCollectionAndDoesNotLeakRawQdrantErrors() { + PersonalKnowledgeProperties properties = properties(); + properties.setQdrantCollection("../enterprise"); + assertThrows(IllegalArgumentException.class, () -> PersonalVectorStore.forTest(properties, mapper, request -> ok("{}"))); + + PersonalVectorStore store = PersonalVectorStore.forTest(properties(), mapper, + request -> new PersonalVectorStore.TransportResponse(500, "secret vector and api-key")); + IllegalStateException error = assertThrows(IllegalStateException.class, () -> store.ensureCollection(2)); + assertEquals("PERSONAL_VECTOR_STORE_UNAVAILABLE", error.getMessage()); + assertFalse(error.getMessage().contains("secret")); + + PersonalVectorStore malformed = PersonalVectorStore.forTest(properties(), mapper, + request -> new PersonalVectorStore.TransportResponse(200, null)); + assertTrue(malformed.query(new PersonalOwner("t", 1, null), "[1,2]", 5).isEmpty()); + + PersonalKnowledgeProperties injected = properties(); + injected.setQdrantApiKey("secret\r\nX-Evil: yes"); + List requests = new ArrayList<>(); + PersonalVectorStore safe = PersonalVectorStore.forTest(injected, mapper, request -> { + requests.add(request); + return ok("{}"); + }); + assertTrue(safe.query(new PersonalOwner("t", 1, null), "[1,2]", 5).isEmpty()); + assertTrue(requests.isEmpty()); + } + + private PersonalVectorStore fixture(List seen, PersonalVectorStore.HttpTransport delegate) { + return PersonalVectorStore.forTest(properties(), mapper, request -> { seen.add(request); return delegate.send(request); }); + } + + private PersonalKnowledgeProperties properties() { + PersonalKnowledgeProperties properties = new PersonalKnowledgeProperties(); + properties.setQdrantCollection("aihr_personal_knowledge"); + return properties; + } + + private PersonalVectorStore.TransportResponse ok(String body) { + return new PersonalVectorStore.TransportResponse(200, body); + } + + private JsonNode read(String body) { + try { return mapper.readTree(body); } catch (Exception e) { throw new AssertionError(e); } + } + + private void assertFilter(JsonNode filter, String tenant, long owner, Long itemId) { + Map values = new java.util.HashMap<>(); + filter.path("must").forEach(node -> values.put(node.path("key").asText(), node.path("match").path("value"))); + assertEquals(tenant, values.get("tenant_id").asText()); + assertTrue(values.get("owner_user_id").isIntegralNumber()); + assertEquals(owner, values.get("owner_user_id").asLong()); + if (itemId != null) assertEquals(itemId.longValue(), values.get("item_id").asLong()); + } +}