Skip to content

Commit 2fb9b4f

Browse files
committed
Share search property filter dispatch
Signed-off-by: Arnab Nandy <arnab_nandy7@yahoo.com>
1 parent d40d889 commit 2fb9b4f

2 files changed

Lines changed: 135 additions & 53 deletions

File tree

embabel-agent-rag/embabel-agent-rag-core/src/main/kotlin/com/embabel/agent/rag/tools/CoreSearchTools.kt

Lines changed: 40 additions & 53 deletions
Original file line numberDiff line numberDiff line change
@@ -78,6 +78,25 @@ internal fun List<SimilarityResult<out Retrievable>>.withNeighbours(
7878
return this + extra
7979
}
8080

81+
@Suppress("UNCHECKED_CAST")
82+
internal fun <T : Retrievable> VectorSearch.vectorSearchWithFilter(
83+
request: TextSimilaritySearchRequest,
84+
clazz: Class<T>,
85+
metadataFilter: PropertyFilter?,
86+
entityFilter: EntityFilter?,
87+
): List<SimilarityResult<T>> = when {
88+
metadataFilter == null && entityFilter == null -> vectorSearch(request, clazz)
89+
this is FilteringVectorSearch -> vectorSearchWithFilter(request, clazz, metadataFilter, entityFilter)
90+
else -> PostFilteringSearch.search(
91+
request,
92+
metadataFilter,
93+
entityFilter,
94+
TopKInflationStrategy.DEFAULT,
95+
) { inflatedRequest ->
96+
vectorSearch(inflatedRequest, clazz)
97+
} as List<SimilarityResult<T>>
98+
}
99+
81100
internal class VectorSearchTools @JvmOverloads constructor(
82101
private val vectorSearch: VectorSearch,
83102
private val searchFor: List<Class<out Retrievable>> = listOf(Chunk::class.java),
@@ -119,36 +138,10 @@ internal class VectorSearchTools @JvmOverloads constructor(
119138

120139
private fun searchForAllTypes(request: TextSimilaritySearchRequest): List<SimilarityResult<out Retrievable>> {
121140
val allResults = searchFor.flatMap { clazz ->
122-
searchWithFilter(request, clazz)
141+
vectorSearch.vectorSearchWithFilter(request, clazz, metadataFilter, entityFilter)
123142
}
124143
return deduplicateByIdKeepingHighestScore(allResults)
125144
}
126-
127-
@Suppress("UNCHECKED_CAST")
128-
private fun <T : Retrievable> searchWithFilter(
129-
request: TextSimilaritySearchRequest,
130-
clazz: Class<T>,
131-
): List<SimilarityResult<T>> {
132-
if (metadataFilter == null && entityFilter == null) {
133-
return vectorSearch.vectorSearch(request, clazz)
134-
}
135-
136-
// If backend supports native filtering, use it
137-
if (vectorSearch is FilteringVectorSearch) {
138-
return vectorSearch.vectorSearchWithFilter(request, clazz, metadataFilter, entityFilter)
139-
}
140-
141-
// Fallback: inflate topK, search, post-filter, take topK
142-
// Note: PostFilteringSearch requires Datum constraint, so we cast
143-
return PostFilteringSearch.search(
144-
request,
145-
metadataFilter,
146-
entityFilter,
147-
TopKInflationStrategy.DEFAULT
148-
) { inflatedRequest ->
149-
vectorSearch.vectorSearch(inflatedRequest, clazz)
150-
} as List<SimilarityResult<T>>
151-
}
152145
}
153146

154147
/**
@@ -378,6 +371,25 @@ internal class SectionReadingTools @JvmOverloads constructor(
378371
}
379372
}
380373

374+
@Suppress("UNCHECKED_CAST")
375+
internal fun <T : Retrievable> TextSearch.textSearchWithFilter(
376+
request: TextSimilaritySearchRequest,
377+
clazz: Class<T>,
378+
metadataFilter: PropertyFilter?,
379+
entityFilter: EntityFilter?,
380+
): List<SimilarityResult<T>> = when {
381+
metadataFilter == null && entityFilter == null -> textSearch(request, clazz)
382+
this is FilteringTextSearch -> textSearchWithFilter(request, clazz, metadataFilter, entityFilter)
383+
else -> PostFilteringSearch.search(
384+
request,
385+
metadataFilter,
386+
entityFilter,
387+
TopKInflationStrategy.DEFAULT,
388+
) { inflatedRequest ->
389+
textSearch(inflatedRequest, clazz)
390+
} as List<SimilarityResult<T>>
391+
}
392+
381393
/**
382394
* Tools to perform text search operations with the syntax supported by
383395
* the backing [TextSearch] store.
@@ -478,36 +490,11 @@ internal class TextSearchTools @JvmOverloads constructor(
478490

479491
private fun searchForAllTypes(request: TextSimilaritySearchRequest): List<SimilarityResult<out Retrievable>> {
480492
val allResults = searchFor.flatMap { clazz ->
481-
searchWithFilter(request, clazz)
493+
textSearch.textSearchWithFilter(request, clazz, metadataFilter, entityFilter)
482494
}
483495
return deduplicateByIdKeepingHighestScore(allResults)
484496
}
485497

486-
@Suppress("UNCHECKED_CAST")
487-
private fun <T : Retrievable> searchWithFilter(
488-
request: TextSimilaritySearchRequest,
489-
clazz: Class<T>,
490-
): List<SimilarityResult<T>> {
491-
if (metadataFilter == null && entityFilter == null) {
492-
return textSearch.textSearch(request, clazz)
493-
}
494-
495-
// If backend supports native filtering, use it
496-
if (textSearch is FilteringTextSearch) {
497-
return textSearch.textSearchWithFilter(request, clazz, metadataFilter, entityFilter)
498-
}
499-
500-
// Fallback: inflate topK, search, post-filter, take topK
501-
return PostFilteringSearch.search(
502-
request,
503-
metadataFilter,
504-
entityFilter,
505-
TopKInflationStrategy.DEFAULT
506-
) { inflatedRequest ->
507-
textSearch.textSearch(inflatedRequest, clazz)
508-
} as List<SimilarityResult<T>>
509-
}
510-
511498
companion object {
512499
/**
513500
* Compose the tool's top-level description from the store's

embabel-agent-rag/embabel-agent-rag-core/src/test/kotlin/com/embabel/agent/rag/tools/ToolishRagTest.kt

Lines changed: 95 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616
package com.embabel.agent.rag.tools
1717

1818
import com.embabel.agent.api.tool.Tool
19+
import com.embabel.agent.filter.PropertyFilter
1920
import com.embabel.agent.rag.model.Chunk
2021
import com.embabel.agent.rag.model.ContentElement
2122
import com.embabel.agent.rag.model.NamedEntityData.Companion.ENTITY_LABEL
@@ -462,6 +463,28 @@ class ToolishRagTest {
462463
}
463464
}
464465

466+
@Nested
467+
inner class SearchWithFilterExtensionsTests {
468+
469+
@Test
470+
fun `filter helpers should be callable on CoreSearchOperations`() {
471+
val searchOperations = mockk<CoreSearchOperations>()
472+
val request = TextSimilaritySearchRequest("test query", 0.5, 5)
473+
every {
474+
searchOperations.vectorSearch(request, Chunk::class.java)
475+
} returns emptyList()
476+
every {
477+
searchOperations.textSearch(request, Chunk::class.java)
478+
} returns emptyList()
479+
480+
val vectorResults = searchOperations.vectorSearchWithFilter(request, Chunk::class.java, null, null)
481+
val textResults = searchOperations.textSearchWithFilter(request, Chunk::class.java, null, null)
482+
483+
assertTrue(vectorResults.isEmpty())
484+
assertTrue(textResults.isEmpty())
485+
}
486+
}
487+
465488
@Nested
466489
inner class TextSearchToolsTests {
467490

@@ -506,6 +529,78 @@ class ToolishRagTest {
506529
assertEquals("0 results:", result)
507530
}
508531

532+
@Test
533+
fun `textSearch should use native filtering when supported`() {
534+
val textSearch = mockk<FilteringTextSearch>()
535+
val metadataFilter = PropertyFilter.eq("ownerId", "alice")
536+
val chunk = createChunk("chunk1", "Alice's content")
537+
every {
538+
textSearch.textSearchWithFilter(
539+
any<TextSimilaritySearchRequest>(),
540+
Chunk::class.java,
541+
metadataFilter,
542+
null,
543+
)
544+
} returns listOf(SimpleSimilaritySearchResult(match = chunk, score = 0.9))
545+
val tools = TextSearchTools(textSearch, metadataFilter = metadataFilter)
546+
547+
val result = tools.textSearch("test query", 5, 0.5)
548+
549+
verify(exactly = 1) {
550+
textSearch.textSearchWithFilter(
551+
match<TextSimilaritySearchRequest> { it.topK == 5 },
552+
Chunk::class.java,
553+
metadataFilter,
554+
null,
555+
)
556+
}
557+
assertTrue(result.contains("Alice's content"))
558+
}
559+
560+
@Test
561+
fun `textSearch should inflate topK and post-filter when native filtering is unavailable`() {
562+
val textSearch = mockk<TextSearch>()
563+
val metadataFilter = PropertyFilter.eq("ownerId", "alice")
564+
val included = Chunk(
565+
id = "included",
566+
text = "Alice's content",
567+
parentId = "parent",
568+
metadata = mapOf("ownerId" to "alice"),
569+
)
570+
val excluded = Chunk(
571+
id = "excluded",
572+
text = "Bob's content",
573+
parentId = "parent",
574+
metadata = mapOf("ownerId" to "bob"),
575+
)
576+
val truncated = Chunk(
577+
id = "truncated",
578+
text = "Alice's lower-ranked content",
579+
parentId = "parent",
580+
metadata = mapOf("ownerId" to "alice"),
581+
)
582+
every {
583+
textSearch.textSearch(any<TextSimilaritySearchRequest>(), Chunk::class.java)
584+
} returns listOf(
585+
SimpleSimilaritySearchResult(match = excluded, score = 0.9),
586+
SimpleSimilaritySearchResult(match = included, score = 0.8),
587+
SimpleSimilaritySearchResult(match = truncated, score = 0.7),
588+
)
589+
val tools = TextSearchTools(textSearch, metadataFilter = metadataFilter)
590+
591+
val result = tools.textSearch("test query", 1, 0.5)
592+
593+
verify(exactly = 1) {
594+
textSearch.textSearch(
595+
match<TextSimilaritySearchRequest> { it.topK == 3 },
596+
Chunk::class.java,
597+
)
598+
}
599+
assertTrue(result.contains("Alice's content"))
600+
assertFalse(result.contains("Bob's content"))
601+
assertFalse(result.contains("Alice's lower-ranked content"))
602+
}
603+
509604
}
510605

511606
@Nested

0 commit comments

Comments
 (0)