GCC Code Coverage Report


Directory: src/
File: src/standalone/query.c
Date: 2026-09-30 11:11:31
Exec Total Coverage
Lines: 141 158 89.2%
Functions: 5 5 100.0%
Branches: 55 78 70.5%

Line Branch Exec Source
1 /*
2 * Copyright (c) 2026 Tiger Data, Inc.
3 * Licensed under the PostgreSQL License. See LICENSE for details.
4 *
5 * query.c - Zero-allocation query execution
6 *
7 * All buffers are pre-allocated in PrismQueryCtx. The query hot path
8 * uses only pre-allocated memory and arena reset (no malloc/free).
9 *
10 * For paged posting lists, delegates to PrismQueryState (shared with
11 * PG). Flat posting lists and brute-force remain standalone-only.
12 *
13 * Memory layout:
14 * memctx (long-lived) — owns all query context buffers
15 * └── arena (child) — transient per-query allocations, reset each query
16 */
17
18 #include <math.h>
19 #include <stdio.h>
20 #include <string.h>
21
22 #include "algo/topk.h"
23 #include "algo/vecops.h"
24 #include "core/memory.h"
25 #include "index/posting_page.h"
26 #include "index/posting_scan.h"
27 #include "index/query_scan.h"
28 #include "quant/rabitq.h"
29 #include "standalone/query.h"
30
31 /* ----------------------------------------------------------------
32 * Query context internals
33 * ---------------------------------------------------------------- */
34
35 struct PrismQueryCtx
36 {
37 PrismIndex *idx;
38
39 /* Shared search context (paged mode) */
40 PrismQueryState search;
41 bool has_query_state;
42
43 /* Reranking (standalone-specific) */
44 VsTopK rerank_topk;
45 VsTopKEntry *rerank_buf;
46 uint32_t rerank_cap;
47
48 /* Flat/brute-force fallback buffers */
49 float *query_buf;
50 PrismCentroidResult *beam_results;
51 PrismCentroidScratch *centroid_scratch;
52 VsTopK topk;
53 PrismPostingScan posting_scan;
54 float *pt_query;
55 float *pt_cents_buf;
56 RaBitQQueryState beam_qs;
57 RaBitQQueryState cluster_qs;
58 float *beam_transformed;
59 float *cluster_transformed;
60 uint8_t *beam_query_bits;
61 uint8_t *cluster_query_bits;
62
63 /* Long-lived memory context for all query context buffers.
64 * Deleting this frees everything at once (no individual frees). */
65 VsMemCtx memctx;
66
67 /* Child arena for transient per-query allocations (beam search
68 * internals). Reset per query — no create/delete overhead. */
69 VsMemCtx arena;
70
71 /* Limits */
72 uint32_t max_k;
73 uint32_t max_nprobe;
74 };
75
76 /* ----------------------------------------------------------------
77 * Create / destroy
78 * ---------------------------------------------------------------- */
79
80 PrismQueryCtx *
81 40 prism_query_ctx_create(PrismIndex *idx, uint32_t max_k, uint32_t max_nprobe)
82 {
83
4/6
✓ Branch 0 taken 38 times.
✓ Branch 1 taken 2 times.
✓ Branch 2 taken 38 times.
✗ Branch 3 not taken.
✗ Branch 4 not taken.
✓ Branch 5 taken 38 times.
40 if (idx == NULL || max_k == 0 || max_nprobe == 0)
84 2 return NULL;
85
86 38 VsMemCtx memctx = vs_memctx_create(NULL, "query_ctx");
87 38 VsMemCtx old_ctx = vs_memctx_switch(memctx);
88
89 38 PrismQueryCtx *ctx = vs_alloc0(sizeof(PrismQueryCtx));
90 38 ctx->idx = idx;
91 38 ctx->max_k = max_k;
92 38 ctx->max_nprobe = max_nprobe;
93 38 ctx->memctx = memctx;
94
95 38 Dimension dim = idx->base.dim;
96
97 /* Paged mode: use shared PrismQueryState */
98
3/4
✓ Branch 0 taken 36 times.
✓ Branch 1 taken 2 times.
✓ Branch 2 taken 36 times.
✗ Branch 3 not taken.
38 if (idx->has_posting_data && idx->base.params != NULL &&
99
2/2
✓ Branch 0 taken 30 times.
✓ Branch 1 taken 6 times.
36 idx->posting_fmt == PRISM_POSTING_FMT_PAGES)
100 {
101 30 prism_query_state_init(&ctx->search, &idx->base, max_k, max_nprobe);
102 30 ctx->has_query_state = true;
103
104
2/2
✓ Branch 0 taken 4 times.
✓ Branch 1 taken 26 times.
30 if (idx->base.fastscan)
105 4 prism_posting_scan_enable_fastscan(
106 &ctx->search.pscan, idx->base.fastscan);
107 }
108 else
109 {
110 /* Flat/brute-force: allocate standalone buffers */
111 8 uint32_t packed_bytes = VS_RABITQ_BYTES(dim);
112
113 8 ctx->query_buf = vs_alloc(dim * sizeof(float));
114 8 ctx->beam_results = vs_alloc(max_nprobe * sizeof(PrismCentroidResult));
115 8 ctx->centroid_scratch = prism_centroid_scratch_create(dim, max_nprobe);
116 8 vs_topk_init(&ctx->topk, max_k);
117
118
2/2
✓ Branch 0 taken 6 times.
✓ Branch 1 taken 2 times.
8 if (idx->has_posting_data)
119 {
120 12 VsStorage *storage = (idx->posting_fmt == PRISM_POSTING_FMT_PAGES)
121 ? &idx->posting_storage.base
122
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 6 times.
6 : NULL;
123 12 char *page_base = (idx->posting_fmt == PRISM_POSTING_FMT_PAGES)
124 ? idx->posting_storage.pages
125
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 6 times.
6 : NULL;
126 12 uint32_t max_per_page = (idx->posting_fmt ==
127 PRISM_POSTING_FMT_PAGES)
128 ✗ ? prism_posting_max_entries(dim)
129
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 6 times.
6 : idx->max_cluster_size;
130
131 6 prism_posting_scan_init(
132 &ctx->posting_scan,
133 storage,
134 page_base,
135 6 idx->base.params,
136 dim,
137 max_per_page);
138
139
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 6 times.
6 if (idx->base.fastscan)
140 ✗ prism_posting_scan_enable_fastscan(
141 &ctx->posting_scan, idx->base.fastscan);
142 }
143
144 8 ctx->pt_cents_buf = vs_alloc((size_t)max_nprobe * dim * sizeof(float));
145 8 ctx->pt_query = vs_alloc_aligned(dim * sizeof(float), 64);
146 8 ctx->beam_transformed = vs_alloc_aligned(dim * sizeof(float), 64);
147 8 ctx->cluster_transformed = vs_alloc_aligned(dim * sizeof(float), 64);
148 8 ctx->beam_query_bits = vs_alloc_aligned(packed_bytes, 64);
149 8 ctx->cluster_query_bits = vs_alloc_aligned(packed_bytes, 64);
150
151 8 ctx->beam_qs.transformed = ctx->beam_transformed;
152 8 ctx->beam_qs.query_bits = ctx->beam_query_bits;
153 8 ctx->cluster_qs.transformed = ctx->cluster_transformed;
154 8 ctx->cluster_qs.query_bits = ctx->cluster_query_bits;
155
156 8 vs_rabitq_init_query_constants(&ctx->beam_qs, dim);
157 8 vs_rabitq_init_query_constants(&ctx->cluster_qs, dim);
158 }
159
160 /* Reranking buffers (used for both paged and flat) */
161 38 vs_topk_init(&ctx->rerank_topk, max_k);
162 38 ctx->rerank_cap = max_k * 16;
163 38 ctx->rerank_buf = vs_alloc(ctx->rerank_cap * sizeof(VsTopKEntry));
164
165 /* Child arena for transient per-query allocations */
166 38 ctx->arena = vs_memctx_create(memctx, "query_arena");
167
168 38 vs_memctx_switch(old_ctx);
169 38 return ctx;
170 }
171
172 void
173 38 prism_query_ctx_destroy(PrismQueryCtx *ctx)
174 {
175
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 38 times.
38 if (ctx == NULL)
176 ✗ return;
177
178
2/2
✓ Branch 0 taken 30 times.
✓ Branch 1 taken 8 times.
38 if (ctx->has_query_state)
179 30 prism_query_state_cleanup(&ctx->search);
180 else
181 {
182
2/2
✓ Branch 0 taken 6 times.
✓ Branch 1 taken 2 times.
8 if (ctx->idx->has_posting_data)
183 6 prism_posting_scan_cleanup(&ctx->posting_scan);
184 8 vs_topk_cleanup(&ctx->topk);
185 8 prism_centroid_scratch_free(ctx->centroid_scratch);
186 8 ctx->centroid_scratch = NULL;
187 }
188
189 38 vs_topk_cleanup(&ctx->rerank_topk);
190
191 /* Deleting memctx frees ctx itself, all buffers, and the arena */
192 38 vs_memctx_delete(ctx->memctx);
193 }
194
195 /* ----------------------------------------------------------------
196 * Paged query (via shared PrismQueryState)
197 * ---------------------------------------------------------------- */
198
199 static uint32_t
200 3704 exec_paged(
201 PrismQueryCtx *ctx,
202 const float *query,
203 uint32_t k,
204 uint32_t nprobe,
205 VsDistanceMode mode,
206 bool rerank,
207 uint32_t *result_ids)
208 {
209 3704 VsMemCtx old = vs_memctx_switch(ctx->memctx);
210 3704 prism_query_execute(&ctx->search, query, k, nprobe, mode, rerank, NULL);
211 3704 vs_memctx_switch(old);
212
213 3704 uint32_t nresults = ctx->search.nresults;
214
2/2
✓ Branch 0 taken 4836 times.
✓ Branch 1 taken 3704 times.
8540 for (uint32_t i = 0; i < nresults; i++)
215 {
216 4836 uint32_t ci = ctx->search.result_order[i];
217 4836 result_ids[i] = prism_posting_decode_vector_id(
218 4836 ctx->search.candidates[ci].id);
219 }
220
221 3704 return nresults;
222 }
223
224 /* ----------------------------------------------------------------
225 * Flat/brute-force query (standalone-only fallback)
226 * ---------------------------------------------------------------- */
227
228 static uint32_t
229 24 exec_fallback(
230 PrismQueryCtx *ctx,
231 const float *query,
232 uint32_t k,
233 uint32_t nprobe,
234 VsDistanceMode mode,
235 bool rerank,
236 uint32_t *result_ids)
237 {
238 24 PrismIndex *idx = ctx->idx;
239 24 Dimension dim = idx->base.dim;
240
241
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 24 times.
24 if (k > ctx->max_k)
242 ✗ k = ctx->max_k;
243
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 24 times.
24 if (nprobe > ctx->max_nprobe)
244 ✗ nprobe = ctx->max_nprobe;
245
246 24 VsMemCtx old_ctx = vs_memctx_switch(ctx->arena);
247 24 vs_memctx_reset(ctx->arena);
248
249 /* Normalize query */
250 24 const float *qvec = query;
251
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 24 times.
24 if (idx->base.metric == DISTANCE_COSINE)
252 {
253 ✗ memcpy(ctx->query_buf, query, dim * sizeof(float));
254 ✗ float norm = vs_l2_norm(ctx->query_buf, dim);
255 ✗ if (norm > 0.0f)
256 ✗ vec32_scale(ctx->query_buf, 1.0f / norm, ctx->query_buf, dim);
257 ✗ qvec = ctx->query_buf;
258 }
259
260 /* Rotate query */
261 24 RaBitQQueryState *qs = NULL;
262
1/2
✓ Branch 0 taken 24 times.
✗ Branch 1 not taken.
24 if (idx->base.params != NULL)
263 {
264 24 vs_rabitq_rotate(idx->base.params, qvec, ctx->pt_query);
265
266
2/2
✓ Branch 0 taken 22 times.
✓ Branch 1 taken 2 times.
24 if (idx->base.centroid_format == PRISM_CENTROID_FMT_RABITQ)
267 {
268 22 vs_rabitq_init_query_state(
269 &ctx->beam_qs,
270 22 ctx->pt_query,
271 22 idx->base.pt_global_mean,
272 dim,
273 mode);
274 22 qs = &ctx->beam_qs;
275 }
276 }
277
278 /* Beam search */
279 24 PrismCentroidSearchState state = {
280 .qstate = qs,
281 .query = qvec,
282 24 .storage = &idx->centroid_storage.base,
283 .beam_width = nprobe,
284 .nprobe = nprobe,
285 .dim = dim,
286 24 .metric = idx->base.metric,
287 .error_scale = 0.0f,
288 24 .scratch = ctx->centroid_scratch,
289 };
290
291 24 PrismCentroidSearchStats beam_stats = {0};
292 24 uint32_t n_results = prism_centroid_beam_search(
293 &state,
294 idx->base.first_centroid,
295 24 idx->base.nlevels,
296 ctx->beam_results,
297 NULL,
298 &beam_stats);
299
300 24 vs_memctx_switch(old_ctx);
301
302 /* Reset top-K */
303 24 vs_topk_reset_to_k(&ctx->topk, k);
304
305 /* Flat mode or brute-force */
306
3/4
✓ Branch 0 taken 22 times.
✓ Branch 1 taken 2 times.
✓ Branch 2 taken 22 times.
✗ Branch 3 not taken.
24 if (idx->has_posting_data && idx->base.params != NULL)
307 {
308
2/2
✓ Branch 0 taken 220 times.
✓ Branch 1 taken 22 times.
242 for (uint32_t j = 0; j < n_results; j++)
309 {
310 220 uint32_t li = (uint32_t)ctx->beam_results[j].posting_head;
311
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 220 times.
220 if (li >= idx->nlist)
312 ✗ continue;
313
314 220 const float *pt_cent = idx->pt_centroids + (size_t)li * dim;
315 220 vs_rabitq_init_query_state(
316 220 &ctx->cluster_qs, ctx->pt_query, pt_cent, dim, mode);
317
318 220 prism_posting_scan_begin_flat(
319 220 &ctx->posting_scan, &ctx->cluster_qs, idx->flat_pages[li]);
320 220 prism_posting_scan_cluster(&ctx->posting_scan, &ctx->topk);
321 220 prism_posting_scan_end_cluster(&ctx->posting_scan);
322 }
323 }
324 else
325 {
326
2/2
✓ Branch 0 taken 10 times.
✓ Branch 1 taken 2 times.
12 for (uint32_t j = 0; j < n_results; j++)
327 {
328 10 uint32_t li = (uint32_t)ctx->beam_results[j].posting_head;
329
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 10 times.
10 if (li >= idx->nlist)
330 ✗ continue;
331
332 10 PrismClusterList *cl = &idx->clusters[li];
333
2/2
✓ Branch 0 taken 232 times.
✓ Branch 1 taken 10 times.
242 for (uint32_t vi = 0; vi < cl->count; vi++)
334 {
335 232 uint32_t vid = cl->ids[vi];
336 232 const float *vec = idx->all_vectors + (size_t)vid * dim;
337 232 Distance d = vs_l2_distance_squared(qvec, vec, dim);
338 232 vs_topk_insert(&ctx->topk, d, 0.0f, vid);
339 }
340 }
341 }
342
343 /* Extract results */
344 uint32_t count;
345
1/2
✓ Branch 0 taken 24 times.
✗ Branch 1 not taken.
24 if (rerank)
346 {
347
2/2
✓ Branch 0 taken 4 times.
✓ Branch 1 taken 20 times.
24 if (ctx->topk.cand_count > ctx->rerank_cap)
348 {
349 4 ctx->rerank_cap = ctx->topk.cand_count;
350 4 ctx->rerank_buf = vs_realloc(
351 4 ctx->rerank_buf, ctx->rerank_cap * sizeof(VsTopKEntry));
352 }
353
354 uint32_t n_cands;
355 24 vs_topk_extract_sorted(&ctx->topk, ctx->rerank_buf, &n_cands);
356
357 24 vs_topk_reset_to_k(&ctx->rerank_topk, k);
358
359
2/2
✓ Branch 0 taken 7902 times.
✓ Branch 1 taken 24 times.
7926 for (uint32_t i = 0; i < n_cands; i++)
360 {
361 7902 uint32_t vid = prism_posting_decode_vector_id(
362 7902 ctx->rerank_buf[i].id);
363 7902 const float *vec = idx->all_vectors + (size_t)vid * dim;
364 7902 Distance d = vs_l2_distance_squared(qvec, vec, dim);
365 7902 vs_topk_insert(&ctx->rerank_topk, d, 0.0f, ctx->rerank_buf[i].id);
366 }
367
368 24 vs_topk_extract_sorted(&ctx->rerank_topk, ctx->rerank_buf, &count);
369 }
370 else
371 {
372 ✗ if (ctx->topk.cand_count > ctx->rerank_cap)
373 {
374 ✗ ctx->rerank_cap = ctx->topk.cand_count;
375 ✗ ctx->rerank_buf = vs_realloc(
376 ✗ ctx->rerank_buf, ctx->rerank_cap * sizeof(VsTopKEntry));
377 }
378
379 ✗ vs_topk_extract_sorted(&ctx->topk, ctx->rerank_buf, &count);
380 }
381
382 24 uint32_t out = count < k ? count : k;
383
2/2
✓ Branch 0 taken 220 times.
✓ Branch 1 taken 24 times.
244 for (uint32_t i = 0; i < out; i++)
384 220 result_ids[i] = prism_posting_decode_vector_id(ctx->rerank_buf[i].id);
385
386 24 return out;
387 }
388
389 /* ----------------------------------------------------------------
390 * Query execution — dispatch to paged or fallback
391 * ---------------------------------------------------------------- */
392
393 uint32_t
394 3730 prism_query_exec(
395 PrismQueryCtx *ctx,
396 const float *query,
397 uint32_t k,
398 uint32_t nprobe,
399 VsDistanceMode mode,
400 bool rerank,
401 uint32_t *result_ids)
402 {
403
5/8
✓ Branch 0 taken 3728 times.
✓ Branch 1 taken 2 times.
✓ Branch 2 taken 3728 times.
✗ Branch 3 not taken.
✓ Branch 4 taken 3728 times.
✗ Branch 5 not taken.
✗ Branch 6 not taken.
✓ Branch 7 taken 3728 times.
3730 if (ctx == NULL || query == NULL || result_ids == NULL || k == 0)
404 2 return 0;
405
406
2/2
✓ Branch 0 taken 3704 times.
✓ Branch 1 taken 24 times.
3728 if (ctx->has_query_state)
407 3704 return exec_paged(ctx, query, k, nprobe, mode, rerank, result_ids);
408 else
409 24 return exec_fallback(ctx, query, k, nprobe, mode, rerank, result_ids);
410 }
411