use anyhow::Result;
use crate::entities::profile::ToolId;
use crate::entities::rag::{RagDocument, RagHit};
use crate::shared::api::EmbedRole;
use crate::shared::config::{
DEFAULT_CHUNK_MAX_CHARS, DEFAULT_CHUNK_OVERLAP_CHARS, DEFAULT_CHUNK_TARGET_CHARS, RagSettings,
};
use super::{Tool, ToolContext, ToolOutcome};
const MIN_STITCH_OVERLAP: usize = 24;
const DEFAULT_TOP_K: usize = 5;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ChunkParams {
pub target: usize,
pub overlap: usize,
pub max: usize,
}
impl Default for ChunkParams {
fn default() -> Self {
Self {
target: DEFAULT_CHUNK_TARGET_CHARS,
overlap: DEFAULT_CHUNK_OVERLAP_CHARS,
max: DEFAULT_CHUNK_MAX_CHARS,
}
}
}
impl ChunkParams {
pub fn from_settings(rag: &RagSettings) -> Self {
let target = if rag.chunk_target_chars == 0 {
DEFAULT_CHUNK_TARGET_CHARS
} else {
rag.chunk_target_chars
};
let max = rag.chunk_max_chars.max(target);
Self {
target,
overlap: rag.chunk_overlap_chars.min(target.saturating_sub(1)),
max,
}
}
}
pub struct RagAdd;
#[async_trait::async_trait]
impl Tool for RagAdd {
fn id(&self) -> ToolId {
"rag_add".into()
}
fn group(&self) -> crate::features::tools::meta::ToolGroup {
crate::features::tools::meta::ToolGroup::Memory
}
fn ui_label(&self) -> &'static str {
"add to knowledge base"
}
fn description(&self, loc: &crate::shared::i18n::Locale) -> String {
loc.t("tool.rag_add.desc").into()
}
fn parameters(&self, loc: &crate::shared::i18n::Locale) -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {
"text": {"type": "string"},
"source": {"type": "string", "description": loc.t("tool.rag_add.param.source")}
},
"required": ["text"]
})
}
async fn invoke(&self, ctx: &ToolContext, args: serde_json::Value) -> Result<ToolOutcome> {
let text = args
.get("text")
.and_then(|v| v.as_str())
.ok_or_else(|| anyhow::anyhow!(ctx.loc.t("tool.rag_add.err.text_string")))?;
let source = args
.get("source")
.and_then(|v| v.as_str())
.map(str::to_string)
.unwrap_or_else(|| ctx.loc.t("tool.rag_add.no_source").to_string());
let chunks = chunk_text(text, ctx.chunk_params);
if chunks.is_empty() {
anyhow::bail!(ctx.loc.t("tool.rag_add.err.no_content"));
}
let embeddings = ctx
.embedder
.embed(chunks.clone(), EmbedRole::Passage)
.await?;
if embeddings.len() != chunks.len() {
anyhow::bail!(ctx.loc.t("tool.rag_add.err.embed_count"));
}
for (chunk, embedding) in chunks.iter().zip(embeddings) {
let doc = RagDocument::new(ctx.profile_id, &source, chunk, embedding);
ctx.storage.db().rag_insert(&doc)?;
}
ctx.storage
.db()
.rag_source_append(ctx.profile_id, &source, text, chrono::Utc::now())?;
Ok(ToolOutcome::text(ctx.loc.tf(
"tool.rag_add.result.added",
&[("n", &chunks.len().to_string())],
)))
}
}
pub struct RagSearch;
#[async_trait::async_trait]
impl Tool for RagSearch {
fn id(&self) -> ToolId {
"rag_search".into()
}
fn group(&self) -> crate::features::tools::meta::ToolGroup {
crate::features::tools::meta::ToolGroup::Memory
}
fn ui_label(&self) -> &'static str {
"search knowledge base"
}
fn description(&self, loc: &crate::shared::i18n::Locale) -> String {
loc.t("tool.rag_search.desc").into()
}
fn parameters(&self, _loc: &crate::shared::i18n::Locale) -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {
"query": {"type": "string"},
"top_k": {"type": "integer", "minimum": 1}
},
"required": ["query"]
})
}
async fn invoke(&self, ctx: &ToolContext, args: serde_json::Value) -> Result<ToolOutcome> {
let query = args
.get("query")
.and_then(|v| v.as_str())
.filter(|s| !s.trim().is_empty())
.ok_or_else(|| anyhow::anyhow!(ctx.loc.t("tool.rag_search.err.query_empty")))?;
let k = args
.get("top_k")
.and_then(|v| v.as_u64())
.map(|n| n as usize)
.unwrap_or(DEFAULT_TOP_K);
if ctx
.storage
.db()
.rag_is_stale(ctx.profile_id)
.unwrap_or(false)
{
return Ok(ToolOutcome::text(ctx.loc.t("tool.rag_search.err.stale")));
}
let mut embeddings = ctx
.embedder
.embed(vec![query.to_string()], EmbedRole::Query)
.await?;
let query_vec = embeddings
.pop()
.ok_or_else(|| anyhow::anyhow!(ctx.loc.t("tool.rag_search.err.no_query_vec")))?;
let hits = ctx.storage.db().rag_search(ctx.profile_id, &query_vec, k)?;
if hits.is_empty() {
return Ok(ToolOutcome::text(ctx.loc.t("tool.rag_search.result.empty")));
}
let passages = dedup_passages(stitch_hits(hits));
let mut out = ctx.loc.tf(
"tool.rag_search.result.header",
&[("n", &passages.len().to_string())],
);
for (i, p) in passages.iter().enumerate() {
out.push_str(&format!("\n\n{}. [{}]\n{}", i + 1, p.source, p.text));
}
out.push('\n');
let mut seen: std::collections::HashSet<uuid::Uuid> = std::collections::HashSet::new();
let mut linked: Vec<String> = Vec::new();
for p in &passages {
let notes = ctx
.storage
.db()
.notes_citing_source(ctx.profile_id, &p.source)
.unwrap_or_default();
for n in notes {
if !seen.insert(n.id) {
continue;
}
let mark = if super::notes::is_self_note(&n) {
format!("{} ", ctx.loc.t("notes.mark.self"))
} else {
String::new()
};
linked.push(format!("- {mark}(id={}) {}", n.id, n.content));
}
}
if !linked.is_empty() {
out.push('\n');
out.push_str(ctx.loc.t("tool.rag_search.block.linked_notes"));
out.push_str(":\n");
out.push_str(&linked.join("\n"));
out.push('\n');
}
Ok(ToolOutcome::text(out.trim_end().to_string()))
}
}
fn clen(s: &str) -> usize {
s.chars().count()
}
pub(crate) fn chunk_text(text: &str, params: ChunkParams) -> Vec<String> {
let units = segment_units(text, params);
pack_units(&units, params.target, params.overlap)
}
pub(crate) fn chunk_markdown(text: &str, params: ChunkParams) -> Vec<String> {
let sections = split_sections(text);
let mut chunks = Vec::new();
for (heading, body) in §ions {
let units = segment_units(body, params);
let packed = if units.is_empty() {
vec![String::new()]
} else {
pack_units(&units, params.target, params.overlap)
};
for p in packed {
let chunk = match (heading.is_empty(), p.is_empty()) {
(true, _) => p,
(false, true) => heading.clone(),
(false, false) => format!("{heading}\n{p}"),
};
let chunk = chunk.trim();
if !chunk.is_empty() {
chunks.push(chunk.to_string());
}
}
}
if chunks.is_empty() {
return chunk_text(text, params);
}
chunks
}
fn segment_units(text: &str, params: ChunkParams) -> Vec<String> {
let mut units = Vec::new();
for paragraph in text.split("\n\n") {
let p = paragraph.trim();
if p.is_empty() {
continue;
}
if clen(p) <= params.target {
units.push(p.to_string());
continue;
}
for sentence in split_sentences(p) {
if clen(&sentence) <= params.max {
units.push(sentence);
} else {
units.extend(break_long(&sentence, params.max));
}
}
}
units
}
pub(crate) fn split_sentences(paragraph: &str) -> Vec<String> {
let chars: Vec<char> = paragraph.chars().collect();
let mut out = Vec::new();
let mut start = 0;
let mut i = 0;
while i < chars.len() {
let is_end = matches!(chars[i], '.' | '!' | '?' | '…' | '。' | '!' | '?');
let next_ws = chars.get(i + 1).map(|c| c.is_whitespace()).unwrap_or(true);
if is_end && next_ws {
let seg: String = chars[start..=i].iter().collect();
let seg = seg.trim();
if !seg.is_empty() {
out.push(seg.to_string());
}
let mut j = i + 1;
while j < chars.len() && chars[j].is_whitespace() {
j += 1;
}
start = j;
i = j;
continue;
}
i += 1;
}
if start < chars.len() {
let seg: String = chars[start..].iter().collect();
let seg = seg.trim();
if !seg.is_empty() {
out.push(seg.to_string());
}
}
out
}
fn break_long(s: &str, max: usize) -> Vec<String> {
let mut out = Vec::new();
let mut cur = String::new();
let mut cur_len = 0usize;
for word in s.split_whitespace() {
let wlen = clen(word);
if wlen > max {
flush_window(&mut cur, &mut cur_len, &mut out);
out.extend(tear_word(word, max));
continue;
}
let add = if cur.is_empty() { wlen } else { wlen + 1 };
if cur_len + add > max && !cur.is_empty() {
out.push(std::mem::take(&mut cur));
cur.push_str(word);
cur_len = wlen;
} else {
if !cur.is_empty() {
cur.push(' ');
cur_len += 1;
}
cur.push_str(word);
cur_len += wlen;
}
}
if !cur.is_empty() {
out.push(cur);
}
out
}
fn flush_window(cur: &mut String, cur_len: &mut usize, out: &mut Vec<String>) {
if !cur.is_empty() {
out.push(std::mem::take(cur));
*cur_len = 0;
}
}
fn tear_word(word: &str, max: usize) -> Vec<String> {
let chars: Vec<char> = word.chars().collect();
chars.chunks(max).map(|w| w.iter().collect()).collect()
}
fn pack_units(units: &[String], target: usize, overlap: usize) -> Vec<String> {
let mut chunks = Vec::new();
let mut cur: Vec<&str> = Vec::new();
let mut cur_len = 0usize;
for unit in units {
let ulen = clen(unit);
let add = if cur.is_empty() { ulen } else { ulen + 1 };
if !cur.is_empty() && cur_len + add > target {
chunks.push(cur.join("\n"));
let (tail, tlen) = overlap_tail(&cur, overlap);
cur = tail;
cur_len = tlen;
}
if !cur.is_empty() {
cur_len += 1;
}
cur.push(unit);
cur_len += ulen;
}
if !cur.is_empty() {
chunks.push(cur.join("\n"));
}
chunks
}
fn overlap_tail<'a>(cur: &[&'a str], overlap: usize) -> (Vec<&'a str>, usize) {
let mut tail: Vec<&str> = Vec::new();
let mut tlen = 0usize;
for &u in cur.iter().rev() {
let a = if tail.is_empty() {
clen(u)
} else {
clen(u) + 1
};
if tlen + a > overlap && !tail.is_empty() {
break;
}
tail.push(u);
tlen += a;
}
tail.reverse();
(tail, tlen)
}
fn split_sections(text: &str) -> Vec<(String, String)> {
let mut sections: Vec<(String, String)> = Vec::new();
let mut heading = String::new();
let mut body = String::new();
let mut fence: Option<&str> = None;
for line in text.lines() {
let trimmed = line.trim_start();
if let Some(f) = fence {
if trimmed.starts_with(f) {
fence = None;
}
body.push_str(line);
body.push('\n');
continue;
}
if let Some(f) = fence_open(trimmed) {
fence = Some(f);
body.push_str(line);
body.push('\n');
continue;
}
if is_atx_heading(trimmed) {
flush_section(&mut sections, &mut heading, &mut body);
heading = trimmed.trim_end().to_string();
} else {
body.push_str(line);
body.push('\n');
}
}
flush_section(&mut sections, &mut heading, &mut body);
sections
}
fn fence_open(trimmed: &str) -> Option<&'static str> {
if trimmed.starts_with("```") {
Some("```")
} else if trimmed.starts_with("~~~") {
Some("~~~")
} else {
None
}
}
fn flush_section(sections: &mut Vec<(String, String)>, heading: &mut String, body: &mut String) {
if !heading.is_empty() || !body.trim().is_empty() {
sections.push((std::mem::take(heading), std::mem::take(body)));
}
}
fn is_atx_heading(line: &str) -> bool {
let hashes = line.chars().take_while(|&c| c == '#').count();
(1..=6).contains(&hashes) && line.chars().nth(hashes) == Some(' ')
}
pub(crate) struct StitchedPassage {
pub source: String,
pub text: String,
pub distance: f32,
}
pub(crate) fn stitch_hits(hits: Vec<RagHit>) -> Vec<StitchedPassage> {
let mut by_source: Vec<(String, Vec<(String, f32)>)> = Vec::new();
for h in hits {
match by_source.iter_mut().find(|(s, _)| *s == h.source) {
Some(g) => g.1.push((h.chunk_text, h.distance)),
None => by_source.push((h.source, vec![(h.chunk_text, h.distance)])),
}
}
let mut passages = Vec::new();
for (source, mut items) in by_source {
while let Some((i, j, text, dist)) = find_mergeable(&items) {
let (hi, lo) = (i.max(j), i.min(j));
items.remove(hi);
items.remove(lo);
items.push((text, dist));
}
for (text, distance) in items {
passages.push(StitchedPassage {
source: source.clone(),
text,
distance,
});
}
}
passages.sort_by(|a, b| {
a.distance
.partial_cmp(&b.distance)
.unwrap_or(std::cmp::Ordering::Equal)
});
passages
}
pub(crate) fn dedup_passages(passages: Vec<StitchedPassage>) -> Vec<StitchedPassage> {
let mut kept: Vec<StitchedPassage> = Vec::with_capacity(passages.len());
let mut kept_norms: Vec<String> = Vec::with_capacity(passages.len());
for p in passages {
let norm = normalize_passage(&p.text);
if kept_norms.iter().any(|k| k.contains(&norm)) {
continue;
}
kept_norms.push(norm);
kept.push(p);
}
kept
}
fn normalize_passage(text: &str) -> String {
text.to_lowercase()
.split_whitespace()
.collect::<Vec<_>>()
.join(" ")
}
fn find_mergeable(items: &[(String, f32)]) -> Option<(usize, usize, String, f32)> {
for i in 0..items.len() {
for j in 0..items.len() {
if i == j {
continue;
}
if let Some(text) = merge_overlap(&items[i].0, &items[j].0) {
return Some((i, j, text, items[i].1.min(items[j].1)));
}
}
}
None
}
fn merge_overlap(a: &str, b: &str) -> Option<String> {
if let Some(k) = overlap_len(a, b) {
let tail: String = b.chars().skip(k).collect();
return Some(format!("{a}{tail}"));
}
if let (Some(ha), Some((hb, rest_b))) = (leading_heading(a), strip_leading_heading(b))
&& ha == hb
&& let Some(k) = overlap_len(a, &rest_b)
{
let tail: String = rest_b.chars().skip(k).collect();
return Some(format!("{a}{tail}"));
}
None
}
fn overlap_len(a: &str, b: &str) -> Option<usize> {
let ac: Vec<char> = a.chars().collect();
let bc: Vec<char> = b.chars().collect();
let max = ac.len().min(bc.len());
let mut k = max;
while k >= MIN_STITCH_OVERLAP {
if ac[ac.len() - k..] == bc[..k] {
return Some(k);
}
k -= 1;
}
None
}
fn leading_heading(s: &str) -> Option<String> {
let first = s.lines().next()?;
is_atx_heading(first.trim_start()).then(|| first.trim_end().to_string())
}
fn strip_leading_heading(s: &str) -> Option<(String, String)> {
let mut lines = s.splitn(2, '\n');
let first = lines.next()?;
if is_atx_heading(first.trim_start()) {
Some((
first.trim_end().to_string(),
lines.next().unwrap_or("").to_string(),
))
} else {
None
}
}
#[cfg(test)]
mod tests {
use super::super::testkit::ctx_with_storage;
use super::*;
use uuid::Uuid;
fn no_cyr(s: &str) -> bool {
!s.chars()
.any(|c| ('а'..='я').contains(&c) || ('А'..='Я').contains(&c))
}
#[test]
fn rag_tool_descriptions_are_localized() {
use crate::shared::i18n::{Lang, locale};
let (ru, en) = (locale(Lang::Ru), locale(Lang::En));
for (r, e) in [
(RagAdd.description(ru), RagAdd.description(en)),
(RagSearch.description(ru), RagSearch.description(en)),
] {
assert_ne!(r, e, "description not localized: {r}");
assert!(no_cyr(&e), "Cyrillic in en description: {e}");
}
}
#[tokio::test]
async fn rag_add_result_localized_for_all_langs() {
use crate::shared::i18n::{Lang, locale};
for &lang in Lang::ALL {
let (_d, _s, mut ctx) = ctx_with_storage(Uuid::new_v4());
ctx.loc = locale(lang);
let out = RagAdd
.invoke(&ctx, serde_json::json!({"text": "hello world alpha beta"}))
.await
.unwrap();
let prefix = locale(lang)
.t("tool.rag_add.result.added")
.split("{n}")
.next()
.unwrap();
assert!(out.result.starts_with(prefix), "{lang:?}: {}", out.result);
}
}
#[test]
fn chunking_groups_small_paragraphs() {
let chunks = chunk_text(
"первый абзац\n\nвторой абзац\n\nтретий абзац",
ChunkParams::default(),
);
assert_eq!(chunks.len(), 1, "{chunks:?}");
assert!(chunks[0].contains("первый абзац"));
assert!(chunks[0].contains("третий абзац"));
}
#[test]
fn chunking_splits_long_paragraph_with_overlap() {
let sentence = "Это предложение средней длины для проверки чанкинга. ";
let text = sentence.repeat(60); let params = ChunkParams::default();
let chunks = chunk_text(&text, params);
assert!(
chunks.len() >= 2,
"expected several chunks: {}",
chunks.len()
);
for c in &chunks {
assert!(clen(c) <= params.max + params.overlap, "{}", clen(c));
}
assert!(
overlap_len(&chunks[0], &chunks[1]).is_some(),
"expected an overlap between adjacent chunks"
);
}
#[test]
fn chunking_never_breaks_mid_word() {
let text = format!("короткое начало {} конец", "ё".repeat(2500));
let chunks = chunk_text(&text, ChunkParams::default());
assert!(chunks.iter().any(|c| c.contains("короткое начало")));
assert!(chunks.iter().any(|c| c.contains("конец")));
}
#[test]
fn chunk_params_from_settings_respects_config() {
let text = "Это предложение средней длины для проверки чанкинга. ".repeat(20);
let big = chunk_text(&text, ChunkParams::default());
let small = chunk_text(
&text,
ChunkParams::from_settings(&RagSettings {
chunk_target_chars: 200,
chunk_overlap_chars: 40,
chunk_max_chars: 400,
}),
);
assert!(
small.len() > big.len(),
"a smaller target → more chunks: small={} big={}",
small.len(),
big.len()
);
}
#[test]
fn chunk_params_from_settings_sanitizes_invalid() {
let p = ChunkParams::from_settings(&RagSettings {
chunk_target_chars: 0,
chunk_overlap_chars: 9999,
chunk_max_chars: 10,
});
assert_eq!(p.target, DEFAULT_CHUNK_TARGET_CHARS);
assert!(p.overlap < p.target);
assert!(p.max >= p.target, "ceiling not smaller than target");
}
#[test]
fn markdown_chunks_carry_their_heading() {
let md = "# Заголовок\n\nтекст раздела один\n\n## Подраздел\n\nтекст подраздела";
let chunks = chunk_markdown(md, ChunkParams::default());
assert!(chunks.iter().any(|c| c.starts_with("# Заголовок")));
assert!(chunks.iter().any(|c| c.starts_with("## Подраздел")));
assert!(chunks.iter().all(|c| c.starts_with('#')));
}
#[test]
fn markdown_ignores_hash_inside_code_fence() {
let md = "# Реальный заголовок\n\n```python\n# это комментарий, не заголовок\nx = 1\n```";
let sections = split_sections(md);
assert_eq!(sections.len(), 1, "{sections:?}");
assert_eq!(sections[0].0, "# Реальный заголовок");
assert!(sections[0].1.contains("# это комментарий"));
}
#[test]
fn stitch_merges_overlapping_neighbors() {
let mk = |text: &str, d: f32| RagHit {
id: Uuid::new_v4(),
source: "doc.md".into(),
chunk_text: text.into(),
distance: d,
};
let a = "альфа бета гамма дельта эпсилон дзета";
let b = "гамма дельта эпсилон дзета эта тета йота";
let merged = stitch_hits(vec![mk(a, 0.2), mk(b, 0.3)]);
assert_eq!(merged.len(), 1, "should stitch into one fragment");
assert_eq!(
merged[0].text,
"альфа бета гамма дельта эпсилон дзета эта тета йота"
);
assert_eq!(merged[0].distance, 0.2, "the best distance is taken");
}
#[test]
fn stitch_keeps_unrelated_hits_separate() {
let mk = |src: &str, text: &str| RagHit {
id: Uuid::new_v4(),
source: src.into(),
chunk_text: text.into(),
distance: 0.5,
};
let out = stitch_hits(vec![
mk("a.txt", "совершенно разный текст один"),
mk("b.txt", "никак не связанный текст два"),
]);
assert_eq!(out.len(), 2, "different sources don't stitch");
}
fn passage(source: &str, text: &str, distance: f32) -> StitchedPassage {
StitchedPassage {
source: source.into(),
text: text.into(),
distance,
}
}
#[test]
fn dedup_drops_identical_from_different_sources_keeps_higher_ranked() {
let out = dedup_passages(vec![
passage("a.txt", "столица франции — париж", 0.1),
passage("b.txt", "столица франции — париж", 0.3),
]);
assert_eq!(
out.len(),
1,
"an identical duplicate from another source is removed"
);
assert_eq!(
out[0].source, "a.txt",
"the more relevant one (smaller distance) remains"
);
}
#[test]
fn dedup_drops_passage_contained_in_higher_ranked() {
let out = dedup_passages(vec![
passage("a.txt", "полный текст с деталями про париж и францию", 0.1),
passage("b.txt", "париж и францию", 0.4),
]);
assert_eq!(
out.len(),
1,
"a subset of the more relevant passage is removed"
);
assert_eq!(out[0].source, "a.txt");
}
#[test]
fn dedup_keeps_distinct_passages_in_order() {
let out = dedup_passages(vec![
passage("a.txt", "первый совершенно уникальный текст", 0.1),
passage("b.txt", "второй никак не связанный текст", 0.2),
]);
assert_eq!(out.len(), 2, "distinct passages are left alone");
assert_eq!(out[0].source, "a.txt");
assert_eq!(out[1].source, "b.txt");
}
#[test]
fn dedup_is_whitespace_and_case_insensitive() {
let out = dedup_passages(vec![
passage("a.txt", "Столица Франции —\nПариж", 0.1),
passage("b.txt", "столица франции — париж", 0.3),
]);
assert_eq!(out.len(), 1, "equal up to case/whitespace → dedup");
assert_eq!(out[0].source, "a.txt");
}
#[test]
fn dedup_keeps_lower_ranked_superset() {
let out = dedup_passages(vec![
passage("a.txt", "париж", 0.1),
passage("b.txt", "париж — столица франции и крупный город", 0.3),
]);
assert_eq!(out.len(), 2, "a lower-ranked superset isn't removed");
assert_eq!(out[0].source, "a.txt");
assert_eq!(out[1].source, "b.txt");
}
#[test]
fn dedup_empty_input_empty_output() {
assert!(dedup_passages(Vec::new()).is_empty());
}
#[tokio::test]
async fn add_then_search_returns_relevant_chunk() {
let (_d, _s, ctx) = ctx_with_storage(Uuid::new_v4());
RagAdd
.invoke(
&ctx,
serde_json::json!({
"text": "кошки любят рыбу\n\nсобаки любят кости",
"source": "факты"
}),
)
.await
.unwrap();
let out = RagSearch
.invoke(&ctx, serde_json::json!({"query": "кошки рыба", "top_k": 1}))
.await
.unwrap();
assert!(
out.result.contains("кошки любят рыбу"),
"got: {}",
out.result
);
assert!(out.result.contains("факты"));
}
#[tokio::test]
async fn search_refuses_on_a_stale_knowledge_base() {
let profile = Uuid::new_v4();
let (_d, storage, ctx) = ctx_with_storage(profile);
RagAdd
.invoke(
&ctx,
serde_json::json!({"text": "кошки любят рыбу", "source": "факты"}),
)
.await
.unwrap();
let ok = RagSearch
.invoke(&ctx, serde_json::json!({"query": "кошки"}))
.await
.unwrap();
assert!(ok.result.contains("кошки любят рыбу"), "got: {}", ok.result);
storage.db().set_rag_stale_profiles(&[profile]).unwrap();
let stale = RagSearch
.invoke(&ctx, serde_json::json!({"query": "кошки"}))
.await
.unwrap();
assert!(
!stale.result.contains("кошки любят рыбу"),
"a stale base must not return passages: {}",
stale.result
);
assert!(
stale.result.contains("/reindex"),
"the refusal must name the fix: {}",
stale.result
);
storage.db().clear_rag_stale_profile(profile).unwrap();
let healed = RagSearch
.invoke(&ctx, serde_json::json!({"query": "кошки"}))
.await
.unwrap();
assert!(healed.result.contains("кошки любят рыбу"));
}
#[tokio::test]
async fn search_numbers_passages_and_separates_them() {
let (_d, _s, ctx) = ctx_with_storage(Uuid::new_v4());
for (text, source) in [
("кошки любят рыбу\nи спят на солнце", "про-кошек"),
("собаки любят кости\nи гулять во дворе", "про-собак"),
] {
RagAdd
.invoke(&ctx, serde_json::json!({"text": text, "source": source}))
.await
.unwrap();
}
let out = RagSearch
.invoke(&ctx, serde_json::json!({"query": "любят", "top_k": 5}))
.await
.unwrap()
.result;
assert!(out.contains("1. ["), "{out}");
assert!(out.contains("2. ["), "{out}");
assert!(
out.contains("]\nкошки любят рыбу") || out.contains("]\nсобаки любят кости"),
"{out}"
);
assert!(
out.contains("\n\n2. ["),
"passages must be separated by a blank line: {out}"
);
}
#[tokio::test]
async fn search_surfaces_notes_citing_matched_source() {
use crate::entities::note::Note;
let profile = Uuid::new_v4();
let (_d, storage, ctx) = ctx_with_storage(profile);
RagAdd
.invoke(
&ctx,
serde_json::json!({"text": "кошки любят рыбу", "source": "факты"}),
)
.await
.unwrap();
let note = Note::new(profile, "мой вывод о кошках", vec![]);
let nid = note.id;
storage.db().note_insert(¬e).unwrap();
storage
.db()
.note_cite_source_insert(profile, nid, "факты")
.unwrap();
let out = RagSearch
.invoke(&ctx, serde_json::json!({"query": "кошки рыба", "top_k": 1}))
.await
.unwrap();
assert!(out.result.contains("Заметки со ссылкой на эти источники"));
assert!(out.result.contains("мой вывод о кошках"));
}
#[tokio::test]
async fn search_isolated_by_profile() {
let dir = tempfile::tempdir().unwrap();
let storage = std::sync::Arc::new(
crate::shared::storage::Storage::open_in_memory(
crate::shared::paths::Paths::with_root(dir.path()),
)
.unwrap(),
);
let engine: std::sync::Arc<dyn crate::shared::api::EngineBackend> =
std::sync::Arc::new(crate::shared::api::mock::MockBackend::scripted(vec![]));
let embedder: std::sync::Arc<dyn crate::shared::api::Embedder> =
std::sync::Arc::new(crate::shared::api::mock::MockEmbedder::new(16));
let a = Uuid::new_v4();
let b = Uuid::new_v4();
let deps = crate::features::tools::ToolDeps {
storage: storage.clone(),
engine: engine.clone(),
embedder: embedder.clone(),
};
let mk = |pid| super::super::testkit::ctx_with_deps(pid, deps.clone());
RagAdd
.invoke(&mk(a), serde_json::json!({"text": "секрет профиля A"}))
.await
.unwrap();
let out = RagSearch
.invoke(&mk(b), serde_json::json!({"query": "секрет профиля A"}))
.await
.unwrap();
assert!(out.result.contains("ничего не найдено"));
}
}