fix(personal): preserve scoped hybrid retrieval
This commit is contained in:
+20
-3
@@ -1,5 +1,6 @@
|
||||
package org.dromara.aihr.personal.service;
|
||||
|
||||
import lombok.extern.slf4j.Slf4j;
|
||||
import org.dromara.aihr.personal.domain.PersonalAssistantDto.PersonalSearchRequest;
|
||||
import org.dromara.aihr.personal.domain.PersonalAssistantDto.SearchHitResponse;
|
||||
import org.dromara.aihr.personal.domain.PersonalAssistantDto.SearchScope;
|
||||
@@ -7,6 +8,7 @@ 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.dao.DataAccessException;
|
||||
import org.springframework.jdbc.core.JdbcTemplate;
|
||||
import org.springframework.jdbc.core.RowMapper;
|
||||
import org.springframework.stereotype.Service;
|
||||
@@ -22,6 +24,7 @@ import java.util.Map;
|
||||
import java.util.Optional;
|
||||
|
||||
@Service
|
||||
@Slf4j
|
||||
public class PersonalRetrievalService {
|
||||
|
||||
private static final int MAX_LIMIT = 50;
|
||||
@@ -36,7 +39,8 @@ public class PersonalRetrievalService {
|
||||
public PersonalRetrievalService(JdbcTemplate jdbcTemplate, PersonalVectorStore vectorStore,
|
||||
ObjectProvider<QueryEmbeddingProvider> embeddingProviders,
|
||||
PersonalKnowledgeProperties properties) {
|
||||
this(jdbcTemplate, vectorStore, embeddingProviders.getIfAvailable(() -> query -> Optional.empty()), properties);
|
||||
this(jdbcTemplate, vectorStore,
|
||||
embeddingProviders.orderedStream().findFirst().orElseGet(() -> query -> Optional.empty()), properties);
|
||||
}
|
||||
|
||||
public PersonalRetrievalService(JdbcTemplate jdbcTemplate, PersonalVectorStore vectorStore,
|
||||
@@ -54,24 +58,36 @@ public class PersonalRetrievalService {
|
||||
return List.of();
|
||||
}
|
||||
|
||||
List<SearchHitResponse> fulltext = fulltext(owner, validated);
|
||||
List<SearchHitResponse> fulltext;
|
||||
try {
|
||||
fulltext = fulltext(owner, validated);
|
||||
} catch (DataAccessException ex) {
|
||||
log.warn("event=personal_fulltext_fallback tenantId={} ownerUserId={} exception={}",
|
||||
owner.tenantId(), owner.userId(), ex.getClass().getSimpleName());
|
||||
fulltext = List.of();
|
||||
}
|
||||
Optional<String> vectorJson;
|
||||
try {
|
||||
vectorJson = embeddingProvider.embed(validated.query());
|
||||
} catch (RuntimeException ex) {
|
||||
log.warn("event=personal_embedding_fallback tenantId={} ownerUserId={} exception={}",
|
||||
owner.tenantId(), owner.userId(), ex.getClass().getSimpleName());
|
||||
vectorJson = Optional.empty();
|
||||
}
|
||||
if (vectorJson.isEmpty() || vectorJson.get().isBlank()) {
|
||||
return fulltext;
|
||||
}
|
||||
try {
|
||||
List<PersonalVectorStore.VectorMatch> vectorMatches = vectorStore.query(owner, vectorJson.get(), validated.limit());
|
||||
List<PersonalVectorStore.VectorMatch> vectorMatches = vectorStore.query(owner, vectorJson.get(),
|
||||
validated.limit(), validated.dateFrom(), validated.dateTo(), validated.itemIds());
|
||||
if (vectorMatches.isEmpty()) {
|
||||
return fulltext;
|
||||
}
|
||||
List<SearchHitResponse> hydrated = hydrate(owner, vectorMatches, validated);
|
||||
return mergeRrf(fulltext, hydrated, validated.limit());
|
||||
} catch (RuntimeException ex) {
|
||||
log.warn("event=personal_vector_hydration_fallback tenantId={} ownerUserId={} exception={}",
|
||||
owner.tenantId(), owner.userId(), ex.getClass().getSimpleName());
|
||||
return fulltext;
|
||||
}
|
||||
}
|
||||
@@ -225,6 +241,7 @@ public class PersonalRetrievalService {
|
||||
|
||||
private static void requireOwner(PersonalOwner owner) {
|
||||
if (owner == null || owner.tenantId() == null || owner.tenantId().isBlank() || owner.userId() <= 0) {
|
||||
log.warn("event=personal_retrieval_owner_invalid");
|
||||
throw new IllegalStateException("个人知识空间需要有效登录身份");
|
||||
}
|
||||
}
|
||||
|
||||
+125
-46
@@ -5,6 +5,7 @@ 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 lombok.extern.slf4j.Slf4j;
|
||||
import org.dromara.aihr.personal.support.PersonalKnowledgeProperties;
|
||||
import org.dromara.aihr.personal.support.PersonalOwner;
|
||||
import org.springframework.stereotype.Service;
|
||||
@@ -14,15 +15,16 @@ import java.net.http.HttpClient;
|
||||
import java.net.http.HttpRequest;
|
||||
import java.net.http.HttpResponse;
|
||||
import java.time.Duration;
|
||||
import java.time.LocalDate;
|
||||
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
|
||||
@Slf4j
|
||||
public class PersonalVectorStore {
|
||||
|
||||
private static final String UNAVAILABLE = "PERSONAL_VECTOR_STORE_UNAVAILABLE";
|
||||
@@ -32,7 +34,6 @@ public class PersonalVectorStore {
|
||||
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));
|
||||
@@ -53,40 +54,43 @@ public class PersonalVectorStore {
|
||||
public void ensureCollection(int dimension) {
|
||||
validateDimension(dimension);
|
||||
TransportResponse current = send("GET", collectionPath(), null);
|
||||
CollectionMetadata metadata;
|
||||
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");
|
||||
}
|
||||
metadata = collectionMetadata(current);
|
||||
} 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 {
|
||||
TransportResponse created = send("PUT", collectionPath(), body);
|
||||
if (!success(created.status()) && created.status() != 409) {
|
||||
log.warn("event=personal_vector_collection_create_failed status={}", created.status());
|
||||
throw unavailable();
|
||||
}
|
||||
setOrValidateDimension(dimension);
|
||||
ensurePayloadIndex("tenant_id", "keyword");
|
||||
ensurePayloadIndex("owner_user_id", "integer");
|
||||
ensurePayloadIndex("item_id", "integer");
|
||||
TransportResponse verified = send("GET", collectionPath(), null);
|
||||
if (!success(verified.status())) {
|
||||
log.warn("event=personal_vector_collection_verify_failed status={}", verified.status());
|
||||
throw unavailable();
|
||||
}
|
||||
metadata = collectionMetadata(verified);
|
||||
} else {
|
||||
log.warn("event=personal_vector_collection_read_failed status={}", current.status());
|
||||
throw unavailable();
|
||||
}
|
||||
if (metadata.dimension() != dimension) {
|
||||
throw new IllegalStateException("PERSONAL_VECTOR_DIMENSION_MISMATCH");
|
||||
}
|
||||
ensurePayloadIndex("tenant_id", "keyword", metadata.payloadSchema());
|
||||
ensurePayloadIndex("owner_user_id", "integer", metadata.payloadSchema());
|
||||
ensurePayloadIndex("item_id", "integer", metadata.payloadSchema());
|
||||
ensurePayloadIndex("captured_at", "datetime", metadata.payloadSchema());
|
||||
}
|
||||
|
||||
/**
|
||||
* 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.
|
||||
* call this method only after a real embedding provider returns a vector. Collection metadata and indexes are
|
||||
* verified before every mutation so a rejected request cannot poison process-local dimension state.
|
||||
*/
|
||||
public void upsert(PersonalOwner owner, VectorPoint point, String vectorJson) {
|
||||
requireOwner(owner);
|
||||
@@ -97,7 +101,7 @@ public class PersonalVectorStore {
|
||||
throw new IllegalArgumentException("PERSONAL_VECTOR_CAPTURED_AT_REQUIRED");
|
||||
}
|
||||
ArrayNode vector = parseVector(vectorJson);
|
||||
setOrValidateDimension(vector.size());
|
||||
ensureCollection(vector.size());
|
||||
|
||||
ObjectNode payload = objectMapper.createObjectNode();
|
||||
payload.put("tenant_id", owner.tenantId());
|
||||
@@ -105,7 +109,7 @@ public class PersonalVectorStore {
|
||||
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());
|
||||
payload.put("source_type", point.source() == null ? "" : point.source());
|
||||
payload.put("captured_at", point.capturedAt().toString());
|
||||
ObjectNode qdrantPoint = objectMapper.createObjectNode();
|
||||
qdrantPoint.put("id", point.fragmentId());
|
||||
@@ -113,16 +117,34 @@ public class PersonalVectorStore {
|
||||
qdrantPoint.set("payload", payload);
|
||||
ObjectNode body = objectMapper.createObjectNode();
|
||||
body.putArray("points").add(qdrantPoint);
|
||||
try {
|
||||
requireMutation(send("PUT", collectionPath() + "/points?wait=true", body));
|
||||
} catch (RuntimeException ex) {
|
||||
log.warn("event=personal_vector_upsert_failed tenantId={} ownerUserId={} exception={}",
|
||||
owner.tenantId(), owner.userId(), ex.getClass().getSimpleName());
|
||||
throw ex;
|
||||
}
|
||||
}
|
||||
|
||||
public List<VectorMatch> query(PersonalOwner owner, String vectorJson, int limit) {
|
||||
return query(owner, vectorJson, limit, null, null, List.of());
|
||||
}
|
||||
|
||||
public List<VectorMatch> query(PersonalOwner owner, String vectorJson, int limit, LocalDate dateFrom,
|
||||
LocalDate dateTo, List<Long> itemIds) {
|
||||
requireOwner(owner);
|
||||
ArrayNode vector = parseVector(vectorJson);
|
||||
setOrValidateDimension(vector.size());
|
||||
validateDimension(vector.size());
|
||||
if (dateFrom != null && dateTo != null && dateFrom.isAfter(dateTo)) {
|
||||
throw new IllegalArgumentException("PERSONAL_VECTOR_DATE_RANGE_INVALID");
|
||||
}
|
||||
List<Long> scopedItems = itemIds == null ? List.of() : itemIds.stream().distinct().toList();
|
||||
if (scopedItems.size() > 100 || scopedItems.stream().anyMatch(id -> id == null || id <= 0)) {
|
||||
throw new IllegalArgumentException("PERSONAL_VECTOR_ITEM_SCOPE_INVALID");
|
||||
}
|
||||
ObjectNode body = objectMapper.createObjectNode();
|
||||
body.set("query", vector);
|
||||
body.set("filter", ownerFilter(owner, null));
|
||||
body.set("filter", scopedFilter(owner, dateFrom, dateTo, scopedItems));
|
||||
body.put("limit", Math.max(1, Math.min(50, limit)));
|
||||
body.put("with_payload", true);
|
||||
body.put("with_vector", false);
|
||||
@@ -130,9 +152,13 @@ public class PersonalVectorStore {
|
||||
try {
|
||||
response = send("POST", collectionPath() + "/points/query", body);
|
||||
} catch (IllegalStateException ex) {
|
||||
log.warn("event=personal_vector_query_fallback tenantId={} ownerUserId={} exception={}",
|
||||
owner.tenantId(), owner.userId(), ex.getClass().getSimpleName());
|
||||
return List.of();
|
||||
}
|
||||
if (!success(response.status())) {
|
||||
log.warn("event=personal_vector_query_fallback tenantId={} ownerUserId={} status={}",
|
||||
owner.tenantId(), owner.userId(), response.status());
|
||||
return List.of();
|
||||
}
|
||||
try {
|
||||
@@ -151,6 +177,8 @@ public class PersonalVectorStore {
|
||||
}
|
||||
return List.copyOf(matches);
|
||||
} catch (Exception ex) {
|
||||
log.warn("event=personal_vector_query_malformed tenantId={} ownerUserId={} exception={}",
|
||||
owner.tenantId(), owner.userId(), ex.getClass().getSimpleName());
|
||||
return List.of();
|
||||
}
|
||||
}
|
||||
@@ -161,30 +189,84 @@ public class PersonalVectorStore {
|
||||
throw new IllegalArgumentException("PERSONAL_ITEM_ID_INVALID");
|
||||
}
|
||||
ObjectNode body = objectMapper.createObjectNode();
|
||||
body.set("filter", ownerFilter(owner, itemId));
|
||||
body.set("filter", scopedFilter(owner, null, null, List.of(itemId)));
|
||||
try {
|
||||
TransportResponse response = send("POST", collectionPath() + "/points/delete?wait=true", body);
|
||||
if (response.status() != 404) {
|
||||
requireMutation(response);
|
||||
}
|
||||
} catch (RuntimeException ex) {
|
||||
log.warn("event=personal_vector_delete_failed tenantId={} ownerUserId={} exception={}",
|
||||
owner.tenantId(), owner.userId(), ex.getClass().getSimpleName());
|
||||
throw ex;
|
||||
}
|
||||
}
|
||||
|
||||
private void ensurePayloadIndex(String field, String schema) {
|
||||
private void ensurePayloadIndex(String field, String schema, JsonNode payloadSchema) {
|
||||
if (payloadIndexMatches(payloadSchema, field, schema)) {
|
||||
return;
|
||||
}
|
||||
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) {
|
||||
if (success(response.status())) {
|
||||
return;
|
||||
}
|
||||
if (response.status() == 409) {
|
||||
TransportResponse verified = send("GET", collectionPath(), null);
|
||||
if (success(verified.status())
|
||||
&& payloadIndexMatches(collectionMetadata(verified).payloadSchema(), field, schema)) {
|
||||
return;
|
||||
}
|
||||
}
|
||||
log.warn("event=personal_vector_payload_index_failed field={} status={}", field, response.status());
|
||||
throw unavailable();
|
||||
}
|
||||
|
||||
private boolean payloadIndexMatches(JsonNode payloadSchema, String field, String schema) {
|
||||
JsonNode entry = payloadSchema.path(field);
|
||||
String actual = entry.isTextual() ? entry.asText() : entry.path("data_type").asText("");
|
||||
return schema.equalsIgnoreCase(actual);
|
||||
}
|
||||
|
||||
private CollectionMetadata collectionMetadata(TransportResponse response) {
|
||||
try {
|
||||
JsonNode result = objectMapper.readTree(response.body()).path("result");
|
||||
JsonNode size = result.path("config").path("params").path("vectors").path("size");
|
||||
if (!size.canConvertToInt() || size.asInt() <= 0) {
|
||||
throw unavailable();
|
||||
}
|
||||
return new CollectionMetadata(size.asInt(), result.path("payload_schema"));
|
||||
} catch (Exception ex) {
|
||||
log.warn("event=personal_vector_transport_failed exception={}", ex.getClass().getSimpleName());
|
||||
throw unavailable();
|
||||
}
|
||||
}
|
||||
|
||||
private ObjectNode ownerFilter(PersonalOwner owner, Long itemId) {
|
||||
private ObjectNode scopedFilter(PersonalOwner owner, LocalDate dateFrom, LocalDate dateTo, List<Long> itemIds) {
|
||||
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));
|
||||
if (itemIds != null && !itemIds.isEmpty()) {
|
||||
ObjectNode condition = objectMapper.createObjectNode();
|
||||
condition.put("key", "item_id");
|
||||
ArrayNode any = condition.putObject("match").putArray("any");
|
||||
itemIds.forEach(any::add);
|
||||
must.add(condition);
|
||||
}
|
||||
if (dateFrom != null || dateTo != null) {
|
||||
ObjectNode condition = objectMapper.createObjectNode();
|
||||
condition.put("key", "captured_at");
|
||||
ObjectNode range = condition.putObject("range");
|
||||
if (dateFrom != null) {
|
||||
range.put("gte", dateFrom.atStartOfDay().toString());
|
||||
}
|
||||
if (dateTo != null) {
|
||||
range.put("lt", dateTo.plusDays(1).atStartOfDay().toString());
|
||||
}
|
||||
must.add(condition);
|
||||
}
|
||||
return filter;
|
||||
}
|
||||
@@ -220,18 +302,6 @@ public class PersonalVectorStore {
|
||||
}
|
||||
}
|
||||
|
||||
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");
|
||||
@@ -240,6 +310,7 @@ public class PersonalVectorStore {
|
||||
|
||||
private void requireOwner(PersonalOwner owner) {
|
||||
if (owner == null || owner.tenantId() == null || owner.tenantId().isBlank() || owner.userId() <= 0) {
|
||||
log.warn("event=personal_vector_owner_invalid");
|
||||
throw new IllegalStateException("个人知识空间需要有效登录身份");
|
||||
}
|
||||
}
|
||||
@@ -265,6 +336,7 @@ public class PersonalVectorStore {
|
||||
|
||||
private void requireMutation(TransportResponse response) {
|
||||
if (!success(response.status())) {
|
||||
log.warn("event=personal_vector_mutation_rejected status={}", response.status());
|
||||
throw unavailable();
|
||||
}
|
||||
}
|
||||
@@ -300,7 +372,7 @@ public class PersonalVectorStore {
|
||||
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()))
|
||||
HttpRequest.Builder builder = HttpRequest.newBuilder(endpointUri(base, request.path()))
|
||||
.timeout(Duration.ofSeconds(seconds));
|
||||
request.headers().forEach(builder::header);
|
||||
builder.method(request.method(), request.body().isEmpty() ? HttpRequest.BodyPublishers.noBody()
|
||||
@@ -310,6 +382,10 @@ public class PersonalVectorStore {
|
||||
};
|
||||
}
|
||||
|
||||
static URI endpointUri(URI base, String path) {
|
||||
return URI.create(base.toString() + path);
|
||||
}
|
||||
|
||||
private static String firstNonBlank(String... values) {
|
||||
for (String value : values) {
|
||||
if (value != null && !value.isBlank()) {
|
||||
@@ -336,4 +412,7 @@ public class PersonalVectorStore {
|
||||
|
||||
public record VectorMatch(long fragmentId, double score) {
|
||||
}
|
||||
|
||||
private record CollectionMetadata(int dimension, JsonNode payloadSchema) {
|
||||
}
|
||||
}
|
||||
|
||||
+54
-2
@@ -10,12 +10,16 @@ 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.beans.factory.ObjectProvider;
|
||||
import org.springframework.dao.DataAccessResourceFailureException;
|
||||
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 java.util.Optional;
|
||||
import java.util.stream.Stream;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.*;
|
||||
import static org.mockito.ArgumentMatchers.*;
|
||||
@@ -85,6 +89,8 @@ class PersonalRetrievalServiceTest {
|
||||
assertTrue(hydration.contains("i.captured_at >= ?"));
|
||||
assertTrue(hydration.contains("i.captured_at < ?"));
|
||||
assertTrue(hydration.contains("i.id in (?)"));
|
||||
verify(vectors).query(any(), eq("[0.1,0.2]"), eq(10), eq(LocalDate.of(2026, 1, 1)),
|
||||
eq(LocalDate.of(2026, 1, 31)), eq(List.of(55L)));
|
||||
}
|
||||
|
||||
@Test
|
||||
@@ -92,13 +98,58 @@ class PersonalRetrievalServiceTest {
|
||||
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"));
|
||||
when(broken.query(any(), anyString(), anyInt(), nullable(LocalDate.class), nullable(LocalDate.class), anyList()))
|
||||
.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));
|
||||
}
|
||||
|
||||
@Test
|
||||
void fulltextFailureStillAllowsScopedVectorHydration() {
|
||||
JdbcTemplate jdbc = mock(JdbcTemplate.class);
|
||||
when(jdbc.query(anyString(), any(RowMapper.class), any(Object[].class)))
|
||||
.thenThrow(new DataAccessResourceFailureException("mysql fulltext unavailable"))
|
||||
.thenReturn(List.of(hit("20", "Vector")));
|
||||
PersonalVectorStore vectors = vectorStore(List.of(new PersonalVectorStore.VectorMatch(20, .9)));
|
||||
PersonalRetrievalService service = service(jdbc, query -> Optional.of("[1,2]"), vectors);
|
||||
|
||||
List<SearchHitResponse> hits = service.search(new PersonalOwner("t", 1, null),
|
||||
new PersonalSearchRequest("q", null, null, null, null, 10));
|
||||
|
||||
assertEquals(List.of("20"), hits.stream().map(SearchHitResponse::sourceId).toList());
|
||||
verify(jdbc, times(2)).query(anyString(), any(RowMapper.class), any(Object[].class));
|
||||
}
|
||||
|
||||
@Test
|
||||
void springConstructorUsesNoopForZeroProvidersAndOrderedFirstForMultiple() {
|
||||
JdbcTemplate jdbc = mock(JdbcTemplate.class);
|
||||
when(jdbc.query(anyString(), any(RowMapper.class), any(Object[].class))).thenReturn(List.of());
|
||||
PersonalVectorStore vectors = mock(PersonalVectorStore.class);
|
||||
PersonalKnowledgeProperties properties = new PersonalKnowledgeProperties();
|
||||
|
||||
ObjectProvider<PersonalRetrievalService.QueryEmbeddingProvider> none = mock(ObjectProvider.class);
|
||||
when(none.orderedStream()).thenReturn(Stream.empty());
|
||||
PersonalRetrievalService noProvider = new PersonalRetrievalService(jdbc, vectors, none, properties);
|
||||
assertDoesNotThrow(() -> noProvider.search(new PersonalOwner("t", 1, null),
|
||||
new PersonalSearchRequest("q", null, null, null, null, 10)));
|
||||
verifyNoInteractions(vectors);
|
||||
|
||||
reset(jdbc, vectors);
|
||||
when(jdbc.query(anyString(), any(RowMapper.class), any(Object[].class))).thenReturn(List.of());
|
||||
PersonalRetrievalService.QueryEmbeddingProvider first = mock(PersonalRetrievalService.QueryEmbeddingProvider.class);
|
||||
PersonalRetrievalService.QueryEmbeddingProvider second = mock(PersonalRetrievalService.QueryEmbeddingProvider.class);
|
||||
when(first.embed("q")).thenReturn(Optional.empty());
|
||||
ObjectProvider<PersonalRetrievalService.QueryEmbeddingProvider> multiple = mock(ObjectProvider.class);
|
||||
when(multiple.orderedStream()).thenReturn(Stream.of(first, second));
|
||||
PersonalRetrievalService selected = new PersonalRetrievalService(jdbc, vectors, multiple, properties);
|
||||
selected.search(new PersonalOwner("t", 1, null),
|
||||
new PersonalSearchRequest("q", null, null, null, null, 10));
|
||||
verify(first).embed("q");
|
||||
verifyNoInteractions(second);
|
||||
}
|
||||
|
||||
private PersonalRetrievalService service(JdbcTemplate jdbc, PersonalRetrievalService.QueryEmbeddingProvider provider,
|
||||
PersonalVectorStore vectors) {
|
||||
PersonalKnowledgeProperties properties = new PersonalKnowledgeProperties();
|
||||
@@ -107,7 +158,8 @@ class PersonalRetrievalServiceTest {
|
||||
|
||||
private PersonalVectorStore vectorStore(List<PersonalVectorStore.VectorMatch> matches) {
|
||||
PersonalVectorStore vectors = mock(PersonalVectorStore.class);
|
||||
when(vectors.query(any(), anyString(), anyInt())).thenReturn(matches);
|
||||
when(vectors.query(any(), anyString(), anyInt(), nullable(LocalDate.class), nullable(LocalDate.class), anyList()))
|
||||
.thenReturn(matches);
|
||||
return vectors;
|
||||
}
|
||||
|
||||
|
||||
+100
-17
@@ -1,14 +1,15 @@
|
||||
package org.dromara.aihr.personal;
|
||||
package org.dromara.aihr.personal.service;
|
||||
|
||||
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.time.LocalDate;
|
||||
import java.net.URI;
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
@@ -28,27 +29,32 @@ class PersonalVectorStoreTest {
|
||||
: ok("{}"));
|
||||
PersonalOwner owner = new PersonalOwner("000001", 42, "ext");
|
||||
|
||||
assertEquals(99, store.query(owner, "[0.1,0.2]", 100).get(0).fragmentId());
|
||||
assertEquals(99, store.query(owner, "[0.1,0.2]", 100,
|
||||
LocalDate.of(2026, 7, 1), LocalDate.of(2026, 7, 2), List.of(7L, 8L)).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);
|
||||
assertFilter(query.path("filter"), "000001", 42, List.of(7L, 8L));
|
||||
Map<String, JsonNode> conditions = conditions(query.path("filter"));
|
||||
assertEquals("2026-07-01T00:00", conditions.get("captured_at").path("range").path("gte").asText());
|
||||
assertEquals("2026-07-03T00:00", conditions.get("captured_at").path("range").path("lt").asText());
|
||||
JsonNode delete = mapper.readTree(seen.get(1).body());
|
||||
assertFilter(delete.path("filter"), "000001", 42, 7L);
|
||||
assertFilter(delete.path("filter"), "000001", 42, List.of(7L));
|
||||
assertFalse(seen.get(0).body().contains("999"), "Qdrant payload owner must not influence authorization filter");
|
||||
}
|
||||
|
||||
@Test
|
||||
void upsertUsesPersonalCollectionAndOwnerScopedPayload() throws Exception {
|
||||
List<PersonalVectorStore.TransportRequest> seen = new ArrayList<>();
|
||||
PersonalVectorStore store = fixture(seen, request -> ok("{}"));
|
||||
PersonalVectorStore store = fixture(seen, request -> request.method().equals("GET")
|
||||
? ok(collectionBody(2, true)) : 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 body = mapper.readTree(seen.get(seen.size() - 1).body());
|
||||
JsonNode point = body.path("points").get(0);
|
||||
assertEquals(13, point.path("id").asLong());
|
||||
assertEquals("t-1", point.path("payload").path("tenant_id").asText());
|
||||
@@ -56,6 +62,8 @@ class PersonalVectorStoreTest {
|
||||
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());
|
||||
assertEquals("file", point.path("payload").path("source_type").asText());
|
||||
assertFalse(point.path("payload").has("source"));
|
||||
assertEquals("2026-07-01T09:00", point.path("payload").path("captured_at").asText());
|
||||
assertFalse(point.path("payload").has("content"));
|
||||
}
|
||||
@@ -77,8 +85,10 @@ class PersonalVectorStoreTest {
|
||||
@Test
|
||||
void collectionCreationAndPayloadIndexesAreStable() throws Exception {
|
||||
List<PersonalVectorStore.TransportRequest> seen = new ArrayList<>();
|
||||
java.util.concurrent.atomic.AtomicInteger gets = new java.util.concurrent.atomic.AtomicInteger();
|
||||
PersonalVectorStore store = fixture(seen, request -> request.method().equals("GET")
|
||||
? new PersonalVectorStore.TransportResponse(404, "") : ok("{}"));
|
||||
? (gets.getAndIncrement() == 0 ? new PersonalVectorStore.TransportResponse(404, "")
|
||||
: ok(collectionBody(2, false))) : ok("{}"));
|
||||
|
||||
store.ensureCollection(2);
|
||||
|
||||
@@ -86,13 +96,68 @@ class PersonalVectorStoreTest {
|
||||
assertEquals(2, mapper.readTree(seen.get(1).body()).path("vectors").path("size").asInt());
|
||||
List<String> 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);
|
||||
assertEquals(List.of("tenant_id", "owner_user_id", "item_id", "captured_at"), indexFields);
|
||||
assertTrue(seen.stream().filter(r -> r.method().equals("GET")).count() >= 2,
|
||||
"collection creation must be followed by metadata verification");
|
||||
}
|
||||
|
||||
@Test
|
||||
void concurrentCollectionAndIndexCreationRereadsMetadata() {
|
||||
List<PersonalVectorStore.TransportRequest> seen = new ArrayList<>();
|
||||
java.util.concurrent.atomic.AtomicInteger gets = new java.util.concurrent.atomic.AtomicInteger();
|
||||
PersonalVectorStore store = fixture(seen, request -> {
|
||||
if (request.method().equals("GET")) {
|
||||
return gets.getAndIncrement() == 0 ? new PersonalVectorStore.TransportResponse(404, "")
|
||||
: ok(collectionBody(2, true));
|
||||
}
|
||||
if (request.path().equals("/collections/aihr_personal_knowledge")) {
|
||||
return new PersonalVectorStore.TransportResponse(409, "already exists");
|
||||
}
|
||||
return ok("{}");
|
||||
});
|
||||
assertDoesNotThrow(() -> store.ensureCollection(2));
|
||||
assertEquals(2, gets.get());
|
||||
|
||||
List<PersonalVectorStore.TransportRequest> indexSeen = new ArrayList<>();
|
||||
java.util.concurrent.atomic.AtomicInteger indexGets = new java.util.concurrent.atomic.AtomicInteger();
|
||||
PersonalVectorStore indexStore = fixture(indexSeen, request -> {
|
||||
if (request.method().equals("GET")) {
|
||||
return ok(collectionBody(2, indexGets.getAndIncrement() > 0));
|
||||
}
|
||||
return new PersonalVectorStore.TransportResponse(409, "already exists");
|
||||
});
|
||||
assertDoesNotThrow(() -> indexStore.ensureCollection(2));
|
||||
assertEquals(5, indexGets.get(), "every concurrent index conflict must reread and verify payload schema");
|
||||
}
|
||||
|
||||
@Test
|
||||
void rejectedFirstQueryDoesNotPoisonLaterVectorDimension() {
|
||||
List<PersonalVectorStore.TransportRequest> seen = new ArrayList<>();
|
||||
java.util.concurrent.atomic.AtomicInteger posts = new java.util.concurrent.atomic.AtomicInteger();
|
||||
PersonalVectorStore store = fixture(seen, request -> posts.getAndIncrement() == 0
|
||||
? new PersonalVectorStore.TransportResponse(400, "wrong dimension")
|
||||
: ok("{\"result\":{\"points\":[{\"score\":0.8,\"payload\":{\"fragment_id\":8}}]}}"));
|
||||
PersonalOwner owner = new PersonalOwner("t", 1, null);
|
||||
|
||||
assertTrue(store.query(owner, "[1,2,3]", 5).isEmpty());
|
||||
assertEquals(8, store.query(owner, "[1,2]", 5).get(0).fragmentId());
|
||||
assertEquals(2, posts.get());
|
||||
}
|
||||
|
||||
@Test
|
||||
void preservesConfiguredQdrantBasePathPrefix() {
|
||||
PersonalKnowledgeProperties properties = properties();
|
||||
properties.setQdrantUrl("https://qdrant.example/internal/api/");
|
||||
URI base = URI.create(properties.getQdrantUrl().replaceFirst("/$", ""));
|
||||
assertEquals(URI.create("https://qdrant.example/internal/api/collections/personal"),
|
||||
PersonalVectorStore.endpointUri(base, "/collections/personal"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void rejectsMalformedNonFiniteAndDimensionMismatchBeforeTransport() {
|
||||
List<PersonalVectorStore.TransportRequest> seen = new ArrayList<>();
|
||||
PersonalVectorStore store = fixture(seen, request -> ok("{}"));
|
||||
PersonalVectorStore store = fixture(seen, request -> request.method().equals("GET")
|
||||
? ok(collectionBody(2, true)) : ok("{}"));
|
||||
PersonalOwner owner = new PersonalOwner("t", 1, null);
|
||||
var point = new PersonalVectorStore.VectorPoint(1, 2, 3, "text", LocalDateTime.now());
|
||||
|
||||
@@ -100,7 +165,7 @@ class PersonalVectorStoreTest {
|
||||
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.upsert(owner, point, "[1,2,3]"));
|
||||
assertThrows(IllegalStateException.class, () -> store.query(null, "[1,2]", 5));
|
||||
}
|
||||
|
||||
@@ -149,12 +214,30 @@ class PersonalVectorStoreTest {
|
||||
try { return mapper.readTree(body); } catch (Exception e) { throw new AssertionError(e); }
|
||||
}
|
||||
|
||||
private void assertFilter(JsonNode filter, String tenant, long owner, Long itemId) {
|
||||
private void assertFilter(JsonNode filter, String tenant, long owner, List<Long> itemIds) {
|
||||
Map<String, JsonNode> values = conditions(filter);
|
||||
assertEquals(tenant, values.get("tenant_id").path("match").path("value").asText());
|
||||
assertTrue(values.get("owner_user_id").path("match").path("value").isIntegralNumber());
|
||||
assertEquals(owner, values.get("owner_user_id").path("match").path("value").asLong());
|
||||
if (itemIds != null) {
|
||||
List<Long> actual = new ArrayList<>();
|
||||
values.get("item_id").path("match").path("any").forEach(v -> actual.add(v.asLong()));
|
||||
assertEquals(itemIds, actual);
|
||||
}
|
||||
}
|
||||
|
||||
private Map<String, JsonNode> conditions(JsonNode filter) {
|
||||
Map<String, JsonNode> 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());
|
||||
filter.path("must").forEach(node -> values.put(node.path("key").asText(), node));
|
||||
return values;
|
||||
}
|
||||
|
||||
private String collectionBody(int dimension, boolean indexes) {
|
||||
String schema = indexes ? "\"payload_schema\":{" +
|
||||
"\"tenant_id\":{\"data_type\":\"keyword\"}," +
|
||||
"\"owner_user_id\":{\"data_type\":\"integer\"}," +
|
||||
"\"item_id\":{\"data_type\":\"integer\"}," +
|
||||
"\"captured_at\":{\"data_type\":\"datetime\"}}" : "\"payload_schema\":{}";
|
||||
return "{\"result\":{\"config\":{\"params\":{\"vectors\":{\"size\":" + dimension + "}}}," + schema + "}}";
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user