use std::{collections::BTreeMap, future::Future, pin::Pin};
use runifold_core::{
CapabilityDescriptor, CapabilityId, CapabilityKind, EffectClass, RiskLevel, Usage,
};
use serde_json::{Value, json};
use crate::{Document, RetrievalContext, RetrievalError};
#[cfg(not(target_arch = "wasm32"))]
pub type RetrievalFuture<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
#[cfg(target_arch = "wasm32")]
pub type RetrievalFuture<'a, T> = Pin<Box<dyn Future<Output = T> + 'a>>;
#[derive(Clone, Debug, PartialEq)]
pub struct RetrieverDescriptor {
pub id: CapabilityId,
pub name: String,
pub version: String,
pub effect: EffectClass,
pub risk: RiskLevel,
pub metadata: BTreeMap<String, Value>,
}
impl RetrieverDescriptor {
pub fn read_only(name: impl Into<String>) -> Self {
Self {
id: CapabilityId::new(),
name: name.into(),
version: "1".into(),
effect: EffectClass::ReadOnly,
risk: RiskLevel::Low,
metadata: BTreeMap::new(),
}
}
pub fn capability(&self) -> CapabilityDescriptor {
CapabilityDescriptor {
id: self.id,
name: self.name.clone(),
version: self.version.clone(),
kind: CapabilityKind::Extension("runifold.retrieval".into()),
input_schema: json!({
"type": "object",
"required": ["query", "limit"],
"properties": {
"query": {"type": "string"},
"limit": {"type": "integer", "minimum": 1}
}
}),
output_schema: json!({
"type": "array",
"items": {
"type": "object",
"required": ["id", "score"],
"properties": {
"id": {"type": "string"},
"score": {"type": "number"}
}
}
}),
effect: self.effect,
risk: self.risk,
metadata: self.metadata.clone(),
}
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct RetrievalQuery {
pub text: String,
pub limit: usize,
}
impl RetrievalQuery {
pub fn new(text: impl Into<String>, limit: usize) -> Result<Self, RetrievalError> {
let text = text.into();
if text.trim().is_empty() {
return Err(RetrievalError::EmptyQuery);
}
if limit == 0 {
return Err(RetrievalError::ZeroLimit);
}
Ok(Self { text, limit })
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct RetrievedDocument {
pub document: Document,
pub score: f64,
}
#[derive(Clone, Debug, Default, PartialEq)]
pub struct RetrievalResponse {
pub documents: Vec<RetrievedDocument>,
pub usage: Usage,
}
pub trait Retriever: Send + Sync {
fn descriptor(&self) -> &RetrieverDescriptor;
fn retrieve(
&self,
query: RetrievalQuery,
context: RetrievalContext,
) -> RetrievalFuture<'_, Result<RetrievalResponse, RetrievalError>>;
}