use std::sync::Arc;
use std::time::Instant;
use async_trait::async_trait;
use toolkit_macros::domain_model;
use tracing::{info, instrument};
use uuid::Uuid;
use crate::domain::error::{ChatEngineError, Result};
use crate::domain::message::{Message, MessagePart, MessageRole, message_text};
use crate::domain::ports::MessageRepo;
use crate::domain::ports::SessionRepo;
use crate::domain::search::{
Cursor, MAX_QUERY_LENGTH, MessageRef, SearchError, SearchPage, SearchQuery, SearchResult,
SessionMeta, make_snippet, sanitize_for_tsquery,
};
use crate::domain::service::session_service::Identity;
#[domain_model]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SearchScope {
Session,
CrossSession,
}
impl SearchScope {
#[must_use]
pub fn as_str(self) -> &'static str {
match self {
Self::Session => "session",
Self::CrossSession => "cross_session",
}
}
}
#[domain_model]
#[derive(Debug, Clone)]
pub struct ParsedQuery {
pub raw: String,
pub tsquery: String,
}
pub fn parse_search_query(raw: &str) -> std::result::Result<ParsedQuery, SearchError> {
let trimmed = raw.trim();
if trimmed.is_empty() {
return Err(SearchError::QueryRequired);
}
if trimmed.chars().count() > MAX_QUERY_LENGTH {
return Err(SearchError::QueryTooLong);
}
let tsquery = sanitize_for_tsquery(trimmed);
if tsquery.is_empty() {
return Err(SearchError::QueryRequired);
}
Ok(ParsedQuery {
raw: trimmed.to_string(),
tsquery,
})
}
#[domain_model]
#[derive(Debug, Clone)]
pub struct BackendHit {
pub message_id: Uuid,
pub session_id: Uuid,
pub parent_message_id: Option<Uuid>,
pub role: MessageRole,
pub parts: Vec<MessagePart>,
pub created_at: time::OffsetDateTime,
pub rank: f32,
}
#[domain_model]
#[derive(Debug, Clone)]
pub struct SearchScopeFilter {
pub tenant_id: String,
pub user_id: String,
pub session_id: Option<Uuid>,
}
#[async_trait]
pub trait SearchBackend: Send + Sync {
async fn search(
&self,
scope: &SearchScopeFilter,
query: &ParsedQuery,
cursor: Option<&Cursor>,
skip: u32,
limit: u32,
) -> std::result::Result<(Vec<BackendHit>, u64), ChatEngineError>;
}
#[domain_model]
#[derive(Debug, Default)]
pub struct InMemorySearchBackend {
rows: Vec<(SearchScopeFilter, Message)>,
}
impl InMemorySearchBackend {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn push(&mut self, scope: SearchScopeFilter, message: Message) {
self.rows.push((scope, message));
}
}
#[async_trait]
impl SearchBackend for InMemorySearchBackend {
async fn search(
&self,
scope: &SearchScopeFilter,
query: &ParsedQuery,
cursor: Option<&Cursor>,
skip: u32,
limit: u32,
) -> std::result::Result<(Vec<BackendHit>, u64), ChatEngineError> {
let needle = query.raw.to_lowercase();
let mut matches: Vec<BackendHit> = self
.rows
.iter()
.filter(|(s, _)| {
s.tenant_id == scope.tenant_id
&& s.user_id == scope.user_id
&& match scope.session_id {
Some(sid) => s.session_id == Some(sid),
None => true,
}
})
.filter(|(_, m)| !m.is_hidden_from_user)
.filter(|(_, m)| {
let text = message_text(&m.parts);
text.to_lowercase().contains(&needle)
})
.map(|(_, m)| BackendHit {
message_id: m.message_id,
session_id: m.session_id,
parent_message_id: m.parent_message_id,
role: m.role.clone(),
parts: m.parts.clone(),
created_at: m.created_at,
rank: 0.0,
})
.collect();
matches.sort_by(|a, b| {
b.created_at
.cmp(&a.created_at)
.then_with(|| b.message_id.cmp(&a.message_id))
});
let total = matches.len() as u64;
let matches = apply_cursor_desc(matches, cursor);
let skip = skip as usize;
let limit = limit as usize;
if skip >= matches.len() {
return Ok((Vec::new(), total));
}
let end = (skip + limit).min(matches.len());
Ok((matches[skip..end].to_vec(), total))
}
}
fn apply_cursor_desc(matches: Vec<BackendHit>, cursor: Option<&Cursor>) -> Vec<BackendHit> {
let Some(c) = cursor else {
return matches;
};
if let Some(c_ts) = c.created_at {
return matches
.into_iter()
.filter(|h| {
h.created_at < c_ts || (h.created_at == c_ts && h.message_id < c.message_id)
})
.collect();
}
match matches.iter().position(|h| h.message_id == c.message_id) {
Some(idx) => matches.into_iter().skip(idx + 1).collect(),
None => matches,
}
}
#[domain_model]
#[derive(Clone)]
pub struct SearchService {
sessions: Arc<dyn SessionRepo>,
messages: Arc<dyn MessageRepo>,
backend: Arc<dyn SearchBackend>,
}
impl SearchService {
#[must_use]
pub fn new(
sessions: Arc<dyn SessionRepo>,
messages: Arc<dyn MessageRepo>,
backend: Arc<dyn SearchBackend>,
) -> Self {
Self {
sessions,
messages,
backend,
}
}
#[instrument(skip(self, identity, query), fields(session_id = %session_id))]
pub async fn search_in_session(
&self,
identity: &Identity,
session_id: Uuid,
query: &SearchQuery,
) -> Result<SearchPage> {
let started = Instant::now();
let parsed =
parse_search_query(query.q.as_deref().unwrap_or("")).map_err(ChatEngineError::from)?;
let owned = self
.sessions
.find_by_id(&identity.tenant_id, &identity.user_id, session_id)
.await?
.ok_or_else(|| ChatEngineError::not_found("session", session_id))?;
let scope = SearchScopeFilter {
tenant_id: identity.tenant_id.clone(),
user_id: identity.user_id.clone(),
session_id: Some(owned.session_id),
};
let page = self
.run(&scope, &parsed, query, SearchScope::Session)
.await?;
let duration_ms = started.elapsed().as_millis() as u64;
info!(
target: "chat_engine::search",
scope = SearchScope::Session.as_str(),
session_id = %session_id,
query_length = parsed.raw.chars().count(),
result_count = page.items.len(),
duration_ms,
"search.completed"
);
Ok(page)
}
#[instrument(skip(self, identity, query))]
pub async fn search_across_sessions(
&self,
identity: &Identity,
query: &SearchQuery,
) -> Result<SearchPage> {
let started = Instant::now();
let parsed =
parse_search_query(query.q.as_deref().unwrap_or("")).map_err(ChatEngineError::from)?;
let scope = SearchScopeFilter {
tenant_id: identity.tenant_id.clone(),
user_id: identity.user_id.clone(),
session_id: None,
};
let page = self
.run(&scope, &parsed, query, SearchScope::CrossSession)
.await?;
let duration_ms = started.elapsed().as_millis() as u64;
info!(
target: "chat_engine::search",
scope = SearchScope::CrossSession.as_str(),
query_length = parsed.raw.chars().count(),
result_count = page.items.len(),
duration_ms,
"search.completed"
);
Ok(page)
}
async fn run(
&self,
scope: &SearchScopeFilter,
parsed: &ParsedQuery,
query: &SearchQuery,
kind: SearchScope,
) -> Result<SearchPage> {
let limit = query.effective_top();
let skip = if query.cursor.is_some() {
0
} else {
query.effective_skip()
};
let context_radius = query.effective_context_radius();
let cursor = match query.cursor.as_deref() {
Some(raw) => Some(Cursor::decode(raw).map_err(ChatEngineError::from)?),
None => None,
};
let (hits, total) = self
.backend
.search(scope, parsed, cursor.as_ref(), skip, limit + 1)
.await?;
let mut hits = hits;
let has_more = hits.len() as u32 > limit;
if has_more {
hits.truncate(limit as usize);
}
let next_cursor = if has_more {
hits.last()
.map(|h| Cursor::new(h.rank, h.message_id, h.created_at).encode())
} else {
None
};
let mut items = Vec::with_capacity(hits.len());
for hit in hits {
let context_messages = self
.load_context_window(hit.session_id, hit.created_at, context_radius)
.await?;
let parent_chain = self
.load_parent_chain(hit.session_id, hit.parent_message_id)
.await?;
let snippet = make_snippet(&message_text(&hit.parts), &parsed.raw);
let session_metadata = match kind {
SearchScope::CrossSession => self.load_session_meta(scope, hit.session_id).await?,
SearchScope::Session => None,
};
items.push(SearchResult {
message_id: hit.message_id,
session_id: hit.session_id,
content_snippet: snippet,
rank: hit.rank,
context_messages,
parent_chain,
session_metadata,
});
}
Ok(SearchPage {
items,
total_count: total,
next_cursor,
per_page: limit,
})
}
async fn load_context_window(
&self,
session_id: Uuid,
anchor: time::OffsetDateTime,
radius: u32,
) -> Result<Vec<MessageRef>> {
if radius == 0 {
return Ok(Vec::new());
}
let all = self.messages.list_active_path(session_id).await?;
let mut before: Vec<&Message> = Vec::new();
let mut after: Vec<&Message> = Vec::new();
for m in &all {
if m.is_hidden_from_user {
continue;
}
if m.created_at < anchor {
before.push(m);
} else if m.created_at > anchor {
after.push(m);
}
}
let before_skip = before.len().saturating_sub(radius as usize);
let before_slice = &before[before_skip..];
let after_take = (radius as usize).min(after.len());
let after_slice = &after[..after_take];
let mut out = Vec::with_capacity(before_slice.len() + after_slice.len());
for m in before_slice.iter().chain(after_slice.iter()) {
out.push(MessageRef {
message_id: m.message_id,
role: m.role.clone(),
parts: m.parts.clone(),
created_at: m.created_at,
});
}
Ok(out)
}
async fn load_parent_chain(
&self,
session_id: Uuid,
parent_message_id: Option<Uuid>,
) -> Result<Vec<MessageRef>> {
let Some(mut cursor) = parent_message_id else {
return Ok(Vec::new());
};
let all = self.messages.list_active_path(session_id).await?;
let mut chain: Vec<MessageRef> = Vec::new();
let max_depth = 256;
for _ in 0..max_depth {
let Some(m) = all.iter().find(|m| m.message_id == cursor) else {
break;
};
chain.push(MessageRef {
message_id: m.message_id,
role: m.role.clone(),
parts: m.parts.clone(),
created_at: m.created_at,
});
match m.parent_message_id {
Some(p) => cursor = p,
None => break,
}
}
chain.reverse();
Ok(chain)
}
async fn load_session_meta(
&self,
scope: &SearchScopeFilter,
session_id: Uuid,
) -> Result<Option<SessionMeta>> {
let row = self
.sessions
.find_by_id(&scope.tenant_id, &scope.user_id, session_id)
.await?;
let Some(row) = row else { return Ok(None) };
let title = row
.metadata
.as_ref()
.and_then(|v| v.get("title"))
.and_then(|t| t.as_str())
.map(std::string::ToString::to_string);
let tags = row
.metadata
.as_ref()
.and_then(|v| v.get("tags"))
.and_then(|v| v.as_array())
.map(|arr| {
arr.iter()
.filter_map(|v| v.as_str().map(std::string::ToString::to_string))
.collect::<Vec<_>>()
})
.unwrap_or_default();
Ok(Some(SessionMeta {
session_id,
title,
tags,
}))
}
}
impl SearchScopeFilter {
#[must_use]
pub fn new(
tenant_id: impl Into<String>,
user_id: impl Into<String>,
session_id: Option<Uuid>,
) -> Self {
Self {
tenant_id: tenant_id.into(),
user_id: user_id.into(),
session_id,
}
}
}
#[cfg(test)]
#[path = "search_service_tests.rs"]
mod search_service_tests;