use std::time::{SystemTime, UNIX_EPOCH};
use kindling_store::SqliteKindlingStore;
use kindling_types::{
CandidateResult, Id, PinResult, PinTargetType, RetrieveOptions, RetrieveProvenance,
RetrieveResult, RetrievedEntity, Timestamp,
};
use crate::error::ProviderResult;
use crate::provider::RetrievalProvider;
const DEFAULT_MAX_CANDIDATES: u32 = 10;
pub fn retrieve(
store: &SqliteKindlingStore,
provider: &impl RetrievalProvider,
options: &RetrieveOptions,
) -> ProviderResult<RetrieveResult> {
retrieve_at(store, provider, options, now_ms())
}
pub fn retrieve_at(
store: &SqliteKindlingStore,
provider: &impl RetrievalProvider,
options: &RetrieveOptions,
now: Timestamp,
) -> ProviderResult<RetrieveResult> {
let max_candidates = options.max_candidates.unwrap_or(DEFAULT_MAX_CANDIDATES);
let include_redacted = options.include_redacted.unwrap_or(false);
let pins = store.list_active_pins(Some(&options.scope_ids), Some(now))?;
let mut pin_results: Vec<PinResult> = Vec::new();
let mut pinned_ids: Vec<Id> = Vec::new();
for pin in pins {
let target = match pin.target_type {
PinTargetType::Observation => store
.get_observation_by_id(&pin.target_id)?
.map(RetrievedEntity::Observation),
PinTargetType::Summary => store
.get_summary_by_id(&pin.target_id)?
.map(RetrievedEntity::Summary),
};
if let Some(target) = target {
if let RetrievedEntity::Observation(obs) = &target {
if obs.redacted && !include_redacted {
continue;
}
}
let target_id = entity_id(&target).to_string();
if !pinned_ids.contains(&target_id) {
pinned_ids.push(target_id);
}
pin_results.push(PinResult { pin, target });
}
}
let mut current_summary = None;
if let Some(session_id) = options
.scope_ids
.session_id
.as_deref()
.filter(|s| !s.is_empty())
{
if let Some(capsule) = store.get_open_capsule_for_session(session_id)? {
if let Some(summary) = store.get_latest_summary_for_capsule(&capsule.id)? {
if !pinned_ids.contains(&summary.id) {
pinned_ids.push(summary.id.clone());
}
current_summary = Some(summary);
}
}
}
let provider_results = provider.search(
&kindling_types::ProviderSearchOptions {
query: options.query.clone(),
scope_ids: options.scope_ids.clone(),
max_results: Some(max_candidates),
exclude_ids: Some(pinned_ids),
include_redacted: Some(include_redacted),
},
now,
)?;
let total_candidates = provider_results.len() as u32;
let candidates: Vec<CandidateResult> = provider_results
.into_iter()
.map(|result| CandidateResult {
entity: result.entity,
score: result.score,
match_context: result.match_context,
})
.collect();
let provenance = RetrieveProvenance {
query: options.query.clone(),
scope_ids: options.scope_ids.clone(),
total_candidates,
returned_candidates: candidates.len() as u32,
truncated_due_to_token_budget: false,
provider_used: provider.name().to_string(),
};
Ok(RetrieveResult {
pins: pin_results,
current_summary,
candidates,
provenance,
})
}
fn entity_id(entity: &RetrievedEntity) -> &str {
match entity {
RetrievedEntity::Observation(obs) => &obs.id,
RetrievedEntity::Summary(sum) => &sum.id,
}
}
fn now_ms() -> Timestamp {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.expect("system clock before Unix epoch")
.as_millis() as Timestamp
}