1use super::{
4 credential::{Credential, Principal, hash_token, is_bearer_token_byte},
5 protocol::{
6 self, DeleteRequest, ErrorResponse, ExportRequest, ExportResponse, ListResponse,
7 PutRequest, PutResponse, ReadRequest, ReadResponse, RemoteErrorCode, RemoteRole,
8 ScanRequest, ScanResponse, SessionResponse, SyncRequest,
9 },
10};
11use crate::{MemoryError, MemoryKey, MemoryLimits, MemoryRecord, MemoryStore};
12use axum::{
13 Json, Router,
14 body::Body,
15 extract::{DefaultBodyLimit, State, rejection::JsonRejection},
16 http::{HeaderMap, Response, StatusCode, header},
17 response::IntoResponse,
18 routing::{get, post},
19};
20#[cfg(feature = "native-server")]
21use std::time::Duration;
22use std::{
23 collections::{HashMap, HashSet},
24 future::Future,
25 sync::Arc,
26};
27use thiserror::Error;
28#[cfg(feature = "native-server")]
29use tower::limit::ConcurrencyLimitLayer;
30use tracing::info;
31use web_time::Instant;
32
33pub(crate) const MAX_JSON_BODY_BYTES: usize = 2 * 1024 * 1024;
35
36#[cfg(feature = "native-server")]
37const MAX_IN_FLIGHT_REQUESTS: usize = 64;
38#[cfg(feature = "native-server")]
39const STORE_OPERATION_TIMEOUT: Duration = Duration::from_secs(30);
40
41#[derive(Debug, Error)]
43pub enum ServerBuildError {
44 #[error("at least one memory credential is required")]
46 NoCredentials,
47 #[error("duplicate bearer tokens are not allowed")]
49 DuplicateBearerToken,
50}
51
52struct ServerState<S> {
53 store_factory: Arc<dyn Fn(String) -> S + Send + Sync>,
54 principals: HashMap<[u8; 32], Principal>,
55}
56
57#[derive(Clone)]
63pub struct MemoryServer<S> {
64 state: Arc<ServerState<S>>,
65}
66
67impl<S: MemoryStore> MemoryServer<S> {
68 pub fn new(
73 store_factory: impl Fn(String) -> S + Send + Sync + 'static,
74 credentials: impl IntoIterator<Item = Credential>,
75 ) -> Result<Self, ServerBuildError> {
76 let mut principals = HashMap::new();
77 for credential in credentials {
78 let (token_hash, principal) = credential.into_hashed_principal();
79 if principals.insert(token_hash, principal).is_some() {
80 return Err(ServerBuildError::DuplicateBearerToken);
81 }
82 }
83 if principals.is_empty() {
84 return Err(ServerBuildError::NoCredentials);
85 }
86 Ok(Self {
87 state: Arc::new(ServerState {
88 store_factory: Arc::new(store_factory),
89 principals,
90 }),
91 })
92 }
93
94 pub fn router(&self) -> Router {
99 let router = Router::new()
100 .route(&route(protocol::SESSION_PATH), get(session))
101 .route(&route(protocol::SCAN_PATH), post(scan))
102 .route(&route(protocol::READ_PATH), post(read))
103 .route(&route(protocol::LIST_PATH), post(list))
104 .route(&route(protocol::PUT_PATH), post(put))
105 .route(&route(protocol::DELETE_PATH), post(delete))
106 .route(&route(protocol::SYNC_PATH), post(sync))
107 .route(&route(protocol::EXPORT_PATH), post(export))
108 .layer(DefaultBodyLimit::max(MAX_JSON_BODY_BYTES));
109 #[cfg(feature = "native-server")]
110 let router = router.layer(ConcurrencyLimitLayer::new(MAX_IN_FLIGHT_REQUESTS));
111 router.with_state(self.state.clone())
112 }
113}
114
115fn route(path: &str) -> String {
116 format!("/{path}")
117}
118
119async fn session<S: MemoryStore>(
120 State(state): State<Arc<ServerState<S>>>,
121 headers: HeaderMap,
122) -> Response<Body> {
123 let principal = match authenticate(&state, &headers) {
124 Ok(principal) => principal,
125 Err(error) => return error.into_response(),
126 };
127 let operation = OperationTrace::new("session", &principal);
128 let response = Json(SessionResponse {
129 protocol_version: crate::VERSION,
130 namespace: principal.namespace.clone(),
131 role: principal.role,
132 })
133 .into_response();
134 operation.success(OperationCounts::default());
135 response
136}
137
138async fn scan<S: MemoryStore>(
139 State(state): State<Arc<ServerState<S>>>,
140 headers: HeaderMap,
141 payload: Result<Json<ScanRequest>, JsonRejection>,
142) -> Response<Body> {
143 let principal = match authenticate(&state, &headers) {
144 Ok(principal) => principal,
145 Err(error) => return error.into_response(),
146 };
147 let operation = OperationTrace::new("scan", &principal);
148 let request = match json_payload(payload) {
149 Ok(request) => request,
150 Err(error) => return operation.error_response(error, OperationCounts::default()),
151 };
152 let counts = OperationCounts::input(1);
153 if request.query.len() > MemoryLimits::PRODUCTION.query_bytes
154 || request.limit == 0
155 || request.limit > MemoryLimits::PRODUCTION.scan_results
156 {
157 let error = if request.query.len() > MemoryLimits::PRODUCTION.query_bytes {
158 ApiError::new(
159 StatusCode::PAYLOAD_TOO_LARGE,
160 RemoteErrorCode::QueryTooLarge,
161 )
162 } else {
163 ApiError::bad_request()
164 };
165 return operation.error_response(error, counts);
166 }
167 let store = (state.store_factory)(principal.namespace);
168 match run_store(
169 operation,
170 counts,
171 store.scan(&request.query, request.limit),
172 |scan| OperationCounts::candidates(scan.candidates.len()),
173 )
174 .await
175 {
176 Ok(scan) => Json(ScanResponse {
177 candidates: scan.candidates,
178 })
179 .into_response(),
180 Err(error) => error.into_response(),
181 }
182}
183
184async fn read<S: MemoryStore>(
185 State(state): State<Arc<ServerState<S>>>,
186 headers: HeaderMap,
187 payload: Result<Json<ReadRequest>, JsonRejection>,
188) -> Response<Body> {
189 let principal = match authenticate(&state, &headers) {
190 Ok(principal) => principal,
191 Err(error) => return error.into_response(),
192 };
193 let operation = OperationTrace::new("read", &principal);
194 let request = match json_payload(payload) {
195 Ok(request) => request,
196 Err(error) => return operation.error_response(error, OperationCounts::default()),
197 };
198 let counts = OperationCounts::input(request.ids.len().saturating_add(request.keys.len()));
199 if counts.input_count > MemoryLimits::PRODUCTION.records
200 || request.ids.iter().any(|id| *id <= 0)
201 || request.keys.iter().any(|key| !valid_key(key))
202 {
203 return operation.error_response(ApiError::bad_request(), counts);
204 }
205 let store = (state.store_factory)(principal.namespace);
206 match run_store(
207 operation,
208 counts,
209 store.read(&request.ids, &request.keys),
210 |memories| OperationCounts::records(memories.len()),
211 )
212 .await
213 {
214 Ok(memories) => Json(ReadResponse { memories }).into_response(),
215 Err(error) => error.into_response(),
216 }
217}
218
219async fn list<S: MemoryStore>(
220 State(state): State<Arc<ServerState<S>>>,
221 headers: HeaderMap,
222) -> Response<Body> {
223 let principal = match authenticate(&state, &headers) {
224 Ok(principal) => principal,
225 Err(error) => return error.into_response(),
226 };
227 let operation = OperationTrace::new("list", &principal);
228 let store = (state.store_factory)(principal.namespace);
229 match run_store(
230 operation,
231 OperationCounts::default(),
232 store.list(),
233 |memories| OperationCounts::records(memories.len()),
234 )
235 .await
236 {
237 Ok(memories) => Json(ListResponse { memories }).into_response(),
238 Err(error) => error.into_response(),
239 }
240}
241
242async fn put<S: MemoryStore>(
243 State(state): State<Arc<ServerState<S>>>,
244 headers: HeaderMap,
245 payload: Result<Json<PutRequest>, JsonRejection>,
246) -> Response<Body> {
247 let principal = match authenticate(&state, &headers) {
248 Ok(principal) => principal,
249 Err(error) => return error.into_response(),
250 };
251 let operation = OperationTrace::new("put", &principal);
252 if principal.role != RemoteRole::Writer {
253 return operation.error_response(
254 ApiError::new(StatusCode::FORBIDDEN, RemoteErrorCode::Forbidden),
255 OperationCounts::default(),
256 );
257 }
258 let request = match json_payload(payload) {
259 Ok(request) => request,
260 Err(error) => return operation.error_response(error, OperationCounts::default()),
261 };
262 let counts = OperationCounts::input(1);
263 if request.content.trim().is_empty() {
264 return operation.error_response(ApiError::bad_request(), counts);
265 }
266 if request.content.len() > MemoryLimits::PRODUCTION.content_bytes {
267 return operation.error_response(
268 ApiError::new(
269 StatusCode::PAYLOAD_TOO_LARGE,
270 RemoteErrorCode::ContentTooLarge,
271 ),
272 counts,
273 );
274 }
275 if request.replacement.as_ref().is_some_and(|key| {
276 !valid_key(key) || key.namespace.as_deref() != Some(principal.namespace.as_str())
277 }) {
278 return operation.error_response(
279 ApiError::new(StatusCode::FORBIDDEN, RemoteErrorCode::Forbidden),
280 counts,
281 );
282 }
283 let store = (state.store_factory)(principal.namespace);
284 match run_store(
285 operation,
286 counts,
287 store.put(&request.content, request.replacement),
288 |_| OperationCounts::records(1),
289 )
290 .await
291 {
292 Ok(memory) => Json(PutResponse { memory }).into_response(),
293 Err(error) => error.into_response(),
294 }
295}
296
297async fn delete<S: MemoryStore>(
298 State(state): State<Arc<ServerState<S>>>,
299 headers: HeaderMap,
300 payload: Result<Json<DeleteRequest>, JsonRejection>,
301) -> Response<Body> {
302 let principal = match authenticate(&state, &headers) {
303 Ok(principal) => principal,
304 Err(error) => return error.into_response(),
305 };
306 let operation = OperationTrace::new("delete", &principal);
307 if principal.role != RemoteRole::Writer {
308 return operation.error_response(
309 ApiError::new(StatusCode::FORBIDDEN, RemoteErrorCode::Forbidden),
310 OperationCounts::default(),
311 );
312 }
313 let request = match json_payload(payload) {
314 Ok(request) => request,
315 Err(error) => return operation.error_response(error, OperationCounts::default()),
316 };
317 let counts = OperationCounts::input(1);
318 if !valid_key(&request.key)
319 || request.key.namespace.as_deref() != Some(principal.namespace.as_str())
320 {
321 return operation.error_response(
322 ApiError::new(StatusCode::FORBIDDEN, RemoteErrorCode::Forbidden),
323 counts,
324 );
325 }
326 let store = (state.store_factory)(principal.namespace);
327 match run_store(operation, counts, store.delete(request.key), |_| {
328 OperationCounts::records(1)
329 })
330 .await
331 {
332 Ok(()) => Json(serde_json::json!({})).into_response(),
333 Err(error) => error.into_response(),
334 }
335}
336
337async fn sync<S: MemoryStore>(
338 State(state): State<Arc<ServerState<S>>>,
339 headers: HeaderMap,
340 payload: Result<Json<SyncRequest>, JsonRejection>,
341) -> Response<Body> {
342 let principal = match authenticate(&state, &headers) {
343 Ok(principal) => principal,
344 Err(error) => return error.into_response(),
345 };
346 let operation = OperationTrace::new("sync", &principal);
347 if principal.role != RemoteRole::Writer {
348 return operation.error_response(
349 ApiError::new(StatusCode::FORBIDDEN, RemoteErrorCode::Forbidden),
350 OperationCounts::default(),
351 );
352 }
353 let request = match json_payload(payload) {
354 Ok(request) => request,
355 Err(error) => return operation.error_response(error, OperationCounts::default()),
356 };
357 let counts = OperationCounts::input(request.memories.len());
358 if !valid_snapshot(&request.memories) {
359 return operation.error_response(ApiError::bad_request(), counts);
360 }
361 let store = (state.store_factory)(principal.namespace);
362 match run_store(operation, counts, store.sync(&request.memories), |report| {
363 OperationCounts::report(*report)
364 })
365 .await
366 {
367 Ok(report) => Json(report).into_response(),
368 Err(error) => error.into_response(),
369 }
370}
371
372async fn export<S: MemoryStore>(
373 State(state): State<Arc<ServerState<S>>>,
374 headers: HeaderMap,
375 payload: Result<Json<ExportRequest>, JsonRejection>,
376) -> Response<Body> {
377 let principal = match authenticate(&state, &headers) {
378 Ok(principal) => principal,
379 Err(error) => return error.into_response(),
380 };
381 let operation = OperationTrace::new("export", &principal);
382 let request = match json_payload(payload) {
383 Ok(request) => request,
384 Err(error) => return operation.error_response(error, OperationCounts::default()),
385 };
386 let counts = OperationCounts::input(
387 request
388 .namespaces
389 .as_ref()
390 .map_or(0, |namespaces| namespaces.len()),
391 );
392 if request.limit == 0
393 || request.limit > protocol::MAX_EXPORT_PAGE_RECORDS
394 || request.namespaces.as_ref().is_some_and(|namespaces| {
395 namespaces.is_empty()
396 || namespaces.len() > MemoryLimits::PRODUCTION.records
397 || namespaces
398 .iter()
399 .any(|namespace| !protocol::is_valid_namespace(namespace))
400 })
401 || request.cursor.as_ref().is_some_and(|cursor| {
402 !protocol::is_valid_namespace(&cursor.namespace) || cursor.id <= 0
403 })
404 {
405 return operation.error_response(ApiError::bad_request(), counts);
406 }
407 let store = (state.store_factory)(principal.namespace);
408 match run_store(
409 operation,
410 counts,
411 store.export_page(
412 request.namespaces.as_deref(),
413 request.cursor.as_ref(),
414 request.limit,
415 ),
416 |(memories, _)| OperationCounts::records(memories.len()),
417 )
418 .await
419 {
420 Ok((memories, next_cursor)) => Json(ExportResponse {
421 memories,
422 next_cursor,
423 })
424 .into_response(),
425 Err(error) => error.into_response(),
426 }
427}
428
429fn json_payload<T>(payload: Result<Json<T>, JsonRejection>) -> Result<T, ApiError> {
430 payload.map(|Json(value)| value).map_err(|rejection| {
431 let status = if rejection.status() == StatusCode::PAYLOAD_TOO_LARGE {
432 StatusCode::PAYLOAD_TOO_LARGE
433 } else {
434 StatusCode::BAD_REQUEST
435 };
436 ApiError::new(status, RemoteErrorCode::BadRequest)
437 })
438}
439
440fn authenticate<S>(state: &ServerState<S>, headers: &HeaderMap) -> Result<Principal, ApiError> {
441 let authorization = headers
445 .get(header::AUTHORIZATION)
446 .and_then(|value| value.to_str().ok())
447 .ok_or_else(ApiError::unauthorized)?;
448 let (scheme, token) = authorization
449 .split_once(' ')
450 .filter(|(scheme, token)| {
451 scheme.eq_ignore_ascii_case("bearer")
452 && !token.is_empty()
453 && token.bytes().all(is_bearer_token_byte)
454 })
455 .ok_or_else(ApiError::unauthorized)?;
456 let _ = scheme;
457 let token_hash = hash_token(token);
458 let principal = state
459 .principals
460 .get(&token_hash)
461 .cloned()
462 .ok_or_else(ApiError::unauthorized)?;
463 let asserted_namespace = headers
464 .get(protocol::NAMESPACE_HEADER)
465 .and_then(|value| value.to_str().ok())
466 .ok_or_else(ApiError::namespace_mismatch)?;
467 if asserted_namespace != principal.namespace {
468 return Err(ApiError::namespace_mismatch());
469 }
470 Ok(principal)
471}
472
473fn valid_key(key: &MemoryKey) -> bool {
474 key.id > 0
475 && key.version > 0
476 && key
477 .namespace
478 .as_deref()
479 .is_none_or(protocol::is_valid_namespace)
480}
481
482fn valid_snapshot(memories: &[MemoryRecord]) -> bool {
483 let mut ids = HashSet::with_capacity(memories.len());
484 memories.len() <= MemoryLimits::PRODUCTION.records
485 && memories.iter().all(|memory| {
486 memory.key.is_local()
487 && valid_key(&memory.key)
488 && ids.insert(memory.key.id)
489 && !memory.content.trim().is_empty()
490 && memory.content.len() <= MemoryLimits::PRODUCTION.content_bytes
491 && memory.created_at_ms >= 0
492 && memory.updated_at_ms >= memory.created_at_ms
493 })
494 && memories
495 .iter()
496 .map(|memory| memory.content.len())
497 .try_fold(0usize, usize::checked_add)
498 .is_some_and(|bytes| bytes <= MemoryLimits::PRODUCTION.total_content_bytes)
499}
500
501#[derive(Clone, Copy, Default)]
502struct OperationCounts {
503 candidate_count: usize,
504 record_count: usize,
505 input_count: usize,
506 report_inserted: usize,
507 report_replaced: usize,
508 report_unchanged: usize,
509 report_deleted: usize,
510}
511
512impl OperationCounts {
513 fn candidates(candidate_count: usize) -> Self {
514 Self {
515 candidate_count,
516 ..Self::default()
517 }
518 }
519
520 fn records(record_count: usize) -> Self {
521 Self {
522 record_count,
523 ..Self::default()
524 }
525 }
526
527 fn input(input_count: usize) -> Self {
528 Self {
529 input_count,
530 ..Self::default()
531 }
532 }
533
534 fn report(report: protocol::SyncReport) -> Self {
535 Self {
536 report_inserted: report.inserted,
537 report_replaced: report.replaced,
538 report_unchanged: report.unchanged,
539 report_deleted: report.deleted,
540 ..Self::default()
541 }
542 }
543}
544
545struct OperationTrace {
546 operation: &'static str,
547 namespace: String,
548 role: RemoteRole,
549 started_at: Instant,
550}
551
552impl OperationTrace {
553 fn new(operation: &'static str, principal: &Principal) -> Self {
554 Self {
555 operation,
556 namespace: principal.namespace.clone(),
557 role: principal.role,
558 started_at: Instant::now(),
559 }
560 }
561
562 fn success(self, counts: OperationCounts) {
563 info!(
564 operation = self.operation,
565 namespace = %self.namespace,
566 role = ?self.role,
567 success = true,
568 elapsed_ms = self.started_at.elapsed().as_millis() as u64,
569 candidate_count = counts.candidate_count,
570 record_count = counts.record_count,
571 input_count = counts.input_count,
572 report_inserted = counts.report_inserted,
573 report_replaced = counts.report_replaced,
574 report_unchanged = counts.report_unchanged,
575 report_deleted = counts.report_deleted,
576 "remote memory operation"
577 );
578 }
579
580 fn failure(self, error: ApiError, counts: OperationCounts) {
581 info!(
582 operation = self.operation,
583 namespace = %self.namespace,
584 role = ?self.role,
585 success = false,
586 elapsed_ms = self.started_at.elapsed().as_millis() as u64,
587 status = error.status.as_u16(),
588 error_code = ?error.code,
589 candidate_count = counts.candidate_count,
590 record_count = counts.record_count,
591 input_count = counts.input_count,
592 report_inserted = counts.report_inserted,
593 report_replaced = counts.report_replaced,
594 report_unchanged = counts.report_unchanged,
595 report_deleted = counts.report_deleted,
596 "remote memory operation"
597 );
598 }
599
600 fn error_response(self, error: ApiError, counts: OperationCounts) -> Response<Body> {
601 self.failure(error, counts);
602 error.into_response()
603 }
604}
605
606async fn run_store<T, C>(
607 trace: OperationTrace,
608 request_counts: OperationCounts,
609 operation: impl Future<Output = Result<T, MemoryError>> + Send,
610 success_counts: C,
611) -> Result<T, ApiError>
612where
613 T: Send + 'static,
614 C: FnOnce(&T) -> OperationCounts,
615{
616 #[cfg(not(feature = "native-server"))]
617 let result = operation.await;
618 #[cfg(feature = "native-server")]
619 let result = match tokio::time::timeout(STORE_OPERATION_TIMEOUT, operation).await {
620 Err(_) => {
621 let error = ApiError::unavailable();
622 trace.failure(error, request_counts);
623 return Err(error);
624 }
625 Ok(result) => result,
626 };
627
628 match result {
629 Ok(value) => {
630 let mut counts = success_counts(&value);
631 counts.input_count = request_counts.input_count;
632 trace.success(counts);
633 Ok(value)
634 }
635 Err(source) => {
636 let error = ApiError::from(source);
637 trace.failure(error, request_counts);
638 Err(error)
639 }
640 }
641}
642
643#[derive(Clone, Copy)]
644struct ApiError {
645 status: StatusCode,
646 code: RemoteErrorCode,
647}
648
649impl ApiError {
650 const fn new(status: StatusCode, code: RemoteErrorCode) -> Self {
651 Self { status, code }
652 }
653
654 const fn bad_request() -> Self {
655 Self::new(StatusCode::BAD_REQUEST, RemoteErrorCode::BadRequest)
656 }
657
658 const fn unauthorized() -> Self {
659 Self::new(StatusCode::UNAUTHORIZED, RemoteErrorCode::Unauthorized)
660 }
661
662 const fn namespace_mismatch() -> Self {
663 Self::new(StatusCode::FORBIDDEN, RemoteErrorCode::NamespaceMismatch)
664 }
665
666 const fn unavailable() -> Self {
667 Self::new(
668 StatusCode::SERVICE_UNAVAILABLE,
669 RemoteErrorCode::Unavailable,
670 )
671 }
672
673 const fn internal() -> Self {
674 Self::new(StatusCode::INTERNAL_SERVER_ERROR, RemoteErrorCode::Internal)
675 }
676}
677
678impl From<MemoryError> for ApiError {
679 fn from(error: MemoryError) -> Self {
680 if error.is_retryable() {
681 return Self::unavailable();
682 }
683 match error {
684 MemoryError::EmptyContent => Self::bad_request(),
685 MemoryError::ContentTooLarge { .. } => Self::new(
686 StatusCode::PAYLOAD_TOO_LARGE,
687 RemoteErrorCode::ContentTooLarge,
688 ),
689 MemoryError::QueryTooLarge { .. } => Self::new(
690 StatusCode::PAYLOAD_TOO_LARGE,
691 RemoteErrorCode::QueryTooLarge,
692 ),
693 MemoryError::RecordCapacity { .. } => Self::new(
694 StatusCode::INSUFFICIENT_STORAGE,
695 RemoteErrorCode::RecordCapacity,
696 ),
697 MemoryError::ContentCapacity { .. } | MemoryError::StorageCapacity => Self::new(
698 StatusCode::INSUFFICIENT_STORAGE,
699 RemoteErrorCode::ContentCapacity,
700 ),
701 MemoryError::SecretRejected => Self::bad_request(),
702 MemoryError::Duplicate => Self::new(StatusCode::CONFLICT, RemoteErrorCode::Duplicate),
703 MemoryError::NotFound => Self::new(StatusCode::NOT_FOUND, RemoteErrorCode::NotFound),
704 MemoryError::Conflict => Self::new(StatusCode::CONFLICT, RemoteErrorCode::Conflict),
705 MemoryError::RemoteReadOnly => {
706 Self::new(StatusCode::FORBIDDEN, RemoteErrorCode::Forbidden)
707 }
708 MemoryError::UnsupportedSchemaVersion { .. }
709 | MemoryError::InvalidPagination
710 | MemoryError::Backend { .. }
711 | MemoryError::Unavailable { .. } => Self::internal(),
712 }
713 }
714}
715
716impl IntoResponse for ApiError {
717 fn into_response(self) -> Response<Body> {
718 (self.status, Json(ErrorResponse { code: self.code })).into_response()
719 }
720}