use async_trait::async_trait;
use contextgraph_types::{
Capabilities, ContextQuery, ContextQueryResult, FrameKind, ProviderInfo, Verdict,
VerifyRequest, VerifyResponse,
};
use crate::error::HostError;
#[async_trait]
pub trait ContextProvider: Send + Sync {
fn id(&self) -> &str;
fn info(&self) -> &ProviderInfo;
fn capabilities(&self) -> &Capabilities;
async fn query(&self, query: &ContextQuery) -> Result<ContextQueryResult, HostError>;
async fn verify(&self, request: &VerifyRequest) -> Result<VerifyResponse, HostError> {
Ok(VerifyResponse::uniform(request, Verdict::Unknown))
}
async fn shutdown(&self) -> Result<(), HostError> {
Ok(())
}
}
pub fn frame_kind_name(kind: FrameKind) -> &'static str {
match kind {
FrameKind::Snippet => "snippet",
FrameKind::Symbol => "symbol",
FrameKind::Fact => "fact",
FrameKind::Doc => "doc",
FrameKind::Memory => "memory",
FrameKind::Episode => "episode",
FrameKind::Graph => "graph",
}
}
pub fn capability_matches(caps: &Capabilities, query: &ContextQuery) -> bool {
if query.kinds.is_empty() {
return true;
}
query.kinds.iter().any(|requested| {
let name = frame_kind_name(*requested);
caps.query.kinds.iter().any(|served| served == name)
})
}
#[cfg(test)]
mod tests {
use super::*;
use contextgraph_types::capability::QueryCapability;
fn caps_for(kinds: &[&str]) -> Capabilities {
Capabilities {
query: QueryCapability {
kinds: kinds.iter().map(|k| k.to_string()).collect(),
},
..Capabilities::default()
}
}
fn query_for(kinds: Vec<FrameKind>) -> ContextQuery {
ContextQuery {
goal: "g".into(),
query_text: None,
embedding: None,
kinds,
anchors: vec![],
max_frames: 5,
max_tokens: 1000,
as_of: None,
representation_preferences: vec![],
}
}
#[test]
fn frame_kind_names_match_serde_snake_case() {
for (kind, name) in [
(FrameKind::Snippet, "snippet"),
(FrameKind::Symbol, "symbol"),
(FrameKind::Fact, "fact"),
(FrameKind::Doc, "doc"),
(FrameKind::Memory, "memory"),
(FrameKind::Episode, "episode"),
(FrameKind::Graph, "graph"),
] {
assert_eq!(frame_kind_name(kind), name);
let serde_name = serde_json::to_value(kind).unwrap();
assert_eq!(serde_name, serde_json::Value::String(name.to_string()));
}
}
#[test]
fn an_empty_kind_filter_matches_every_provider() {
let caps = caps_for(&["doc"]);
assert!(capability_matches(&caps, &query_for(vec![])));
}
#[test]
fn a_kind_filter_matches_only_overlapping_providers() {
let doc_provider = caps_for(&["doc", "snippet"]);
assert!(capability_matches(
&doc_provider,
&query_for(vec![FrameKind::Doc])
));
assert!(capability_matches(
&doc_provider,
&query_for(vec![FrameKind::Fact, FrameKind::Snippet])
));
assert!(!capability_matches(
&doc_provider,
&query_for(vec![FrameKind::Memory, FrameKind::Episode])
));
}
}