| 1 | using System.Buffers.Binary; |
| 2 | using System.Globalization; |
| 3 | using System.Text; |
| 4 | using HybridCodebaseIndex.Core.Embeddings; |
| 5 | using Microsoft.Data.Sqlite; |
| 6 | |
| 7 | namespace HybridCodebaseIndex.Core; |
| 8 | |
| 9 | internal static partial class SqliteFtsIndex |
| 10 | { |
| 11 | internal static Task<(SearchResponse response, string? error)> SearchHybridAsync( |
| 12 | string workspaceRoot, |
| 13 | string dbPath, |
| 14 | string query, |
| 15 | int topN, |
| 16 | string? pathPrefix, |
| 17 | IReadOnlyList<string>? excludePathPrefixes, |
| 18 | IReadOnlyList<string>? extensions, |
| 19 | bool semantic, |
| 20 | double alpha, |
| 21 | double beta, |
| 22 | int vecTopK, |
| 23 | CancellationToken cancellationToken) |
| 24 | => Task.Run(() => SearchHybrid(workspaceRoot, dbPath, query, topN, pathPrefix, excludePathPrefixes, extensions, semantic, alpha, beta, vecTopK, cancellationToken), cancellationToken); |
| 25 | |
| 26 | private static (SearchResponse response, string? error) SearchHybrid( |
| 27 | string workspaceRoot, |
| 28 | string dbPath, |
| 29 | string userQuery, |
| 30 | int topN, |
| 31 | string? pathPrefix, |
| 32 | IReadOnlyList<string>? excludePathPrefixes, |
| 33 | IReadOnlyList<string>? extensions, |
| 34 | bool semantic, |
| 35 | double alpha, |
| 36 | double beta, |
| 37 | int vecTopK, |
| 38 | CancellationToken cancellationToken) |
| 39 | { |
| 40 | // FTS always available |
| 41 | var (ftsResp, ftsErr) = Search(workspaceRoot, dbPath, userQuery, topN, pathPrefix, excludePathPrefixes, extensions); |
| 42 | if (!semantic) |
| 43 | return (ftsResp, ftsErr); |
| 44 | if (!string.IsNullOrEmpty(ftsErr)) |
| 45 | return (ftsResp, ftsErr); |
| 46 | |
| 47 | if (!File.Exists(dbPath)) |
| 48 | return (ftsResp, ftsErr); |
| 49 | |
| 50 | workspaceRoot = Path.GetFullPath(workspaceRoot.TrimEnd(Path.DirectorySeparatorChar)); |
| 51 | using var conn = new SqliteConnection($"Data Source={dbPath};Mode=ReadOnly"); |
| 52 | conn.Open(); |
| 53 | |
| 54 | // If vec backend missing, just return FTS. |
| 55 | if (!TableExists(conn, "vectors") && !TableExists(conn, "vec_chunks")) |
| 56 | return (ftsResp, null); |
| 57 | |
| 58 | var settings = IndexSettings.TryLoadFromIndexDirectory(Path.GetDirectoryName(dbPath)!); |
| 59 | if (!settings.SemanticEnabled) |
| 60 | return (ftsResp, null); |
| 61 | |
| 62 | var provider = EmbeddingProviderFactory.Create(settings, Path.GetDirectoryName(dbPath)); |
| 63 | var qv = provider.EmbedAsync(userQuery, cancellationToken).GetAwaiter().GetResult(); |
| 64 | var qn = Norm(qv); |
| 65 | if (qn <= 1e-12) |
| 66 | return (ftsResp, null); |
| 67 | |
| 68 | vecTopK = Math.Clamp(vecTopK, 5, 200); |
| 69 | |
| 70 | var indexDir = Path.GetDirectoryName(dbPath)!; |
| 71 | var sqliteVecLoaded = SqliteVecInterop.TryEnableAndLoad(conn, settings, indexDir, out _); |
| 72 | List<(long chunkRowId, double sim)> vecHitsRaw = |
| 73 | sqliteVecLoaded && TableExists(conn, "vec_chunks") |
| 74 | ? SqliteVecInterop.QueryTopK(conn, qv, vecTopK) |
| 75 | : TableExists(conn, "vectors") |
| 76 | ? VecTopK(conn, qv, qn, vecTopK, settings) |
| 77 | : []; |
| 78 | var vecHits = FilterVecHitsByAllowedChunkExtensions(conn, settings, vecHitsRaw); |
| 79 | |
| 80 | // Merge by hitId (chunk rowid) |
| 81 | var merged = new Dictionary<long, IndexHit>(); |
| 82 | foreach (var h in ftsResp.Hits) |
| 83 | merged[h.HitId] = h; |
| 84 | |
| 85 | foreach (var (chunkId, sim) in vecHits) |
| 86 | { |
| 87 | cancellationToken.ThrowIfCancellationRequested(); |
| 88 | |
| 89 | if (merged.TryGetValue(chunkId, out var existing)) |
| 90 | { |
| 91 | var fused = new IndexHit( |
| 92 | existing.HitId, |
| 93 | existing.Path, |
| 94 | existing.Extension, |
| 95 | existing.HitKind, |
| 96 | RankScore: alpha * (existing.FtsScore ?? existing.RankScore) + beta * sim, |
| 97 | FtsScore: existing.FtsScore ?? existing.RankScore, |
| 98 | VecScore: sim, |
| 99 | existing.Snippet, |
| 100 | existing.LineStart, |
| 101 | existing.LineEnd, |
| 102 | existing.ChunkCharCount, |
| 103 | existing.LastWriteUtcIso); |
| 104 | merged[chunkId] = fused; |
| 105 | continue; |
| 106 | } |
| 107 | |
| 108 | // Fetch minimal chunk metadata for vec-only hits. |
| 109 | if (!TryGetChunk(conn, chunkId, out var hit)) |
| 110 | continue; |
| 111 | |
| 112 | merged[chunkId] = new IndexHit( |
| 113 | hit.HitId, |
| 114 | hit.Path, |
| 115 | hit.Extension, |
| 116 | HitKinds.TextVector, |
| 117 | RankScore: sim, |
| 118 | FtsScore: null, |
| 119 | VecScore: sim, |
| 120 | hit.Snippet, |
| 121 | hit.LineStart, |
| 122 | hit.LineEnd, |
| 123 | hit.ChunkCharCount, |
| 124 | hit.LastWriteUtcIso); |
| 125 | } |
| 126 | |
| 127 | var ordered = merged.Values |
| 128 | .OrderByDescending(static h => h.RankScore) |
| 129 | .Take(topN) |
| 130 | .ToList(); |
| 131 | |
| 132 | return (new SearchResponse(FormatVersion, userQuery, dbPath, ordered), null); |
| 133 | } |
| 134 | |
| 135 | /// <summary>Вектор не должен «протекать» для расширений вне FTS-правил до следующего vec_reindex.</summary> |
| 136 | private static List<(long chunkRowId, double sim)> FilterVecHitsByAllowedChunkExtensions( |
| 137 | SqliteConnection conn, |
| 138 | IndexSettings settings, |
| 139 | List<(long chunkRowId, double sim)> hits) |
| 140 | { |
| 141 | if (hits.Count == 0) |
| 142 | return hits; |
| 143 | |
| 144 | var allowed = settings.GetEffectiveVecExtensionsSet(); |
| 145 | if (allowed.Count == 0) |
| 146 | return []; |
| 147 | |
| 148 | var distinctIds = hits.Select(static h => h.chunkRowId).Distinct().ToArray(); |
| 149 | using var cmd = conn.CreateCommand(); |
| 150 | var sb = new StringBuilder("SELECT rowid, extension FROM chunks WHERE rowid IN ("); |
| 151 | for (var i = 0; i < distinctIds.Length; i++) |
| 152 | { |
| 153 | if (i > 0) |
| 154 | sb.Append(','); |
| 155 | var p = "$i" + i.ToString(CultureInfo.InvariantCulture); |
| 156 | sb.Append(p); |
| 157 | cmd.Parameters.AddWithValue(p, distinctIds[i]); |
| 158 | } |
| 159 | |
| 160 | sb.Append(')'); |
| 161 | |
| 162 | cmd.CommandText = sb.ToString(); |
| 163 | |
| 164 | var extOf = new Dictionary<long, string>(distinctIds.Length); |
| 165 | using (var rr = cmd.ExecuteReader()) |
| 166 | { |
| 167 | while (rr.Read()) |
| 168 | { |
| 169 | var id = rr.GetInt64(0); |
| 170 | var ext = rr.IsDBNull(1) ? "" : rr.GetString(1); |
| 171 | extOf[id] = ext; |
| 172 | } |
| 173 | } |
| 174 | |
| 175 | return hits.Where(h => |
| 176 | extOf.TryGetValue(h.chunkRowId, out var ext) && allowed.Contains(ext)).ToList(); |
| 177 | } |
| 178 | |
| 179 | private static bool TableExists(SqliteConnection conn, string name) |
| 180 | { |
| 181 | using var cmd = conn.CreateCommand(); |
| 182 | cmd.CommandText = "SELECT 1 FROM sqlite_master WHERE type='table' AND name=$n LIMIT 1;"; |
| 183 | cmd.Parameters.AddWithValue("$n", name); |
| 184 | return cmd.ExecuteScalar() is not null; |
| 185 | } |
| 186 | |
| 187 | private static List<(long chunkRowId, double sim)> VecTopK(SqliteConnection conn, float[] qv, double qn, int k, IndexSettings settings) |
| 188 | { |
| 189 | var allowed = settings.GetEffectiveVecExtensionsSet(); |
| 190 | if (allowed.Count == 0) |
| 191 | return []; |
| 192 | |
| 193 | using var cmd = conn.CreateCommand(); |
| 194 | var sb = new StringBuilder( |
| 195 | """ |
| 196 | SELECT v.chunk_rowid, v.dim, v.norm, v.vec FROM vectors v |
| 197 | INNER JOIN chunks c ON c.rowid = v.chunk_rowid |
| 198 | WHERE lower(c.extension) IN ( |
| 199 | """); |
| 200 | |
| 201 | var allowedList = allowed.Select(static e => e.ToLowerInvariant()).OrderBy(static e => e, StringComparer.Ordinal).ToArray(); |
| 202 | for (var i = 0; i < allowedList.Length; i++) |
| 203 | { |
| 204 | if (i > 0) |
| 205 | sb.Append(','); |
| 206 | var p = "$e" + i.ToString(CultureInfo.InvariantCulture); |
| 207 | sb.Append(p); |
| 208 | cmd.Parameters.AddWithValue(p, allowedList[i]); |
| 209 | } |
| 210 | |
| 211 | sb.Append(");"); |
| 212 | cmd.CommandText = sb.ToString(); |
| 213 | |
| 214 | var top = new List<(long id, double sim)>(capacity: k); |
| 215 | using var r = cmd.ExecuteReader(); |
| 216 | while (r.Read()) |
| 217 | { |
| 218 | var id = r.GetInt64(0); |
| 219 | var dim = r.GetInt32(1); |
| 220 | var norm = r.GetDouble(2); |
| 221 | if (dim != qv.Length || norm <= 1e-12) |
| 222 | continue; |
| 223 | var blob = (byte[])r.GetValue(3); |
| 224 | var sim = DotBlob(blob, qv) / (qn * norm); |
| 225 | InsertTopK(top, (id, sim), k); |
| 226 | } |
| 227 | |
| 228 | return top.Select(static x => (x.id, x.sim)).ToList(); |
| 229 | } |
| 230 | |
| 231 | private static void InsertTopK(List<(long id, double sim)> top, (long id, double sim) item, int k) |
| 232 | { |
| 233 | if (top.Count < k) |
| 234 | { |
| 235 | top.Add(item); |
| 236 | top.Sort(static (a, b) => b.sim.CompareTo(a.sim)); |
| 237 | return; |
| 238 | } |
| 239 | |
| 240 | if (item.sim <= top[^1].sim) |
| 241 | return; |
| 242 | |
| 243 | top[^1] = item; |
| 244 | top.Sort(static (a, b) => b.sim.CompareTo(a.sim)); |
| 245 | } |
| 246 | |
| 247 | private static double DotBlob(byte[] blob, float[] qv) |
| 248 | { |
| 249 | double sum = 0; |
| 250 | for (var i = 0; i < qv.Length; i++) |
| 251 | { |
| 252 | var f = BinaryPrimitives.ReadSingleLittleEndian(blob.AsSpan(i * 4, 4)); |
| 253 | sum += f * qv[i]; |
| 254 | } |
| 255 | return sum; |
| 256 | } |
| 257 | |
| 258 | private static double Norm(float[] v) |
| 259 | { |
| 260 | double sum = 0; |
| 261 | foreach (var x in v) |
| 262 | sum += x * x; |
| 263 | return Math.Sqrt(sum); |
| 264 | } |
| 265 | |
| 266 | private static bool TryGetChunk(SqliteConnection conn, long id, out IndexHit hit) |
| 267 | { |
| 268 | using var cmd = conn.CreateCommand(); |
| 269 | cmd.CommandText = """ |
| 270 | SELECT c.rowid, c.path, c.extension, c.line_start, c.line_end, length(c.body), snippet(chunks, 4, '[', ']', ' … ', 24), fs.last_write_utc_ticks |
| 271 | FROM chunks c |
| 272 | LEFT JOIN file_state fs ON fs.path = c.path |
| 273 | WHERE c.rowid = $id |
| 274 | LIMIT 1; |
| 275 | """; |
| 276 | cmd.Parameters.AddWithValue("$id", id); |
| 277 | |
| 278 | using var r = cmd.ExecuteReader(); |
| 279 | if (!r.Read()) |
| 280 | { |
| 281 | hit = null!; |
| 282 | return false; |
| 283 | } |
| 284 | |
| 285 | var path = r.GetString(1); |
| 286 | var ext = r.IsDBNull(2) ? "" : r.GetString(2); |
| 287 | var ls = r.IsDBNull(3) ? 0 : r.GetInt32(3); |
| 288 | var le = r.IsDBNull(4) ? 0 : r.GetInt32(4); |
| 289 | var chars = r.IsDBNull(5) ? 0 : r.GetInt32(5); |
| 290 | var snip = r.IsDBNull(6) ? null : r.GetString(6); |
| 291 | var lastWriteIso = r.IsDBNull(7) ? null : new DateTime(r.GetInt64(7), DateTimeKind.Utc).ToString("O"); |
| 292 | |
| 293 | hit = new IndexHit(id, path, ext, HitKinds.TextVector, 0, FtsScore: null, VecScore: null, snip, ls, le, chars, lastWriteIso); |
| 294 | return true; |
| 295 | } |
| 296 | } |
| 297 | |
| 298 | |