pub mod query;
pub mod types;
mod sse;
use std::collections::VecDeque;
pub use futures_util::stream::{Stream, StreamExt};
use reqwest::Method;
use serde::de::DeserializeOwned;
use crate::client::SearchcraftClient;
use crate::config::Operation;
use crate::error;
use self::sse::SseDecoder;
use self::types::{SearchRequest, SearchResponse, SummaryError, SummaryStreamEvent};
impl SearchcraftClient {
pub async fn search_index<T: DeserializeOwned>(
&self,
index_name: &str,
request: &SearchRequest,
) -> error::Result<SearchResponse<T>> {
let path = format!("index/{index_name}/search");
self.transport
.request::<SearchResponse<T>>(Method::POST, &path, Operation::Read, Some(request))
.await
}
pub async fn search_federation<T: DeserializeOwned>(
&self,
federation_name: &str,
request: &SearchRequest,
) -> error::Result<SearchResponse<T>> {
let path = format!("federation/{federation_name}/search");
self.transport
.request::<SearchResponse<T>>(Method::POST, &path, Operation::Read, Some(request))
.await
}
pub async fn search_summary(
&self,
index_name: &str,
request: &SearchRequest,
) -> error::Result<impl Stream<Item = SummaryStreamEvent>> {
let path = format!("index/{index_name}/search/summary");
let response = self
.transport
.request_stream(Method::POST, &path, Operation::Read, Some(request))
.await?;
Ok(summary_stream(Box::pin(response.bytes_stream())))
}
}
struct SummaryState<S> {
bytes: S,
decoder: SseDecoder,
pending: VecDeque<SummaryStreamEvent>,
finished: bool,
}
fn summary_stream<S, B, E>(chunks: S) -> impl Stream<Item = SummaryStreamEvent>
where
S: Stream<Item = Result<B, E>> + Unpin,
B: AsRef<[u8]>,
E: std::fmt::Display,
{
let state = SummaryState {
bytes: chunks,
decoder: SseDecoder::new(),
pending: VecDeque::new(),
finished: false,
};
Box::pin(futures_util::stream::unfold(
state,
|mut state| async move {
loop {
if let Some(event) = state.pending.pop_front() {
return Some((event, state));
}
if state.finished {
return None;
}
match state.bytes.next().await {
Some(Ok(chunk)) => {
for frame in state.decoder.push(chunk.as_ref()) {
if let Some(event) = decode_summary_frame(&frame) {
state.pending.push_back(event);
}
}
}
Some(Err(e)) => {
state.finished = true;
state
.pending
.push_back(SummaryStreamEvent::Error(SummaryError {
message: format!("stream interrupted: {e}"),
}));
}
None => {
state.finished = true;
if let Some(frame) = state.decoder.finish() {
if let Some(event) = decode_summary_frame(&frame) {
state.pending.push_back(event);
}
}
}
}
}
},
))
}
fn decode_summary_frame(frame: &sse::SseEvent) -> Option<SummaryStreamEvent> {
let name = frame.event.as_deref()?;
if !matches!(name, "metadata" | "delta" | "done" | "error") {
return None;
}
let data: serde_json::Value = if frame.data.is_empty() {
serde_json::Value::Object(serde_json::Map::new())
} else {
match serde_json::from_str(&frame.data) {
Ok(value) => value,
Err(_) => {
return Some(SummaryStreamEvent::Error(SummaryError {
message: format!("invalid JSON in {name} event: {}", frame.data),
}))
}
}
};
let tagged = serde_json::json!({ "type": name, "data": data });
match serde_json::from_value(tagged) {
Ok(event) => Some(event),
Err(_) => Some(SummaryStreamEvent::Error(SummaryError {
message: format!("malformed {name} payload: {data}"),
})),
}
}
#[cfg(test)]
mod tests {
use futures_util::StreamExt;
use wiremock::matchers::{body_json, header, method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
use super::summary_stream;
use crate::search::query::QueryBuilder;
use crate::search::types::SummaryStreamEvent;
fn test_client(base_url: &str) -> crate::SearchcraftClient {
crate::SearchcraftClient::new(base_url, Some("test-read-key"), None::<String>).unwrap()
}
#[tokio::test]
async fn search_index_sends_correct_request() {
let server = MockServer::start().await;
let response_body = serde_json::json!({
"status": 200,
"data": {
"hits": [{
"doc": {"title": "Laptop"},
"document_id": "doc-1",
"score": 0.9,
"source_index": "products"
}],
"count": 1,
"time_taken": 5.0
}
});
Mock::given(method("POST"))
.and(path("/index/products/search"))
.and(header("Authorization", "test-read-key"))
.respond_with(ResponseTemplate::new(200).set_body_json(&response_body))
.expect(1)
.mount(&server)
.await;
let client = test_client(&server.uri());
let request = QueryBuilder::fuzzy()
.term("laptop")
.limit(10)
.build_request();
let response = client
.search_index::<serde_json::Value>("products", &request)
.await
.unwrap();
assert_eq!(response.status, 200);
assert_eq!(response.data.count, 1);
assert_eq!(response.data.hits.len(), 1);
assert_eq!(response.data.hits[0].document_id, "doc-1");
}
#[tokio::test]
async fn search_federation_sends_correct_request() {
let server = MockServer::start().await;
let response_body = serde_json::json!({
"status": 200,
"data": {
"hits": [],
"count": 0,
"time_taken": 2.0
}
});
Mock::given(method("POST"))
.and(path("/federation/global/search"))
.and(header("Authorization", "test-read-key"))
.respond_with(ResponseTemplate::new(200).set_body_json(&response_body))
.expect(1)
.mount(&server)
.await;
let client = test_client(&server.uri());
let request = QueryBuilder::exact().term("test").build_request();
let response = client
.search_federation::<serde_json::Value>("global", &request)
.await
.unwrap();
assert_eq!(response.data.count, 0);
assert!(response.data.hits.is_empty());
}
#[tokio::test]
async fn search_index_handles_auth_error() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/index/products/search"))
.respond_with(
ResponseTemplate::new(401)
.set_body_json(serde_json::json!({"message": "unauthorized"})),
)
.mount(&server)
.await;
let client = test_client(&server.uri());
let request = QueryBuilder::fuzzy().term("test").build_request();
let err = client
.search_index::<serde_json::Value>("products", &request)
.await
.unwrap_err();
assert!(matches!(err, crate::error::Error::Authentication { .. }));
}
#[tokio::test]
async fn search_index_handles_not_found() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/index/nonexistent/search"))
.respond_with(
ResponseTemplate::new(404)
.set_body_json(serde_json::json!({"message": "index not found"})),
)
.mount(&server)
.await;
let client = test_client(&server.uri());
let request = QueryBuilder::fuzzy().term("test").build_request();
let err = client
.search_index::<serde_json::Value>("nonexistent", &request)
.await
.unwrap_err();
assert!(matches!(err, crate::error::Error::NotFound(_)));
}
#[tokio::test]
async fn oversized_limit_is_forwarded_for_the_server_to_clamp() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/index/products/search"))
.and(body_json(serde_json::json!({
"query": { "fuzzy": { "ctx": "laptop" } },
"limit": 500
})))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"status": 200,
"data": { "hits": [], "count": 0, "time_taken": 1.0 }
})))
.expect(1)
.mount(&server)
.await;
let client = test_client(&server.uri());
let request = QueryBuilder::fuzzy()
.term("laptop")
.limit(500)
.build_request();
let response = client
.search_index::<serde_json::Value>("products", &request)
.await
.unwrap();
assert_eq!(response.data.count, 0);
}
async fn collect_summary(chunks: Vec<&'static str>) -> Vec<SummaryStreamEvent> {
let stream = futures_util::stream::iter(
chunks
.into_iter()
.map(|c| Ok::<_, std::io::Error>(c.as_bytes())),
);
summary_stream(Box::pin(stream)).collect().await
}
#[tokio::test]
async fn summary_stream_decodes_a_full_generation() {
let events = collect_summary(vec![
"event: metadata\ndata: {\"results_count\":3,\"cached\":false}\n\n",
"event: delta\ndata: {\"content\":\"Gaming \"}\n\n",
"event: delta\ndata: {\"content\":\"laptops\"}\n\n",
"event: done\ndata: {\"results_count\":3}\n\n",
])
.await;
assert_eq!(events.len(), 4);
match &events[0] {
SummaryStreamEvent::Metadata(m) => {
assert_eq!(m.results_count, 3);
assert!(!m.cached);
}
other => panic!("expected metadata, got {other:?}"),
}
let text: String = events
.iter()
.filter_map(|e| match e {
SummaryStreamEvent::Delta(d) => Some(d.content.as_str()),
_ => None,
})
.collect();
assert_eq!(text, "Gaming laptops");
assert!(matches!(events[3], SummaryStreamEvent::Done(_)));
}
#[tokio::test]
async fn summary_stream_reassembles_chunk_split_frames() {
let events = collect_summary(vec![
"event: delta\ndata: {\"cont",
"ent\":\"split\"}\n\nevent: done\ndata: {\"results_count\":1}\n\n",
])
.await;
assert_eq!(events.len(), 2);
match &events[0] {
SummaryStreamEvent::Delta(d) => assert_eq!(d.content, "split"),
other => panic!("expected delta, got {other:?}"),
}
}
#[tokio::test]
async fn summary_stream_reports_malformed_frames_without_ending() {
let events = collect_summary(vec![
"event: delta\ndata: not json\n\n",
"event: delta\ndata: {\"content\":\"recovered\"}\n\n",
": keep-alive\n\n",
"event: unknown\ndata: {}\n\n",
"event: done\ndata: {\"results_count\":1}\n\n",
])
.await;
assert_eq!(events.len(), 3);
match &events[0] {
SummaryStreamEvent::Error(e) => assert!(e.message.contains("invalid JSON")),
other => panic!("expected error, got {other:?}"),
}
match &events[1] {
SummaryStreamEvent::Delta(d) => assert_eq!(d.content, "recovered"),
other => panic!("expected delta, got {other:?}"),
}
assert!(matches!(events[2], SummaryStreamEvent::Done(_)));
}
#[tokio::test]
async fn summary_stream_reports_wrong_shaped_payloads() {
let events = collect_summary(vec!["event: metadata\ndata: {\"cached\":false}\n\n"]).await;
assert_eq!(events.len(), 1);
match &events[0] {
SummaryStreamEvent::Error(e) => assert!(e.message.contains("malformed metadata")),
other => panic!("expected error, got {other:?}"),
}
}
#[tokio::test]
async fn summary_stream_surfaces_transport_failures_as_events() {
let stream = futures_util::stream::iter(vec![
Ok::<&[u8], std::io::Error>(b"event: delta\ndata: {\"content\":\"hi\"}\n\n"),
Err(std::io::Error::other("connection reset")),
]);
let events: Vec<_> = summary_stream(Box::pin(stream)).collect().await;
assert_eq!(events.len(), 2);
match &events[1] {
SummaryStreamEvent::Error(e) => assert!(e.message.contains("connection reset")),
other => panic!("expected error, got {other:?}"),
}
}
#[tokio::test]
async fn search_summary_sends_correct_request() {
let server = MockServer::start().await;
let body = "event: metadata\ndata: {\"results_count\":1,\"cached\":true}\n\n\
event: delta\ndata: {\"content\":\"A laptop.\"}\n\n\
event: done\ndata: {\"results_count\":1}\n\n";
Mock::given(method("POST"))
.and(path("/index/products/search/summary"))
.and(header("Authorization", "test-read-key"))
.and(header("Accept", "text/event-stream"))
.respond_with(
ResponseTemplate::new(200)
.insert_header("Content-Type", "text/event-stream")
.set_body_string(body),
)
.expect(1)
.mount(&server)
.await;
let client = test_client(&server.uri());
let request = QueryBuilder::fuzzy().term("laptop").build_request();
let events: Vec<_> = client
.search_summary("products", &request)
.await
.unwrap()
.collect()
.await;
assert_eq!(events.len(), 3);
match &events[1] {
SummaryStreamEvent::Delta(d) => assert_eq!(d.content, "A laptop."),
other => panic!("expected delta, got {other:?}"),
}
}
#[tokio::test]
async fn search_summary_surfaces_setup_errors() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/index/products/search/summary"))
.respond_with(
ResponseTemplate::new(403)
.set_body_json(serde_json::json!({"message": "summary not permitted"})),
)
.mount(&server)
.await;
let client = test_client(&server.uri());
let request = QueryBuilder::fuzzy().term("laptop").build_request();
let err = client
.search_summary("products", &request)
.await
.err()
.expect("expected an error");
assert!(matches!(err, crate::error::Error::Authentication { .. }));
}
}