core: гибридный поиск (BM25 + вектор + RRF)

This commit is contained in:
2026-10-02 01:25:23 +03:00
parent c33f94fa13
commit a2f09fffb6
5 changed files with 527 additions and 0 deletions
@@ -0,0 +1,157 @@
package memo.core
import java.io.File
enum class SearchMode { HYBRID, LEX, VEC }
data class Hit(
val path: String,
val line: Int,
val heading: String,
val score: Double,
val text: String,
)
fun interface RefreshHook { fun refresh(root: File) }
class Searcher(
private val db: Db,
private val embedder: Embedder,
private val refresh: RefreshHook? = null,
) {
fun search(
root: File,
query: String,
k: Int = 8,
mode: SearchMode = SearchMode.HYBRID,
): List<Hit> {
if (refresh != null) {
refresh.refresh(root)
}
val lexRows: List<Pair<Long, Double>> =
if (mode == SearchMode.LEX || mode == SearchMode.HYBRID) lexSearch(query) else emptyList()
val vecRows: List<Pair<Long, Double>> =
if (mode == SearchMode.VEC || mode == SearchMode.HYBRID) vecSearch(query) else emptyList()
val scores = HashMap<Long, Double>()
for ((idx, row) in lexRows.withIndex()) {
val rank = idx + 1
scores.merge(row.first, 1.0 / (60.0 + rank), Double::plus)
}
for ((idx, row) in vecRows.withIndex()) {
val rank = idx + 1
scores.merge(row.first, 1.0 / (60.0 + rank), Double::plus)
}
if (scores.isEmpty()) return emptyList()
val topEntries = scores.entries.sortedByDescending { it.value }.take(k)
val orderedRowids = topEntries.map { it.key }
val scoreByRowid = orderedRowids.associateWith { scores[it]!! }
val placeholders = orderedRowids.joinToString(",") { "?" }
val select = db.conn.prepare(
"SELECT id, path, line, heading, text FROM chunks WHERE id IN ($placeholders)"
)
val rowData = HashMap<Long, RowData>()
try {
for ((idx, id) in orderedRowids.withIndex()) {
select.bindLong(idx + 1, id)
}
val rs = select.executeQuery()
try {
while (rs.next()) {
val id = rs.getLong(0)!!
val path = rs.getText(1) ?: ""
val line = rs.getInt(2)!!
val heading = rs.getText(3) ?: ""
val textRaw = rs.getText(4) ?: ""
val text = if (textRaw.length > 1200) textRaw.substring(0, 1200) else textRaw
rowData[id] = RowData(path, line, heading, text)
}
} finally {
rs.close()
}
} finally {
select.close()
}
return orderedRowids.mapNotNull { id ->
val rd = rowData[id] ?: return@mapNotNull null
Hit(
path = rd.path,
line = rd.line,
heading = rd.heading,
score = scoreByRowid[id] ?: 0.0,
text = rd.text,
)
}
}
private fun lexSearch(query: String): List<Pair<Long, Double>> {
val tokens = tokensOf(query)
if (tokens.isEmpty()) return emptyList()
val matchExpr = tokens.joinToString(" OR ") { "\"$it\"" }
val stmt = db.conn.prepare(
"SELECT rowid, bm25(chunks_fts) AS s FROM chunks_fts " +
"WHERE chunks_fts MATCH ? ORDER BY s LIMIT 32"
)
return try {
stmt.bindText(1, matchExpr)
val rs = stmt.executeQuery()
try {
val out = ArrayList<Pair<Long, Double>>()
while (rs.next()) {
out.add(rs.getLong(0)!! to rs.getDouble(1)!!)
}
out
} finally {
rs.close()
}
} catch (_: Throwable) {
emptyList()
} finally {
stmt.close()
}
}
private fun vecSearch(query: String): List<Pair<Long, Double>> {
val stmt = db.conn.prepare(
"SELECT rowid, distance FROM chunks_vec " +
"WHERE embedding MATCH ? ORDER BY distance LIMIT 32"
)
return try {
stmt.bindVector(1, embedder.embed(query))
val rs = stmt.executeQuery()
try {
val out = ArrayList<Pair<Long, Double>>()
while (rs.next()) {
out.add(rs.getLong(0)!! to rs.getDouble(1)!!)
}
out
} finally {
rs.close()
}
} finally {
stmt.close()
}
}
private data class RowData(val path: String, val line: Int, val heading: String, val text: String)
}
private fun tokensOf(query: String): List<String> {
val tokens = ArrayList<String>()
val sb = StringBuilder()
for (ch in query) {
if (Character.isLetterOrDigit(ch)) {
sb.append(ch)
} else {
if (sb.length >= 2) tokens.add(sb.toString())
sb.setLength(0)
}
}
if (sb.length >= 2) tokens.add(sb.toString())
return if (tokens.size <= 8) tokens else tokens.subList(0, 8)
}
@@ -0,0 +1,141 @@
package memo.core
import java.io.File
import java.nio.file.Files
import kotlin.test.Test
import kotlin.test.assertEquals
import kotlin.test.assertTrue
class SearcherTest {
private class Fixture : AutoCloseable {
val tempDir: File = Files.createTempDirectory("memo-searcher-").toFile()
val db = Db(File(tempDir, "index.db").absolutePath)
val embedder: Embedder
init {
db.init()
val modelDir = System.getenv("MEMO_MODEL_DIR") ?: "/root/WORK/memo/models/siglip2"
val modelPath = "$modelDir/text_model_int8.onnx"
val tokenizerPath = "$modelDir/tokenizer.model"
embedder = Embedder(modelPath, tokenizerPath)
}
fun addChunk(path: String, heading: String, line: Int, text: String) {
val conn = db.conn
val insChunk = conn.prepare(
"INSERT INTO chunks(path, heading, line, ord, text, hash) VALUES (?, ?, ?, ?, ?, ?)"
)
try {
insChunk.bindText(1, path)
insChunk.bindText(2, heading)
insChunk.bindLong(3, line.toLong())
insChunk.bindLong(4, 0L)
insChunk.bindText(5, text)
insChunk.bindText(6, "$path#$line")
insChunk.executeUpdate()
} finally {
insChunk.close()
}
val id = conn.lastInsertRowId
val insFts = conn.prepare(
"INSERT INTO chunks_fts(rowid, text, heading) VALUES (?, ?, ?)"
)
try {
insFts.bindLong(1, id)
insFts.bindText(2, text)
insFts.bindText(3, heading)
insFts.executeUpdate()
} finally {
insFts.close()
}
val insVec = conn.prepare(
"INSERT INTO chunks_vec(rowid, embedding) VALUES (?, ?)"
)
try {
insVec.bindLong(1, id)
insVec.bindVector(2, embedder.embed(text))
insVec.executeUpdate()
} finally {
insVec.close()
}
}
override fun close() {
embedder.close()
db.close()
tempDir.deleteRecursively()
}
}
@Test
fun lexModeFindsExactValue() {
Fixture().use { f ->
f.addChunk("/a.md", "A", 1, "Прокси корпоративных доменов на 76.132")
f.addChunk("/b.md", "B", 1, "Список контактов службы поддержки")
f.addChunk("/c.md", "C", 1, "Описание архитектуры сетевого шлюза")
val s = Searcher(f.db, f.embedder)
val hits = s.search(f.tempDir, "76.132", k = 8, mode = SearchMode.LEX)
assertTrue(hits.isNotEmpty(), "ожидались хиты, получено 0")
assertTrue(hits[0].text.contains("76.132"), "первый хит должен содержать 76.132")
}
}
@Test
fun vecModeFindsSemanticMatch() {
Fixture().use { f ->
f.addChunk("/server.md", "Server", 1, "Настройка сервера приложений и конфигурация Tomcat")
f.addChunk("/book.md", "Book", 1, "Аннотация книги по истории Древнего Рима")
f.addChunk("/ci.md", "CI", 1, "Пайплайн выпуска приложения: сборка, тесты, деплой в Kubernetes")
val s = Searcher(f.db, f.embedder)
val hits = s.search(f.tempDir, "настройку выпуска приложения", k = 1, mode = SearchMode.VEC)
assertTrue(hits.isNotEmpty(), "ожидались хиты, получено 0")
assertEquals("/ci.md", hits[0].path, "первый хит должен быть чанк про CI")
}
}
@Test
fun lexModeIgnoresVectorOnlyMatch() {
Fixture().use { f ->
f.addChunk("/p.md", "P", 1, "опрос")
f.addChunk("/s.md", "S", 1, "случай")
f.addChunk("/z.md", "Z", 1, "знание")
val s = Searcher(f.db, f.embedder)
val lex = s.search(f.tempDir, "автомобиль", k = 8, mode = SearchMode.LEX)
val vec = s.search(f.tempDir, "автомобиль", k = 8, mode = SearchMode.VEC)
assertTrue(lex.isEmpty(), "LEX должен быть пустым, получили ${lex.size} хитов")
assertTrue(vec.isNotEmpty(), "VEC должен быть непустым, получили 0")
}
}
@Test
fun hybridCombinesBoth() {
Fixture().use { f ->
f.addChunk("/exact.md", "Exact", 1, "Прокси корпоративных доменов на 76.132")
f.addChunk("/sem.md", "Sem", 1, "Пайплайн выпуска приложения: сборка, тесты, деплой в Kubernetes")
f.addChunk("/other.md", "Other", 1, "Заметки о книге по истории")
val s = Searcher(f.db, f.embedder)
val hits = s.search(f.tempDir, "выпуск приложения 76.132", k = 5, mode = SearchMode.HYBRID)
val paths = hits.map { it.path }.toSet()
assertTrue(paths.contains("/exact.md"), "чанк exact отсутствует в: $paths")
assertTrue(paths.contains("/sem.md"), "чанк sem отсутствует в: $paths")
}
}
@Test
fun resultsRespectKAndTextLength() {
Fixture().use { f ->
f.addChunk("/a.md", "A", 1, "Первый фрагмент про сервер")
f.addChunk("/b.md", "B", 1, "Второй фрагмент про книгу")
f.addChunk("/c.md", "C", 1, "Третий фрагмент про CI")
f.addChunk("/d.md", "D", 1, "Четвёртый фрагмент про настройку")
val s = Searcher(f.db, f.embedder)
val hits = s.search(f.tempDir, "фрагмент", k = 2, mode = SearchMode.HYBRID)
assertEquals(2, hits.size, "должно быть ровно 2 хита при k=2")
for (h in hits) {
assertTrue(h.text.isNotEmpty(), "text не должен быть пустым: $h")
assertTrue(h.text.length <= 1200, "длина text превышает 1200: ${h.text.length}")
}
}
}
}