use anyhow::Result;
use rusqlite::Connection;
use super::types::{MultiHopMeta, SearchRequest, SearchResultSet};
const RAW_FALLBACK_THRESHOLD: usize = 3;
const RAW_FALLBACK_LIMIT: i64 = 10;
pub fn search_memories(conn: &Connection, req: &SearchRequest) -> Result<SearchResultSet> {
let limit = req.limit.max(1);
let query = req.query.as_deref();
if req.multi_hop {
return multi_hop_search(conn, query, req.project.as_deref(), limit, req);
}
let (mut memories, mut explain) = if req.explain {
crate::retrieval::search::search_with_branch_explain(
conn,
query,
req.project.as_deref(),
req.memory_type.as_deref(),
limit + 1,
req.offset.max(0),
req.include_stale,
req.branch.as_deref(),
)?
} else {
(
crate::retrieval::search::search_with_branch(
conn,
query,
req.project.as_deref(),
req.memory_type.as_deref(),
limit + 1,
req.offset.max(0),
req.include_stale,
req.branch.as_deref(),
)?,
None,
)
};
let has_more = memories.len() as i64 > limit;
memories.truncate(limit as usize);
let raw_hits = maybe_fallback_raw(conn, req, memories.len());
if let Some(explain) = explain.as_mut() {
let result_ids: Vec<i64> = memories.iter().map(|memory| memory.id).collect();
explain.retain_result_ids(&result_ids, has_more, limit);
explain.set_raw_fallback_count(raw_hits.len());
}
Ok(SearchResultSet {
memories,
multi_hop: None,
has_more,
explain,
raw_hits,
})
}
fn multi_hop_search(
conn: &Connection,
query: Option<&str>,
project: Option<&str>,
limit: i64,
req: &SearchRequest,
) -> Result<SearchResultSet> {
if let Some(query_text) = query.filter(|query_text| !query_text.is_empty()) {
let mut result = crate::retrieval::search_multihop::search_multi_hop(
conn,
query_text,
project,
limit + 1,
req.offset.max(0),
req.memory_type.as_deref(),
req.branch.as_deref(),
req.include_stale,
)?;
let has_more = result.memories.len() as i64 > limit;
result.memories.truncate(limit as usize);
let raw_hits = maybe_fallback_raw(conn, req, result.memories.len());
Ok(SearchResultSet {
memories: result.memories,
multi_hop: Some(MultiHopMeta {
hops: result.hops,
entities_discovered: result.entities_discovered,
}),
has_more,
explain: None,
raw_hits,
})
} else {
Ok(SearchResultSet {
memories: vec![],
multi_hop: Some(MultiHopMeta {
hops: 1,
entities_discovered: vec![],
}),
has_more: false,
explain: None,
raw_hits: vec![],
})
}
}
fn maybe_fallback_raw(
conn: &Connection,
req: &SearchRequest,
curated_len: usize,
) -> Vec<crate::memory::raw_archive::RawMessage> {
if curated_len >= RAW_FALLBACK_THRESHOLD {
return vec![];
}
let Some(query) = req
.query
.as_deref()
.map(str::trim)
.filter(|q| !q.is_empty())
else {
return vec![];
};
let raw_req = crate::memory::raw_archive::RawSearchRequest {
query: query.to_string(),
project: req.project.clone(),
branch: req.branch.clone(),
role: None,
limit: RAW_FALLBACK_LIMIT,
offset: 0,
};
match crate::memory::raw_archive::search_raw_messages(conn, &raw_req) {
Ok(hits) => hits,
Err(error) => {
crate::log::warn("search", &format!("raw archive fallback failed: {}", error));
vec![]
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::memory::raw_archive::{insert_raw_message, ROLE_USER, SOURCE_HOOK};
#[test]
fn raw_fallback_respects_branch_filter() {
let conn = Connection::open_in_memory().unwrap();
crate::migrate::run_migrations(&conn).unwrap();
insert_raw_message(
&conn,
"s-main",
"/repo",
ROLE_USER,
"fallback needle from main",
SOURCE_HOOK,
Some("main"),
None,
)
.unwrap();
insert_raw_message(
&conn,
"s-feature",
"/repo",
ROLE_USER,
"fallback needle from feature",
SOURCE_HOOK,
Some("feature"),
None,
)
.unwrap();
insert_raw_message(
&conn,
"s-branchless",
"/repo",
ROLE_USER,
"fallback needle from branchless history",
SOURCE_HOOK,
None,
None,
)
.unwrap();
let result = search_memories(
&conn,
&SearchRequest {
query: Some("needle".to_string()),
project: Some("/repo".to_string()),
limit: 10,
branch: Some("main".to_string()),
..SearchRequest::default()
},
)
.unwrap();
let branches: Vec<Option<String>> =
result.raw_hits.into_iter().map(|hit| hit.branch).collect();
assert!(branches.contains(&Some("main".to_string())));
assert!(branches.contains(&None));
assert!(
!branches.contains(&Some("feature".to_string())),
"{branches:?}"
);
}
}