Forge
csharp72c0d4fb
1using System.Buffers.Binary;
2using System.Globalization;
3using System.Text;
4using HybridCodebaseIndex.Core.Embeddings;
5using Microsoft.Data.Sqlite;
6
7namespace HybridCodebaseIndex.Core;
8
9internal 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
View only · write via MCP/CIDE