| Line | Branch | Exec | Source |
|---|---|---|---|
| 1 | /* | ||
| 2 | * Copyright (c) 2026 Tiger Data, Inc. | ||
| 3 | * Licensed under the PostgreSQL License. See LICENSE for details. | ||
| 4 | * | ||
| 5 | * api.c - Public C API for standalone pg_vectorsearch library | ||
| 6 | * | ||
| 7 | * Thin wrapper around PrismIndex (standalone/index.h) and PrismQueryCtx | ||
| 8 | * (standalone/query.h). The handle bundles both so callers get a | ||
| 9 | * single opaque pointer. | ||
| 10 | * | ||
| 11 | * A top-level memory context owns the handle. The index and query | ||
| 12 | * context create their own child contexts internally. | ||
| 13 | */ | ||
| 14 | |||
| 15 | #include <string.h> | ||
| 16 | |||
| 17 | #include "core/memory.h" | ||
| 18 | #include "standalone/api.h" | ||
| 19 | #include "standalone/index.h" | ||
| 20 | #include "standalone/query.h" | ||
| 21 | |||
| 22 | struct VsHandle | ||
| 23 | { | ||
| 24 | VsMemCtx memctx; /* top-level context, owns everything */ | ||
| 25 | PrismIndex *idx; | ||
| 26 | PrismQueryCtx *qctx; | ||
| 27 | }; | ||
| 28 | |||
| 29 | /* ---------------------------------------------------------------- | ||
| 30 | * Parse helpers | ||
| 31 | * ---------------------------------------------------------------- */ | ||
| 32 | |||
| 33 | static DistanceMetric | ||
| 34 | 12 | parse_metric(const char *s) | |
| 35 | { | ||
| 36 |
1/2✗ Branch 0 not taken.
✓ Branch 1 taken 12 times.
|
12 | if (s == NULL) |
| 37 | ✗ | return DISTANCE_L2; | |
| 38 |
3/4✓ Branch 0 taken 10 times.
✓ Branch 1 taken 2 times.
✗ Branch 2 not taken.
✓ Branch 3 taken 10 times.
|
12 | if (strcmp(s, "angular") == 0 || strcmp(s, "cosine") == 0) |
| 39 | 2 | return DISTANCE_COSINE; | |
| 40 | 10 | return DISTANCE_L2; | |
| 41 | } | ||
| 42 | |||
| 43 | static PrismCentroidFormat | ||
| 44 | 12 | parse_centroid_fmt(const char *s) | |
| 45 | { | ||
| 46 |
1/2✗ Branch 0 not taken.
✓ Branch 1 taken 12 times.
|
12 | if (s == NULL) |
| 47 | ✗ | return PRISM_CENTROID_FMT_RABITQ; | |
| 48 |
1/2✗ Branch 0 not taken.
✓ Branch 1 taken 12 times.
|
12 | if (strcmp(s, "float32") == 0) |
| 49 | ✗ | return PRISM_CENTROID_FMT_FLOAT; | |
| 50 |
1/2✗ Branch 0 not taken.
✓ Branch 1 taken 12 times.
|
12 | if (strcmp(s, "float16") == 0) |
| 51 | ✗ | return PRISM_CENTROID_FMT_HALF; | |
| 52 |
1/2✗ Branch 0 not taken.
✓ Branch 1 taken 12 times.
|
12 | if (strcmp(s, "fastscan") == 0) |
| 53 | ✗ | return PRISM_CENTROID_FMT_FASTSCAN; | |
| 54 | 12 | return PRISM_CENTROID_FMT_RABITQ; | |
| 55 | } | ||
| 56 | |||
| 57 | static PrismPostingFormat | ||
| 58 | 12 | parse_posting_fmt(const char *s) | |
| 59 | { | ||
| 60 |
1/4✗ Branch 0 not taken.
✓ Branch 1 taken 12 times.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
|
12 | if (s != NULL && strcmp(s, "flat") == 0) |
| 61 | ✗ | return PRISM_POSTING_FMT_FLAT; | |
| 62 | 12 | return PRISM_POSTING_FMT_PAGES; /* default: pages (matches PG on-disk) */ | |
| 63 | } | ||
| 64 | |||
| 65 | static VsDistanceMode | ||
| 66 | 8 | parse_distance_mode(const char *s) | |
| 67 | { | ||
| 68 |
3/4✓ Branch 0 taken 8 times.
✗ Branch 1 not taken.
✓ Branch 2 taken 2 times.
✓ Branch 3 taken 6 times.
|
8 | if (s != NULL && strcmp(s, "symmetric") == 0) |
| 69 | 2 | return VS_DISTANCE_MODE_SYMMETRIC; | |
| 70 | 6 | return VS_DISTANCE_MODE_ASYMMETRIC; | |
| 71 | } | ||
| 72 | |||
| 73 | /* ---------------------------------------------------------------- | ||
| 74 | * API | ||
| 75 | * ---------------------------------------------------------------- */ | ||
| 76 | |||
| 77 | VsHandle * | ||
| 78 | 12 | vs_handle_create( | |
| 79 | Vec32Source *src, | ||
| 80 | uint32_t nlist, | ||
| 81 | uint32_t fan_out, | ||
| 82 | const char *metric, | ||
| 83 | const char *centroid_fmt, | ||
| 84 | const char *posting_fmt, | ||
| 85 | uint32_t km_nredo, | ||
| 86 | uint32_t km_max_iter, | ||
| 87 | double soar_lambda, | ||
| 88 | double boundary_epsilon, | ||
| 89 | int fastscan, | ||
| 90 | int32_t nworkers, | ||
| 91 | PrismBuildInfo *info) | ||
| 92 | { | ||
| 93 | /* Create context as child of the current context if one exists, | ||
| 94 | * otherwise as a top-level context. This works both when called | ||
| 95 | * from the CLI (which sets up a cli_memctx) and from Python | ||
| 96 | * ctypes (where no context is set). */ | ||
| 97 | 12 | VsMemCtx parent = vs_current_memctx; | |
| 98 | 12 | VsMemCtx memctx = vs_memctx_create(parent, "vs_handle"); | |
| 99 | 12 | VsMemCtx old_ctx = vs_memctx_switch(memctx); | |
| 100 | |||
| 101 | 12 | PrismCentroidFormat fmt = parse_centroid_fmt(centroid_fmt); | |
| 102 | |||
| 103 | 36 | PrismIndexConfig config = { | |
| 104 | .nlist = nlist, | ||
| 105 | .fan_out = fan_out, | ||
| 106 | .centroid_fmt = fmt, | ||
| 107 | 12 | .metric = parse_metric(metric), | |
| 108 | .km_nredo = km_nredo, | ||
| 109 | .km_max_iter = km_max_iter, | ||
| 110 | .soar_lambda = soar_lambda, | ||
| 111 | .boundary_epsilon = boundary_epsilon, | ||
| 112 | .fastscan = fastscan, | ||
| 113 | .nworkers = nworkers, | ||
| 114 | .encode_rabitq = true, | ||
| 115 | 12 | .posting_fmt = parse_posting_fmt(posting_fmt), | |
| 116 | }; | ||
| 117 | |||
| 118 | PrismBuildStats build_stats; | ||
| 119 | 12 | PrismIndex *idx = prism_index_build(src, &config, &build_stats); | |
| 120 |
1/2✗ Branch 0 not taken.
✓ Branch 1 taken 12 times.
|
12 | if (idx == NULL) |
| 121 | { | ||
| 122 | ✗ | vs_memctx_switch(old_ctx); | |
| 123 | ✗ | vs_memctx_delete(memctx); | |
| 124 | ✗ | return NULL; | |
| 125 | } | ||
| 126 | |||
| 127 | /* Fill build stats if requested */ | ||
| 128 |
2/2✓ Branch 0 taken 4 times.
✓ Branch 1 taken 8 times.
|
12 | if (info != NULL) |
| 129 | { | ||
| 130 | 4 | info->stats = build_stats; | |
| 131 | 4 | info->nlist = idx->nlist; | |
| 132 | 4 | info->nlevels = idx->base.nlevels; | |
| 133 | 4 | info->nvecs = idx->nvecs; | |
| 134 | 4 | info->max_cluster = 0; | |
| 135 | 4 | info->min_cluster = UINT32_MAX; | |
| 136 |
2/2✓ Branch 0 taken 36 times.
✓ Branch 1 taken 4 times.
|
40 | for (uint32_t c = 0; c < idx->nlist; c++) |
| 137 | { | ||
| 138 | 36 | uint32_t sz = idx->clusters[c].count; | |
| 139 |
1/2✗ Branch 0 not taken.
✓ Branch 1 taken 36 times.
|
36 | if (sz > info->max_cluster) |
| 140 | ✗ | info->max_cluster = sz; | |
| 141 |
2/2✓ Branch 0 taken 4 times.
✓ Branch 1 taken 32 times.
|
36 | if (sz < info->min_cluster) |
| 142 | 4 | info->min_cluster = sz; | |
| 143 | } | ||
| 144 |
1/2✗ Branch 0 not taken.
✓ Branch 1 taken 4 times.
|
4 | if (idx->nlist == 0) |
| 145 | { | ||
| 146 | /* No centroids, so there is nothing to estimate from. */ | ||
| 147 | ✗ | info->min_cluster = 0; | |
| 148 | } | ||
| 149 |
1/2✓ Branch 0 taken 4 times.
✗ Branch 1 not taken.
|
4 | else if (info->max_cluster == 0) |
| 150 | { | ||
| 151 | /* Pages mode skips cluster list construction — | ||
| 152 | * fall back to estimates */ | ||
| 153 | 4 | info->max_cluster = (idx->nvecs + idx->nlist - 1) / idx->nlist; | |
| 154 | 4 | info->min_cluster = idx->nvecs / idx->nlist; | |
| 155 | } | ||
| 156 | } | ||
| 157 | |||
| 158 | /* Pre-allocate query context: k up to 100, nprobe up to 400 */ | ||
| 159 | 12 | PrismQueryCtx *qctx = prism_query_ctx_create(idx, 100, 400); | |
| 160 |
1/2✗ Branch 0 not taken.
✓ Branch 1 taken 12 times.
|
12 | if (qctx == NULL) |
| 161 | { | ||
| 162 | ✗ | prism_index_destroy(idx); | |
| 163 | ✗ | vs_memctx_switch(old_ctx); | |
| 164 | ✗ | vs_memctx_delete(memctx); | |
| 165 | ✗ | return NULL; | |
| 166 | } | ||
| 167 | |||
| 168 | 12 | VsHandle *handle = vs_alloc(sizeof(VsHandle)); | |
| 169 | 12 | handle->memctx = memctx; | |
| 170 | 12 | handle->idx = idx; | |
| 171 | 12 | handle->qctx = qctx; | |
| 172 | |||
| 173 | 12 | vs_memctx_switch(old_ctx); | |
| 174 | |||
| 175 | 12 | return handle; | |
| 176 | } | ||
| 177 | |||
| 178 | VsHandle * | ||
| 179 | 12 | vs_handle_create_from_array( | |
| 180 | const float *vectors, | ||
| 181 | uint32_t nvecs, | ||
| 182 | uint32_t dim, | ||
| 183 | uint32_t nlist, | ||
| 184 | uint32_t fan_out, | ||
| 185 | const char *metric, | ||
| 186 | const char *centroid_fmt, | ||
| 187 | const char *posting_fmt, | ||
| 188 | uint32_t km_nredo, | ||
| 189 | uint32_t km_max_iter, | ||
| 190 | double soar_lambda, | ||
| 191 | double boundary_epsilon, | ||
| 192 | int fastscan, | ||
| 193 | int32_t nworkers, | ||
| 194 | PrismBuildInfo *info) | ||
| 195 | { | ||
| 196 | VsArraySource array_src; | ||
| 197 | 12 | vs_array_source_init(&array_src, vectors, nvecs, dim); | |
| 198 | 12 | return vs_handle_create( | |
| 199 | &array_src.base, | ||
| 200 | nlist, | ||
| 201 | fan_out, | ||
| 202 | metric, | ||
| 203 | centroid_fmt, | ||
| 204 | posting_fmt, | ||
| 205 | km_nredo, | ||
| 206 | km_max_iter, | ||
| 207 | soar_lambda, | ||
| 208 | boundary_epsilon, | ||
| 209 | fastscan, | ||
| 210 | nworkers, | ||
| 211 | info); | ||
| 212 | } | ||
| 213 | |||
| 214 | uint32_t | ||
| 215 | 10 | vs_handle_query( | |
| 216 | VsHandle *handle, | ||
| 217 | const float *query, | ||
| 218 | uint32_t k, | ||
| 219 | uint32_t nprobe, | ||
| 220 | const char *distance_mode, | ||
| 221 | bool rerank, | ||
| 222 | uint32_t *result_ids) | ||
| 223 | { | ||
| 224 |
2/2✓ Branch 0 taken 2 times.
✓ Branch 1 taken 8 times.
|
10 | if (handle == NULL) |
| 225 | 2 | return 0; | |
| 226 | |||
| 227 | 8 | VsDistanceMode mode = parse_distance_mode(distance_mode); | |
| 228 | |||
| 229 | 8 | return prism_query_exec( | |
| 230 | handle->qctx, query, k, nprobe, mode, rerank, result_ids); | ||
| 231 | } | ||
| 232 | |||
| 233 | void | ||
| 234 | 14 | vs_handle_destroy(VsHandle *handle) | |
| 235 | { | ||
| 236 |
2/2✓ Branch 0 taken 2 times.
✓ Branch 1 taken 12 times.
|
14 | if (handle == NULL) |
| 237 | 2 | return; | |
| 238 | |||
| 239 | /* Destroy query ctx and index first (they manage their own | ||
| 240 | * child contexts), then delete the top-level context which | ||
| 241 | * frees the handle itself. */ | ||
| 242 | 12 | prism_query_ctx_destroy(handle->qctx); | |
| 243 | 12 | prism_index_destroy(handle->idx); | |
| 244 | |||
| 245 | 12 | VsMemCtx memctx = handle->memctx; | |
| 246 | 12 | vs_memctx_delete(memctx); | |
| 247 | } | ||
| 248 |