diff --git a/components/graph-rag-core/src/main/java/com/orgmemory/graphrag/observability/GraphRagEventSink.java b/components/graph-rag-core/src/main/java/com/orgmemory/graphrag/observability/GraphRagEventSink.java index 7b2a9b77..badc62bb 100644 --- a/components/graph-rag-core/src/main/java/com/orgmemory/graphrag/observability/GraphRagEventSink.java +++ b/components/graph-rag-core/src/main/java/com/orgmemory/graphrag/observability/GraphRagEventSink.java @@ -267,6 +267,17 @@ enum Stage { AUTHORIZE, PREPARE_QUERY, RETRIEVE_SNAPSHOT, + QUERY_SEARCH_ENTITIES, + QUERY_SEARCH_RELATIONS, + QUERY_SEARCH_CHUNKS, + QUERY_EXPAND_ENTITY_IDS, + QUERY_LOAD_INCIDENT_RELATIONS, + QUERY_LOAD_ENTITY_CONTRIBUTIONS, + QUERY_LOAD_RELATION_CONTRIBUTIONS, + QUERY_LOAD_VISIBLE_ENTITY_DEGREES, + QUERY_LOAD_VISIBLE_RELATION_WEIGHTS, + QUERY_RANK_CHUNKS, + QUERY_LOAD_CHUNKS, RETRIEVE, RERANK, ASSEMBLE_CONTEXT diff --git a/components/graph-rag-core/src/main/java/com/orgmemory/graphrag/query/LightRagQueryEngine.java b/components/graph-rag-core/src/main/java/com/orgmemory/graphrag/query/LightRagQueryEngine.java index 8d9dbc33..e0005b49 100644 --- a/components/graph-rag-core/src/main/java/com/orgmemory/graphrag/query/LightRagQueryEngine.java +++ b/components/graph-rag-core/src/main/java/com/orgmemory/graphrag/query/LightRagQueryEngine.java @@ -24,6 +24,8 @@ import java.util.Set; import java.util.UUID; import java.util.function.Function; +import java.util.function.Supplier; +import java.util.function.ToIntFunction; import java.util.stream.Collectors; /** @@ -117,8 +119,16 @@ public LightRagPreparedQuery prepare(LightRagQueryRequest request) { public LightRagQueryResult executePrepared( LightRagQueryRequest request, LightRagPreparedQuery prepared) { + return executePrepared(request, prepared, QueryOperationObserver.NO_OP); + } + + public LightRagQueryResult executePrepared( + LightRagQueryRequest request, + LightRagPreparedQuery prepared, + QueryOperationObserver observer) { Objects.requireNonNull(request, "request"); Objects.requireNonNull(prepared, "prepared").requireMatches(request); + Objects.requireNonNull(observer, "observer"); if (request.options().mode() == LightRagQueryMode.BYPASS) { return bypass(request); } @@ -131,10 +141,10 @@ public LightRagQueryResult executePrepared( } Branch local = request.options().mode().usesEntitySeeds() - ? localBranch(request, prepared.lowLevelEmbedding()) + ? localBranch(request, prepared.lowLevelEmbedding(), observer) : Branch.empty(); Branch global = request.options().mode().usesRelationSeeds() - ? globalBranch(request, prepared.highLevelEmbedding()) + ? globalBranch(request, prepared.highLevelEmbedding(), observer) : Branch.empty(); List> entities = @@ -147,7 +157,8 @@ public LightRagQueryResult executePrepared( entities.stream().map(RankedItem::value).toList(), LightRagQueryResult.Origin.ENTITY, Set.of(), - prepared.queryEmbedding()); + prepared.queryEmbedding(), + observer); Set entityChunkIds = entityChunks.stream().map(state -> state.chunk().id()).collect(Collectors.toSet()); List relationChunks = supportChunks( @@ -155,9 +166,10 @@ public LightRagQueryResult executePrepared( relations.stream().map(RankedItem::value).toList(), LightRagQueryResult.Origin.RELATION, entityChunkIds, - prepared.queryEmbedding()); + prepared.queryEmbedding(), + observer); List vectorChunks = request.options().mode().usesChunkSeeds() - ? vectorChunks(request, prepared.queryEmbedding()) + ? vectorChunks(request, prepared.queryEmbedding(), observer) : List.of(); List chunks = interleaveChunks(vectorChunks, entityChunks, relationChunks); @@ -276,13 +288,20 @@ private EmbeddingPlan embed(LightRagQueryRequest request, KeywordPlan keywords) vector(offsets, vectors, Purpose.HIGH_LEVEL)); } - private Branch localBranch(LightRagQueryRequest request, FloatVector lowLevelVector) { + private Branch localBranch( + LightRagQueryRequest request, + FloatVector lowLevelVector, + QueryOperationObserver observer) { if (lowLevelVector == null) { return Branch.empty(); } var search = vectorSearch(request, lowLevelVector, request.options().topK()); - List> seeds = - projection.searchEntities(request.scope(), request.snapshot(), search); + List> seeds = measure( + observer, + QueryOperation.SEARCH_ENTITIES, + search.limit(), + () -> projection.searchEntities(request.scope(), request.snapshot(), search), + Collection::size); if (seeds.isEmpty()) { return Branch.empty(); } @@ -290,18 +309,28 @@ private Branch localBranch(LightRagQueryRequest request, FloatVector lowLevelVec .map(item -> item.value().id()) .collect(Collectors.toCollection(LinkedHashSet::new)); if (request.options().maximumGraphDepth() > 0) { - entityIds.addAll(projection.expandEntityIds( - request.scope(), - request.snapshot(), - entityIds, - request.options().maximumGraphDepth(), - request.options().topK() * 4)); - } - List incident = projection.loadIncidentRelations( - request.scope(), - request.snapshot(), - entityIds, - request.options().topK() * 4); + entityIds.addAll(measure( + observer, + QueryOperation.EXPAND_ENTITY_IDS, + entityIds.size(), + () -> projection.expandEntityIds( + request.scope(), + request.snapshot(), + entityIds, + request.options().maximumGraphDepth(), + request.options().topK() * 4), + Collection::size)); + } + List incident = measure( + observer, + QueryOperation.LOAD_INCIDENT_RELATIONS, + entityIds.size(), + () -> projection.loadIncidentRelations( + request.scope(), + request.snapshot(), + entityIds, + request.options().topK() * 4), + Collection::size); incident.forEach(relation -> { entityIds.add(relation.sourceEntityId()); entityIds.add(relation.targetEntityId()); @@ -309,7 +338,8 @@ private Branch localBranch(LightRagQueryRequest request, FloatVector lowLevelVec PermissionScopedGraphView view = scopedView( request, entityIds, - incident.stream().map(CanonicalRelation::id).toList()); + incident.stream().map(CanonicalRelation::id).toList(), + observer); Map entityViews = index(view.entities(), item -> item.entity().id()); Map relationViews = @@ -319,8 +349,13 @@ private Branch localBranch(LightRagQueryRequest request, FloatVector lowLevelVec RankedItem::score, Math::max, LinkedHashMap::new)); - Map degrees = projection.loadVisibleEntityDegrees( - request.scope(), request.snapshot(), entityIds); + Map degrees = measure( + observer, + QueryOperation.LOAD_VISIBLE_ENTITY_DEGREES, + entityIds.size(), + () -> projection.loadVisibleEntityDegrees( + request.scope(), request.snapshot(), entityIds), + Map::size); List> orderedEntities = entityViews.values().stream() .sorted(Comparator @@ -337,17 +372,24 @@ private Branch localBranch(LightRagQueryRequest request, FloatVector lowLevelVec seedScores.getOrDefault(item.entity().id(), 0.0))) .toList(); List> orderedRelations = - rankIncidentRelations(request, relationViews.values(), degrees); + rankIncidentRelations(request, relationViews.values(), degrees, observer); return new Branch(orderedEntities, orderedRelations, seeds.size()); } - private Branch globalBranch(LightRagQueryRequest request, FloatVector highLevelVector) { + private Branch globalBranch( + LightRagQueryRequest request, + FloatVector highLevelVector, + QueryOperationObserver observer) { if (highLevelVector == null) { return Branch.empty(); } var search = vectorSearch(request, highLevelVector, request.options().topK()); - List> seeds = - projection.searchRelations(request.scope(), request.snapshot(), search); + List> seeds = measure( + observer, + QueryOperation.SEARCH_RELATIONS, + search.limit(), + () -> projection.searchRelations(request.scope(), request.snapshot(), search), + Collection::size); if (seeds.isEmpty()) { return Branch.empty(); } @@ -359,7 +401,8 @@ private Branch globalBranch(LightRagQueryRequest request, FloatVector highLevelV entityIds.add(seed.value().sourceEntityId()); entityIds.add(seed.value().targetEntityId()); }); - PermissionScopedGraphView view = scopedView(request, entityIds, relationIds); + PermissionScopedGraphView view = + scopedView(request, entityIds, relationIds, observer); Map entityViews = index(view.entities(), item -> item.entity().id()); Map relationViews = @@ -400,23 +443,40 @@ private Branch globalBranch(LightRagQueryRequest request, FloatVector highLevelV private PermissionScopedGraphView scopedView( LightRagQueryRequest request, Collection entityIds, - Collection relationIds) { - List entities = projection.loadEntityContributions( - request.scope(), request.snapshot(), entityIds); - List relations = projection.loadRelationContributions( - request.scope(), request.snapshot(), relationIds); + Collection relationIds, + QueryOperationObserver observer) { + List entities = measure( + observer, + QueryOperation.LOAD_ENTITY_CONTRIBUTIONS, + entityIds.size(), + () -> projection.loadEntityContributions( + request.scope(), request.snapshot(), entityIds), + Collection::size); + List relations = measure( + observer, + QueryOperation.LOAD_RELATION_CONTRIBUTIONS, + relationIds.size(), + () -> projection.loadRelationContributions( + request.scope(), request.snapshot(), relationIds), + Collection::size); return PermissionScopedGraphMerger.merge(request.scope(), entities, relations); } private List> rankIncidentRelations( LightRagQueryRequest request, Collection relations, - Map entityDegrees) { + Map entityDegrees, + QueryOperationObserver observer) { List relationIds = relations.stream() .map(item -> item.relation().id()) .toList(); - Map weights = projection.loadVisibleRelationWeights( - request.scope(), request.snapshot(), relationIds); + Map weights = measure( + observer, + QueryOperation.LOAD_VISIBLE_RELATION_WEIGHTS, + relationIds.size(), + () -> projection.loadVisibleRelationWeights( + request.scope(), request.snapshot(), relationIds), + Map::size); return relations.stream() .sorted(Comparator .comparingLong((PermissionScopedGraphView.RelationView item) -> @@ -441,7 +501,8 @@ private List supportChunks( List views, LightRagQueryResult.Origin origin, Set excluded, - FloatVector queryVector) { + FloatVector queryVector, + QueryOperationObserver observer) { List> groups = new ArrayList<>(); for (Object view : views) { List ids = evidenceChunkIds(view).stream() @@ -475,11 +536,18 @@ private List supportChunks( int limit = Math.max(1, request.options().relatedChunkNumber() * deduplicated.size() / 2); try { - ranked = projection.rankChunks( - request.scope(), - request.snapshot(), - vectorSearch(request, queryVector, limit), - deduplicated.stream().flatMap(Collection::stream).toList()); + List candidates = + deduplicated.stream().flatMap(Collection::stream).toList(); + ranked = measure( + observer, + QueryOperation.RANK_CHUNKS, + candidates.size(), + () -> projection.rankChunks( + request.scope(), + request.snapshot(), + vectorSearch(request, queryVector, limit), + candidates), + Collection::size); } catch (RuntimeException providerFailure) { // Provider diagnostics belong to the imperative shell. Core preserves // authorized ordering by falling back to weighted polling. @@ -496,7 +564,13 @@ private List supportChunks( deduplicated, request.options().relatedChunkNumber(), 1); } Map chunks = index( - projection.loadChunks(request.scope(), request.snapshot(), selectedIds), + measure( + observer, + QueryOperation.LOAD_CHUNKS, + selectedIds.size(), + () -> projection.loadChunks( + request.scope(), request.snapshot(), selectedIds), + Collection::size), AuthorizedQueryProjection.Chunk::id); List result = new ArrayList<>(); int order = 1; @@ -516,15 +590,21 @@ private List supportChunks( } private List vectorChunks( - LightRagQueryRequest request, FloatVector queryVector) { + LightRagQueryRequest request, + FloatVector queryVector, + QueryOperationObserver observer) { if (queryVector == null) { return List.of(); } - List> hits = - projection.searchChunks( + List> hits = measure( + observer, + QueryOperation.SEARCH_CHUNKS, + request.options().chunkTopK(), + () -> projection.searchChunks( request.scope(), request.snapshot(), - vectorSearch(request, queryVector, request.options().chunkTopK())); + vectorSearch(request, queryVector, request.options().chunkTopK())), + Collection::size); List result = new ArrayList<>(); int order = 1; for (RankedItem hit : hits) { @@ -998,4 +1078,78 @@ private static Duration elapsed(long startedAt) { 0, System.nanoTime() - startedAt)); } + + private static T measure( + QueryOperationObserver observer, + QueryOperation operation, + int inputCount, + Supplier supplier, + ToIntFunction outputCount) { + long startedAt = System.nanoTime(); + try { + T result = supplier.get(); + observer.observe(new QueryOperationMeasurement( + operation, + QueryOperationOutcome.SUCCEEDED, + elapsed(startedAt), + inputCount, + outputCount.applyAsInt(result))); + return result; + } catch (RuntimeException failure) { + observer.observe(new QueryOperationMeasurement( + operation, + QueryOperationOutcome.FAILED, + elapsed(startedAt), + inputCount, + 0)); + throw failure; + } + } + + /** Bounded storage-operation names; values are safe metric/span dimensions. */ + public enum QueryOperation { + SEARCH_ENTITIES, + SEARCH_RELATIONS, + SEARCH_CHUNKS, + EXPAND_ENTITY_IDS, + LOAD_INCIDENT_RELATIONS, + LOAD_ENTITY_CONTRIBUTIONS, + LOAD_RELATION_CONTRIBUTIONS, + LOAD_VISIBLE_ENTITY_DEGREES, + LOAD_VISIBLE_RELATION_WEIGHTS, + RANK_CHUNKS, + LOAD_CHUNKS + } + + public enum QueryOperationOutcome { + SUCCEEDED, + FAILED + } + + public record QueryOperationMeasurement( + QueryOperation operation, + QueryOperationOutcome outcome, + Duration duration, + int inputCount, + int outputCount) { + + public QueryOperationMeasurement { + Objects.requireNonNull(operation, "operation"); + Objects.requireNonNull(outcome, "outcome"); + Objects.requireNonNull(duration, "duration"); + if (duration.isNegative()) { + throw new IllegalArgumentException("duration must not be negative"); + } + if (inputCount < 0 || outputCount < 0) { + throw new IllegalArgumentException("counts must be non-negative"); + } + } + } + + @FunctionalInterface + public interface QueryOperationObserver { + QueryOperationObserver NO_OP = measurement -> { }; + + void observe(QueryOperationMeasurement measurement); + } } diff --git a/components/graph-rag-testkit/src/test/java/com/orgmemory/graphrag/testkit/LightRagQueryRuntimeConformanceTests.java b/components/graph-rag-testkit/src/test/java/com/orgmemory/graphrag/testkit/LightRagQueryRuntimeConformanceTests.java index a5eb7ceb..5b546ce8 100644 --- a/components/graph-rag-testkit/src/test/java/com/orgmemory/graphrag/testkit/LightRagQueryRuntimeConformanceTests.java +++ b/components/graph-rag-testkit/src/test/java/com/orgmemory/graphrag/testkit/LightRagQueryRuntimeConformanceTests.java @@ -156,6 +156,38 @@ void mixBatchesRawQueryLowKeywordsAndHighKeywordsExactlyOnce() { .anyMatch(signal -> signal.origin() == LightRagQueryResult.Origin.RELATION)); } + @Test + void preparedExecutionReportsBoundedProjectionOperationsWithoutPayload() { + LightRagQueryRequest request = request( + LightRagQueryMode.MIX, + QueryOutputMode.CONTEXT, + false, + true, + false, + trustedKeywords()); + var prepared = engine.prepare(request); + List measurements = + new ArrayList<>(); + + LightRagQueryResult result = + engine.executePrepared(request, prepared, measurements::add); + + assertEquals(LightRagQueryResult.Status.SUCCESS, result.status()); + assertTrue(measurements.stream().anyMatch(measurement -> + measurement.operation() + == LightRagQueryEngine.QueryOperation.SEARCH_ENTITIES)); + assertTrue(measurements.stream().anyMatch(measurement -> + measurement.operation() + == LightRagQueryEngine.QueryOperation.EXPAND_ENTITY_IDS)); + assertTrue(measurements.stream().anyMatch(measurement -> + measurement.operation() + == LightRagQueryEngine.QueryOperation.LOAD_INCIDENT_RELATIONS)); + assertTrue(measurements.stream().allMatch(measurement -> + !measurement.duration().isNegative() + && measurement.inputCount() >= 0 + && measurement.outputCount() >= 0)); + } + @Test void onePreparedQueryReusesKeywordAndEmbeddingEffectsAcrossSnapshotExecutions() { LightRagQueryRequest request = request( diff --git a/core/src/main/java/com/orgmemory/core/knowledge/retrieval/DefaultGraphRagKnowledgeRetrievalService.java b/core/src/main/java/com/orgmemory/core/knowledge/retrieval/DefaultGraphRagKnowledgeRetrievalService.java index 4d1ca4f8..7b67e93e 100644 --- a/core/src/main/java/com/orgmemory/core/knowledge/retrieval/DefaultGraphRagKnowledgeRetrievalService.java +++ b/core/src/main/java/com/orgmemory/core/knowledge/retrieval/DefaultGraphRagKnowledgeRetrievalService.java @@ -636,7 +636,10 @@ private PublishedSpaceQuery queryPublishedSpaces( futures.add(completed.submit(tasks.decorate(() -> new IndexedSnapshotQueryResult( resultIndex, - queryPublishedSpace(request, prepared))))); + queryPublishedSpace( + operationId, + request, + prepared))))); } SnapshotQueryResult[] ordered = new SnapshotQueryResult[requests.size()]; @@ -679,12 +682,18 @@ private PublishedSpaceQuery queryPublishedSpaces( } private SnapshotQueryResult queryPublishedSpace( + UUID operationId, LightRagQueryRequest request, LightRagPreparedQuery prepared) throws Exception { long startedAt = System.nanoTime(); - LightRagQueryResult result = - admission.execute(() -> - engine.executePrepared(request, prepared)); + LightRagQueryResult result = admission.execute(() -> engine.executePrepared( + request, + prepared, + measurement -> emitQueryOperation( + operationId, + request.scope().organizationId(), + request.scope().authorizationFingerprint(), + measurement))); return new SnapshotQueryResult( result, Duration.ofNanos(Math.max( @@ -694,6 +703,54 @@ private SnapshotQueryResult queryPublishedSpace( request.snapshot().namespace()); } + private void emitQueryOperation( + UUID operationId, + UUID organizationId, + String scopeFingerprint, + LightRagQueryEngine.QueryOperationMeasurement measurement) { + GraphRagEventSink.Outcome outcome = switch (measurement.outcome()) { + case SUCCEEDED -> GraphRagEventSink.Outcome.SUCCEEDED; + case FAILED -> GraphRagEventSink.Outcome.FAILED; + }; + safeEmit(new GraphRagEventSink.GraphRagEvent( + operationId, + organizationId, + queryOperationStage(measurement.operation()), + outcome, + measurement.duration(), + measurement.inputCount(), + measurement.outputCount(), + null, + scopeFingerprint, + null, + outcome == GraphRagEventSink.Outcome.FAILED + ? "query_operation_failed" + : null, + Instant.now())); + } + + private static GraphRagEventSink.Stage queryOperationStage( + LightRagQueryEngine.QueryOperation operation) { + return switch (operation) { + case SEARCH_ENTITIES -> GraphRagEventSink.Stage.QUERY_SEARCH_ENTITIES; + case SEARCH_RELATIONS -> GraphRagEventSink.Stage.QUERY_SEARCH_RELATIONS; + case SEARCH_CHUNKS -> GraphRagEventSink.Stage.QUERY_SEARCH_CHUNKS; + case EXPAND_ENTITY_IDS -> GraphRagEventSink.Stage.QUERY_EXPAND_ENTITY_IDS; + case LOAD_INCIDENT_RELATIONS -> + GraphRagEventSink.Stage.QUERY_LOAD_INCIDENT_RELATIONS; + case LOAD_ENTITY_CONTRIBUTIONS -> + GraphRagEventSink.Stage.QUERY_LOAD_ENTITY_CONTRIBUTIONS; + case LOAD_RELATION_CONTRIBUTIONS -> + GraphRagEventSink.Stage.QUERY_LOAD_RELATION_CONTRIBUTIONS; + case LOAD_VISIBLE_ENTITY_DEGREES -> + GraphRagEventSink.Stage.QUERY_LOAD_VISIBLE_ENTITY_DEGREES; + case LOAD_VISIBLE_RELATION_WEIGHTS -> + GraphRagEventSink.Stage.QUERY_LOAD_VISIBLE_RELATION_WEIGHTS; + case RANK_CHUNKS -> GraphRagEventSink.Stage.QUERY_RANK_CHUNKS; + case LOAD_CHUNKS -> GraphRagEventSink.Stage.QUERY_LOAD_CHUNKS; + }; + } + private record IndexedSnapshotQueryResult( int index, SnapshotQueryResult result) { diff --git a/core/src/test/java/com/orgmemory/core/knowledge/retrieval/GraphRagKnowledgeRetrievalServiceTests.java b/core/src/test/java/com/orgmemory/core/knowledge/retrieval/GraphRagKnowledgeRetrievalServiceTests.java index 59c4fddd..082183ff 100644 --- a/core/src/test/java/com/orgmemory/core/knowledge/retrieval/GraphRagKnowledgeRetrievalServiceTests.java +++ b/core/src/test/java/com/orgmemory/core/knowledge/retrieval/GraphRagKnowledgeRetrievalServiceTests.java @@ -125,7 +125,7 @@ void multipleSpacesPrepareOneLogicalQueryBeforeSnapshotRetrieval() { preparedQueryPlan(); when(engine.prepare(any())).thenReturn(prepared); CountDownLatch bothSpacesStarted = new CountDownLatch(2); - when(engine.executePrepared(any(), any())).thenAnswer(invocation -> { + when(engine.executePrepared(any(), any(), any())).thenAnswer(invocation -> { bothSpacesStarted.countDown(); assertTrue( bothSpacesStarted.await(2, TimeUnit.SECONDS), @@ -152,7 +152,7 @@ void multipleSpacesPrepareOneLogicalQueryBeforeSnapshotRetrieval() { assertEquals(List.of(), result.evidence()); verify(engine).prepare(any()); verify(engine, times(2)) - .executePrepared(any(), any()); + .executePrepared(any(), any(), any()); ArgumentCaptor captured = ArgumentCaptor.forClass(GraphRagEventSink.GraphRagEvent.class); verify(events, atLeastOnce()).emit(captured.capture()); @@ -190,7 +190,7 @@ void acquiresAdmissionPermitBeforeExecutingTheSnapshotStoreQuery() { LightRagQueryEngine engine = mock(LightRagQueryEngine.class); LightRagPreparedQuery prepared = preparedQueryPlan(); when(engine.prepare(any())).thenReturn(prepared); - when(engine.executePrepared(any(), any())).thenAnswer(invocation -> { + when(engine.executePrepared(any(), any(), any())).thenAnswer(invocation -> { assertEquals( 0, admission.availablePermits(), @@ -213,7 +213,7 @@ void acquiresAdmissionPermitBeforeExecutingTheSnapshotStoreQuery() { 10, "request-admission-before-checkout"); - verify(engine).executePrepared(any(), any()); + verify(engine).executePrepared(any(), any(), any()); assertEquals(1, admission.availablePermits()); } @@ -232,7 +232,7 @@ void continuousAdmissionDoesNotWaitForAnEarlierSnapshotBatch() CountDownLatch firstStarted = new CountDownLatch(1); CountDownLatch releaseFirst = new CountDownLatch(1); CountDownLatch thirdStarted = new CountDownLatch(1); - when(engine.executePrepared(any(), any())).thenAnswer(invocation -> { + when(engine.executePrepared(any(), any(), any())).thenAnswer(invocation -> { LightRagQueryRequest request = invocation.getArgument(0); Set assets = request.scope().authorizedAssetIds(); if (assets.contains(ASSET_ID)) { @@ -285,7 +285,7 @@ void snapshotFailureCancelsOutstandingContinuouslyAdmittedTasks() CountDownLatch blockersStarted = new CountDownLatch(2); CountDownLatch neverReleased = new CountDownLatch(1); CountDownLatch interrupted = new CountDownLatch(2); - when(engine.executePrepared(any(), any())).thenAnswer(invocation -> { + when(engine.executePrepared(any(), any(), any())).thenAnswer(invocation -> { LightRagQueryRequest request = invocation.getArgument(0); if (request.scope().authorizedAssetIds().contains(SECOND_ASSET_ID)) { assertTrue(blockersStarted.await(2, TimeUnit.SECONDS)); @@ -351,11 +351,20 @@ void retrievalObservationRunsKeywordAndRawQueryPathsWithoutAnswerGeneration() { ? keywordPrepared : bypassPrepared; }); - when(engine.executePrepared(any(), any())).thenReturn(queryResult( - allowed.forKnowledgeSpace(SPACE_ID).authorizationFingerprint(), - grounding, - false, - false)); + when(engine.executePrepared(any(), any(), any())).thenAnswer(invocation -> { + LightRagQueryEngine.QueryOperationObserver observer = invocation.getArgument(2); + observer.observe(new LightRagQueryEngine.QueryOperationMeasurement( + LightRagQueryEngine.QueryOperation.LOAD_INCIDENT_RELATIONS, + LightRagQueryEngine.QueryOperationOutcome.SUCCEEDED, + Duration.ofMillis(17), + 12, + 9)); + return queryResult( + allowed.forKnowledgeSpace(SPACE_ID).authorizationFingerprint(), + grounding, + false, + false); + }); when(engine.consolidateGrounding(any(), any(), any())).thenReturn(preparedGrounding); when(engine.renderGrounding(any(), any(), any())).thenReturn(preparedGrounding); RelationshipAuthorizationSetPort finalAuthorization = @@ -369,6 +378,7 @@ void retrievalObservationRunsKeywordAndRawQueryPathsWithoutAnswerGeneration() { candidate(ENTITY_CHUNK_ID), candidate(RELATION_CHUNK_ID), candidate(CHUNK_ID))); + GraphRagEventSink events = mock(GraphRagEventSink.class); GraphRagKnowledgeRetrievalService service = service( scopes, finalAuthorization, @@ -376,7 +386,7 @@ void retrievalObservationRunsKeywordAndRawQueryPathsWithoutAnswerGeneration() { engine, GraphRagRetrievalPolicy.defaults(), audit, - mock(GraphRagEventSink.class)); + events); GraphRagKnowledgeRetrievalService.RetrievalObservation observation = service.observe( actor, @@ -404,6 +414,25 @@ void retrievalObservationRunsKeywordAndRawQueryPathsWithoutAnswerGeneration() { assertEquals("model", observation.keywordPlan().source()); assertEquals(3, observation.keywordSeededDocuments().size()); assertEquals(3, observation.bypassDocuments().size()); + ArgumentCaptor emitted = + ArgumentCaptor.forClass(GraphRagEventSink.GraphRagEvent.class); + verify(events, atLeastOnce()).emit(emitted.capture()); + GraphRagEventSink.GraphRagEvent operation = emitted.getAllValues().stream() + .filter(event -> event.stage() + == GraphRagEventSink.Stage.QUERY_LOAD_INCIDENT_RELATIONS) + .findFirst() + .orElseThrow(); + GraphRagEventSink.GraphRagEvent retrieval = emitted.getAllValues().stream() + .filter(event -> event.stage() == GraphRagEventSink.Stage.RETRIEVE) + .findFirst() + .orElseThrow(); + assertEquals(Duration.ofMillis(17), operation.duration()); + assertEquals(12, operation.inputCount()); + assertEquals(9, operation.outputCount()); + assertEquals( + allowed.forKnowledgeSpace(SPACE_ID).authorizationFingerprint(), + operation.scopeFingerprint()); + assertEquals(retrieval.operationId(), operation.operationId()); } @Test @@ -483,7 +512,7 @@ void multipleSpacesDoNotInvokePerSpaceReranking() { "request-multi-space-rerank")); verify(engine, never()).prepare(any()); - verify(engine, never()).executePrepared(any(), any()); + verify(engine, never()).executePrepared(any(), any(), any()); } @Test @@ -531,7 +560,7 @@ void revocationBetweenRetrievalAndCitationCausesAFullRetryWithoutEgress() { LightRagGrounding grounding = grounding(); LightRagGroundingAssembler.PreparedGrounding prepared = prepared(grounding); - when(engine.executePrepared(any(), any())).thenReturn(queryResult( + when(engine.executePrepared(any(), any(), any())).thenReturn(queryResult( allowed.forKnowledgeSpace(SPACE_ID) .authorizationFingerprint(), grounding, @@ -585,7 +614,7 @@ void revocationBetweenRetrievalAndCitationCausesAFullRetryWithoutEgress() { assertEquals(List.of(), result.evidence()); verify(engine).prepare(any()); - verify(engine).executePrepared(any(), any()); + verify(engine).executePrepared(any(), any(), any()); verify(finalAuthorization, never()).batchCheck(any()); assertEquals(0, canonical.recheckCount); ArgumentCaptor captured = @@ -624,7 +653,7 @@ void verifiesTheCompleteGraphGroundingBeforeCreatingTheModelInput() { LightRagPreparedQuery queryPlan = preparedQueryPlan(); when(engine.prepare(any())) .thenReturn(queryPlan); - when(engine.executePrepared(any(), any())).thenReturn(queryResult( + when(engine.executePrepared(any(), any(), any())).thenReturn(queryResult( allowed.forKnowledgeSpace(SPACE_ID) .authorizationFingerprint(), grounding, @@ -730,7 +759,7 @@ void contextAssemblyReportsWhatTheAnswerCostAndWhatTheBudgetRefused() { LightRagQueryEngine engine = mock(LightRagQueryEngine.class); LightRagPreparedQuery queryPlan = preparedQueryPlan(); when(engine.prepare(any())).thenReturn(queryPlan); - when(engine.executePrepared(any(), any())).thenReturn(queryResult( + when(engine.executePrepared(any(), any(), any())).thenReturn(queryResult( allowed.forKnowledgeSpace(SPACE_ID).authorizationFingerprint(), grounding, true, @@ -809,7 +838,7 @@ void authorizationModelMismatchCannotReachTheVerifiedRenderer() { LightRagPreparedQuery queryPlan = preparedQueryPlan(); when(engine.prepare(any())) .thenReturn(queryPlan); - when(engine.executePrepared(any(), any())).thenReturn(queryResult( + when(engine.executePrepared(any(), any(), any())).thenReturn(queryResult( allowed.forKnowledgeSpace(SPACE_ID) .authorizationFingerprint(), grounding, @@ -1002,7 +1031,7 @@ private static AuthorizationFixture authorizationFixture( LightRagQueryEngine engine = mock(LightRagQueryEngine.class); LightRagPreparedQuery queryPlan = preparedQueryPlan(); when(engine.prepare(any())).thenReturn(queryPlan); - when(engine.executePrepared(any(), any())).thenReturn(queryResult( + when(engine.executePrepared(any(), any(), any())).thenReturn(queryResult( allowed.forKnowledgeSpace(SPACE_ID) .authorizationFingerprint(), grounding, diff --git a/integrations/graph-rag-postgres/src/main/java/com/orgmemory/graphrag/postgres/PostgresGraphStore.java b/integrations/graph-rag-postgres/src/main/java/com/orgmemory/graphrag/postgres/PostgresGraphStore.java index 2ea39997..683bf269 100644 --- a/integrations/graph-rag-postgres/src/main/java/com/orgmemory/graphrag/postgres/PostgresGraphStore.java +++ b/integrations/graph-rag-postgres/src/main/java/com/orgmemory/graphrag/postgres/PostgresGraphStore.java @@ -445,23 +445,49 @@ public List loadRelationContributions( if (!readable(scope, snapshot, ids)) { return List.of(); } - return jdbc.query( + return support.read(GRAPH_QUERY_STATEMENT_TIMEOUT, () -> jdbc.query( """ + WITH candidate_relations AS MATERIALIZED ( + SELECT relation.* + FROM projection_graph_relations relation + WHERE relation.batch_id = :batchId + AND relation.relation_id IN (:ids) + ), + visible_entities AS MATERIALIZED ( + SELECT DISTINCT contribution.entity_id + FROM projection_graph_entity_contributions contribution + JOIN candidate_relations relation + ON contribution.entity_id IN ( + relation.source_entity_id, + relation.target_entity_id) + WHERE contribution.batch_id = :batchId + AND contribution.organization_id = :organizationId + AND contribution.knowledge_asset_id + IN (:authorizedAssetIds) + ), + visible_relation_contributions AS MATERIALIZED ( + SELECT contribution.* + FROM projection_graph_relation_contributions contribution + JOIN candidate_relations relation + ON relation.relation_id = contribution.relation_id + JOIN visible_entities source_evidence + ON source_evidence.entity_id = relation.source_entity_id + JOIN visible_entities target_evidence + ON target_evidence.entity_id = relation.target_entity_id + WHERE contribution.batch_id = :batchId + AND contribution.organization_id = :organizationId + AND contribution.knowledge_asset_id + IN (:authorizedAssetIds) + ) SELECT contribution.*, relation.source_entity_id, relation.target_entity_id, relation.orientation - FROM ( - """ - + VISIBLE_RELATION_CONTRIBUTIONS - + """ - ) contribution - JOIN projection_graph_relations relation - ON relation.batch_id = contribution.batch_id - AND relation.relation_id = contribution.relation_id - WHERE contribution.relation_id IN (:ids) + FROM visible_relation_contributions contribution + JOIN candidate_relations relation + ON relation.relation_id = contribution.relation_id ORDER BY contribution.contribution_id """, visibility(scope, snapshot).addValue("ids", ids), - (resultSet, rowNumber) -> relationContribution(resultSet)); + (resultSet, rowNumber) -> relationContribution(resultSet))); } @Override @@ -636,24 +662,53 @@ public Map loadVisibleRelationWeights( if (!readable(scope, snapshot, ids)) { return Map.of(); } - Map weights = new LinkedHashMap<>(); - jdbc.query( - """ - SELECT relation_id, sum(weight) AS weight - FROM ( - """ - + VISIBLE_RELATION_CONTRIBUTIONS - + """ - ) visible - WHERE relation_id IN (:ids) - GROUP BY relation_id - ORDER BY relation_id - """, - visibility(scope, snapshot).addValue("ids", ids), - (RowCallbackHandler) resultSet -> weights.put( - resultSet.getObject("relation_id", UUID.class), - resultSet.getDouble("weight"))); - return Map.copyOf(weights); + return support.read(GRAPH_QUERY_STATEMENT_TIMEOUT, () -> { + Map weights = new LinkedHashMap<>(); + jdbc.query( + """ + WITH candidate_relations AS MATERIALIZED ( + SELECT relation.* + FROM projection_graph_relations relation + WHERE relation.batch_id = :batchId + AND relation.relation_id IN (:ids) + ), + visible_entities AS MATERIALIZED ( + SELECT DISTINCT contribution.entity_id + FROM projection_graph_entity_contributions contribution + JOIN candidate_relations relation + ON contribution.entity_id IN ( + relation.source_entity_id, + relation.target_entity_id) + WHERE contribution.batch_id = :batchId + AND contribution.organization_id = :organizationId + AND contribution.knowledge_asset_id + IN (:authorizedAssetIds) + ), + visible_relation_contributions AS MATERIALIZED ( + SELECT contribution.* + FROM projection_graph_relation_contributions contribution + JOIN candidate_relations relation + ON relation.relation_id = contribution.relation_id + JOIN visible_entities source_evidence + ON source_evidence.entity_id = relation.source_entity_id + JOIN visible_entities target_evidence + ON target_evidence.entity_id = relation.target_entity_id + WHERE contribution.batch_id = :batchId + AND contribution.organization_id = :organizationId + AND contribution.knowledge_asset_id + IN (:authorizedAssetIds) + ) + SELECT relation_id, sum(weight) AS weight + FROM visible_relation_contributions + GROUP BY relation_id + ORDER BY relation_id + """, + visibility(scope, snapshot).addValue("ids", ids), + (RowCallbackHandler) resultSet -> weights.put( + resultSet.getObject("relation_id", UUID.class), + resultSet.getDouble("weight"))); + return Map.copyOf(weights); + }); } @Override diff --git a/integrations/graph-rag-postgres/src/test/java/com/orgmemory/graphrag/postgres/PostgresGraphStoreOptionsTests.java b/integrations/graph-rag-postgres/src/test/java/com/orgmemory/graphrag/postgres/PostgresGraphStoreOptionsTests.java index 5d6c7473..44b2948a 100644 --- a/integrations/graph-rag-postgres/src/test/java/com/orgmemory/graphrag/postgres/PostgresGraphStoreOptionsTests.java +++ b/integrations/graph-rag-postgres/src/test/java/com/orgmemory/graphrag/postgres/PostgresGraphStoreOptionsTests.java @@ -2,7 +2,11 @@ import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; +import java.io.IOException; +import java.nio.file.Files; +import java.nio.file.Path; import java.util.Set; import org.junit.jupiter.api.Test; @@ -28,4 +32,27 @@ void approximateIndexRequiresAtLeastOneDimension() { 100, "")); } + + @Test + void hotRelationReadsAreCandidateFirstAndPostgresBounded() throws IOException { + String source = Files.readString(Path.of( + "src/main/java/com/orgmemory/graphrag/postgres/PostgresGraphStore.java")); + + assertHotRead(source, "loadRelationContributions", "loadIncidentRelations"); + assertHotRead(source, "loadVisibleRelationWeights", "discard"); + } + + private static void assertHotRead(String source, String method, String nextMethod) { + int start = source.lastIndexOf("public ", source.indexOf(" " + method + "(")); + int nextName = source.indexOf(" " + nextMethod + "(", start); + int end = source.lastIndexOf("public ", nextName); + String body = source.substring(start, end); + + assertTrue( + body.contains("support.read(GRAPH_QUERY_STATEMENT_TIMEOUT"), + () -> method + " must use the transaction-local PostgreSQL budget"); + assertTrue( + body.contains("WITH candidate_relations AS MATERIALIZED"), + () -> method + " must constrain authorization work to requested relations first"); + } }