Skip to main content

alopex_server/http/
hnsw.rs

1use std::sync::Arc;
2
3use alopex_core::kv::KVTransaction;
4use alopex_core::types::TxnMode;
5use alopex_core::vector::hnsw::{HnswConfig, HnswIndex, HnswSearchResult, HnswStats};
6use alopex_core::vector::Metric;
7use alopex_core::KVStore;
8use axum::extract::Extension;
9use axum::response::Response;
10use axum::Json;
11use serde::{Deserialize, Serialize};
12
13use crate::error::{Result, ServerError};
14use crate::http::{error_response, json_response, RequestContext};
15use crate::server::ServerState;
16
17const DEFAULT_M: usize = 16;
18const DEFAULT_EF_CONSTRUCTION: usize = 200;
19
20#[derive(Debug, Deserialize)]
21pub struct HnswSearchRequest {
22    pub index: String,
23    pub query: Vec<f32>,
24    #[serde(default = "default_k")]
25    pub k: usize,
26}
27
28#[derive(Debug, Deserialize)]
29pub struct HnswUpsertRequest {
30    pub index: String,
31    pub key: Vec<u8>,
32    pub vector: Vec<f32>,
33}
34
35#[derive(Debug, Deserialize)]
36pub struct HnswDeleteRequest {
37    pub index: String,
38    pub key: Vec<u8>,
39}
40
41#[derive(Debug, Deserialize)]
42pub struct HnswCreateRequest {
43    pub index: String,
44    pub dim: usize,
45    pub metric: String,
46}
47
48#[derive(Debug, Deserialize)]
49pub struct HnswDropRequest {
50    pub index: String,
51}
52
53#[derive(Debug, Deserialize)]
54pub struct HnswStatsRequest {
55    pub index: String,
56}
57
58#[derive(Debug, Serialize)]
59pub struct HnswSearchResponse {
60    pub results: Vec<HnswSearchResult>,
61}
62
63#[derive(Debug, Serialize)]
64pub struct HnswStatsResponse {
65    pub stats: HnswStats,
66}
67
68#[derive(Debug, Serialize)]
69pub struct HnswStatusResponse {
70    pub success: bool,
71}
72
73pub async fn search(
74    Extension(state): Extension<Arc<ServerState>>,
75    Extension(ctx): Extension<RequestContext>,
76    Json(request): Json<HnswSearchRequest>,
77) -> Response {
78    match search_impl(state.clone(), request) {
79        Ok(resp) => json_response(resp, state.config.max_response_size, &ctx),
80        Err(err) => error_response(err, &ctx),
81    }
82}
83
84pub async fn upsert(
85    Extension(state): Extension<Arc<ServerState>>,
86    Extension(ctx): Extension<RequestContext>,
87    Json(request): Json<HnswUpsertRequest>,
88) -> Response {
89    match upsert_impl(state.clone(), request) {
90        Ok(resp) => json_response(resp, state.config.max_response_size, &ctx),
91        Err(err) => error_response(err, &ctx),
92    }
93}
94
95pub async fn delete(
96    Extension(state): Extension<Arc<ServerState>>,
97    Extension(ctx): Extension<RequestContext>,
98    Json(request): Json<HnswDeleteRequest>,
99) -> Response {
100    match delete_impl(state.clone(), request) {
101        Ok(resp) => json_response(resp, state.config.max_response_size, &ctx),
102        Err(err) => error_response(err, &ctx),
103    }
104}
105
106pub async fn create(
107    Extension(state): Extension<Arc<ServerState>>,
108    Extension(ctx): Extension<RequestContext>,
109    Json(request): Json<HnswCreateRequest>,
110) -> Response {
111    match create_impl(state.clone(), request) {
112        Ok(resp) => json_response(resp, state.config.max_response_size, &ctx),
113        Err(err) => error_response(err, &ctx),
114    }
115}
116
117pub async fn drop(
118    Extension(state): Extension<Arc<ServerState>>,
119    Extension(ctx): Extension<RequestContext>,
120    Json(request): Json<HnswDropRequest>,
121) -> Response {
122    match drop_impl(state.clone(), request) {
123        Ok(resp) => json_response(resp, state.config.max_response_size, &ctx),
124        Err(err) => error_response(err, &ctx),
125    }
126}
127
128pub async fn stats(
129    Extension(state): Extension<Arc<ServerState>>,
130    Extension(ctx): Extension<RequestContext>,
131    Json(request): Json<HnswStatsRequest>,
132) -> Response {
133    match stats_impl(state.clone(), request) {
134        Ok(resp) => json_response(resp, state.config.max_response_size, &ctx),
135        Err(err) => error_response(err, &ctx),
136    }
137}
138
139fn search_impl(state: Arc<ServerState>, request: HnswSearchRequest) -> Result<HnswSearchResponse> {
140    let mut txn = state.store.begin(TxnMode::ReadOnly)?;
141    let index = HnswIndex::load(&request.index, &mut txn).map_err(map_core_error)?;
142    let (results, _) = index
143        .search(&request.query, request.k, None)
144        .map_err(map_core_error)?;
145    txn.commit_self()?;
146    Ok(HnswSearchResponse { results })
147}
148
149fn upsert_impl(state: Arc<ServerState>, request: HnswUpsertRequest) -> Result<HnswStatusResponse> {
150    let mut txn = state.store.begin(TxnMode::ReadWrite)?;
151    let mut index = HnswIndex::load(&request.index, &mut txn).map_err(map_core_error)?;
152    index
153        .upsert(&request.key, &request.vector, &[])
154        .map_err(map_core_error)?;
155    index.save(&mut txn).map_err(map_core_error)?;
156    txn.commit_self()?;
157    Ok(HnswStatusResponse { success: true })
158}
159
160fn delete_impl(state: Arc<ServerState>, request: HnswDeleteRequest) -> Result<HnswStatusResponse> {
161    let mut txn = state.store.begin(TxnMode::ReadWrite)?;
162    let mut index = HnswIndex::load(&request.index, &mut txn).map_err(map_core_error)?;
163    index.delete(&request.key).map_err(map_core_error)?;
164    index.save(&mut txn).map_err(map_core_error)?;
165    txn.commit_self()?;
166    Ok(HnswStatusResponse { success: true })
167}
168
169fn create_impl(state: Arc<ServerState>, request: HnswCreateRequest) -> Result<HnswStatusResponse> {
170    let metric = parse_metric(&request.metric)?;
171    let config = HnswConfig {
172        dimension: request.dim,
173        metric,
174        m: DEFAULT_M,
175        ef_construction: DEFAULT_EF_CONSTRUCTION,
176    };
177    config.validate().map_err(map_core_error)?;
178
179    let index = HnswIndex::create(&request.index, config).map_err(map_core_error)?;
180    let mut txn = state.store.begin(TxnMode::ReadWrite)?;
181    index.save(&mut txn).map_err(map_core_error)?;
182    txn.commit_self()?;
183    Ok(HnswStatusResponse { success: true })
184}
185
186fn drop_impl(state: Arc<ServerState>, request: HnswDropRequest) -> Result<HnswStatusResponse> {
187    let mut txn = state.store.begin(TxnMode::ReadWrite)?;
188    let index = HnswIndex::load(&request.index, &mut txn).map_err(map_core_error)?;
189    index.drop(&mut txn).map_err(map_core_error)?;
190    txn.commit_self()?;
191    Ok(HnswStatusResponse { success: true })
192}
193
194fn stats_impl(state: Arc<ServerState>, request: HnswStatsRequest) -> Result<HnswStatsResponse> {
195    let mut txn = state.store.begin(TxnMode::ReadOnly)?;
196    let index = HnswIndex::load(&request.index, &mut txn).map_err(map_core_error)?;
197    let stats = index.stats();
198    txn.commit_self()?;
199    Ok(HnswStatsResponse { stats })
200}
201
202fn parse_metric(raw: &str) -> Result<Metric> {
203    match raw {
204        "cosine" => Ok(Metric::Cosine),
205        "l2" => Ok(Metric::L2),
206        "ip" => Ok(Metric::InnerProduct),
207        other => Err(ServerError::BadRequest(format!("unknown metric: {other}"))),
208    }
209}
210
211fn default_k() -> usize {
212    10
213}
214
215fn map_core_error(err: alopex_core::Error) -> ServerError {
216    match err {
217        alopex_core::Error::NotFound => ServerError::NotFound("index not found".into()),
218        alopex_core::Error::InvalidParameter { param, reason } => {
219            ServerError::BadRequest(format!("invalid parameter {param}: {reason}"))
220        }
221        other => ServerError::Core(other),
222    }
223}