use axum::{
Json,
extract::{Path, Request, State},
http::{StatusCode, header},
response::{IntoResponse, Response},
};
use bytes::Bytes;
use notedthat_core::metrics::{label as metric_label, name as metric, outcome as metric_outcome};
use notedthat_core::{
Error as CoreError, KbSlug, Verb,
search::{SearchRequest, SearchResponse, ValidatedRequest},
};
use notedthat_indexer::KeyPredicate;
use std::time::Instant;
use crate::{
error::{ApiError, ApiErrorResponse},
state::AppState,
};
pub const SEARCH_BODY_MAX_BYTES: usize = 64 * 1024;
struct SearchCall {
kb: String,
started: Instant,
outcome: &'static str,
}
impl SearchCall {
fn begin(kb: String) -> Self {
Self {
kb,
started: Instant::now(),
outcome: metric_outcome::CANCELLED,
}
}
fn finish(&mut self, ok: bool) {
self.outcome = if ok {
metric_outcome::OK
} else {
metric_outcome::ERROR
};
}
}
impl Drop for SearchCall {
fn drop(&mut self) {
metrics::histogram!(metric::SEARCH_DURATION, metric_label::KB => self.kb.clone())
.record(self.started.elapsed().as_secs_f64());
metrics::counter!(
metric::SEARCH_REQUESTS,
metric_label::KB => self.kb.clone(),
metric_label::OUTCOME => self.outcome,
)
.increment(1);
}
}
pub async fn search_kb(
State(state): State<AppState>,
Path(kb_slug_raw): Path<String>,
req: Request,
) -> Result<Response, ApiErrorResponse> {
let request_id = crate::middleware::extract_request_id(&req);
let err = |error: ApiError| ApiErrorResponse {
error,
request_id: request_id.clone(),
};
let kb_slug = KbSlug::try_new(kb_slug_raw).map_err(|e| err(ApiError::Core(e)))?;
let access = crate::authz::KbAccess::resolve(&state, kb_slug.as_str(), &req).map_err(err)?;
access.require_any(Verb::Search).map_err(err)?;
let (parts, body) = req.into_parts();
let body_bytes: Bytes = axum::body::to_bytes(body, SEARCH_BODY_MAX_BYTES)
.await
.map_err(|_| {
err(ApiError::Core(CoreError::PayloadTooLarge {
size: SEARCH_BODY_MAX_BYTES as u64 + 1,
limit: SEARCH_BODY_MAX_BYTES as u64,
}))
})?;
let content_type = parts
.headers
.get(header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok());
if content_type.is_none_or(|value| !value.starts_with("application/json")) {
return Err(err(ApiError::Core(CoreError::InvalidInput {
message: "Content-Type must be application/json".into(),
})));
}
let raw: SearchRequest = serde_json::from_slice(&body_bytes).map_err(|e| {
err(ApiError::Core(CoreError::InvalidInput {
message: format!("invalid request body: {e}"),
}))
})?;
let validated = raw
.validate()
.map_err(|e| err(ApiError::Core(CoreError::from(e))))?;
let response = execute_search(&state, &access, validated)
.await
.map_err(err)?;
Ok((StatusCode::OK, Json(response)).into_response())
}
pub(crate) async fn execute_search(
state: &AppState,
access: &crate::authz::KbAccess,
validated: ValidatedRequest,
) -> Result<SearchResponse, ApiError> {
let kb = access.kb();
let filter = access.filter(Verb::Search);
let allows = |key: &str| filter.allows(key);
let key_filter: Option<KeyPredicate<'_>> = if filter.is_allow_all() {
None
} else {
Some(&allows)
};
let mut search = SearchCall::begin(kb.as_str().to_string());
let searched = state.searcher.search(kb, validated, key_filter).await;
search.finish(searched.is_ok());
let mut response = searched.map_err(|e| ApiError::Core(CoreError::from(e)))?;
response
.hits
.retain(|hit| filter.allows(hit.object_key.as_str()));
#[allow(clippy::cast_precision_loss)]
metrics::histogram!(metric::SEARCH_HITS, metric_label::KB => kb.as_str().to_string())
.record(response.hits.len() as f64);
Ok(response)
}
#[cfg(test)]
mod tests {
use super::*;
use axum::{Router, body::Body, body::to_bytes, http::Request, routing::post};
use std::{collections::BTreeMap, sync::Arc};
use tower::util::ServiceExt;
const KB: &str = "notes";
fn app() -> Router {
let mut kbs = BTreeMap::new();
kbs.insert(KB.to_string(), KbSlug::try_new(KB).unwrap());
let (indexer_tx, _) = tokio::sync::mpsc::channel(1024);
let state = AppState {
storage: Arc::new(crate::testing::InMemoryStorage::default()),
access_policies: Arc::new(BTreeMap::from([(
KB.to_string(),
Arc::new(
[notedthat_core::AccessRule::new(
notedthat_core::Who::Anyone,
[Verb::Search],
)]
.into_iter()
.collect::<notedthat_core::AccessPolicy>(),
),
)])),
kb_details: Arc::new(notedthat_core::slug_kb_details(&kbs)),
declared_kbs: Arc::new(kbs),
authenticator: Arc::new(notedthat_core::Authenticator::new("token")),
max_body_size: 16 * 1024 * 1024,
max_patchable_size: 16 * 1024 * 1024,
indexer_tx: (&indexer_tx).into(),
searcher: Arc::new(crate::testing::NoopSearcher),
events: None,
index_health: Arc::new(notedthat_indexer::IndexHealth::new()),
readiness: crate::testing::ready_receiver(),
reconcile: None,
};
Router::new()
.route("/api/v1/knowledgebases/{kb_slug}/search", post(search_kb))
.with_state(state)
}
fn request(body: impl Into<Body>) -> Request<Body> {
Request::builder()
.method("POST")
.uri(format!("/api/v1/knowledgebases/{KB}/search"))
.header(header::CONTENT_TYPE, "application/json")
.body(body.into())
.unwrap()
}
async fn response_json(response: Response) -> serde_json::Value {
let bytes = to_bytes(response.into_body(), SEARCH_BODY_MAX_BYTES + 1024)
.await
.unwrap();
serde_json::from_slice(&bytes).unwrap()
}
#[tokio::test]
async fn valid_request_returns_200() {
let response = app()
.oneshot(request(r#"{"query":"install cargo"}"#))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let json = response_json(response).await;
assert_eq!(json, serde_json::json!({"hits": []}));
}
#[tokio::test]
async fn missing_content_type_returns_400() {
let response = app()
.oneshot(
Request::builder()
.method("POST")
.uri(format!("/api/v1/knowledgebases/{KB}/search"))
.body(Body::from(r#"{"query":"install cargo"}"#))
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let json = response_json(response).await;
assert_eq!(json["error"], "invalid_request");
assert!(json["request_id"].is_string());
}
#[tokio::test]
async fn empty_query_returns_400() {
let response = app().oneshot(request(r#"{"query":""}"#)).await.unwrap();
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let json = response_json(response).await;
assert_eq!(json["error"], "invalid_request");
assert!(json["message"].as_str().unwrap().contains("query"));
}
#[tokio::test]
async fn body_too_large_returns_413() {
let body = serde_json::json!({"query": "x".repeat(SEARCH_BODY_MAX_BYTES + 1)}).to_string();
let response = app().oneshot(request(body)).await.unwrap();
assert_eq!(response.status(), StatusCode::PAYLOAD_TOO_LARGE);
let json = response_json(response).await;
assert_eq!(json["error"], "payload_too_large");
assert!(json["request_id"].is_string());
}
}