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}