use crate::error::ContextError;
use crate::ids::{ProviderId, RunId, SessionId};
use async_trait::async_trait;
use std::sync::Arc;
pub struct ContextRequest {
pub session_id: SessionId,
pub run_id: RunId,
pub turn: u32,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ContextFailurePolicy {
Ignore,
Warn,
Block,
}
#[derive(Clone, Debug)]
pub struct ContextItem {
pub source: String,
pub version: String,
pub text: String,
pub budget_tokens: Option<u64>,
pub persistable: bool,
pub allowed_providers: Option<Vec<ProviderId>>,
pub failure: ContextFailurePolicy,
}
#[async_trait]
pub trait ContextProvider: Send + Sync {
fn id(&self) -> &str;
fn failure_policy(&self) -> ContextFailurePolicy {
ContextFailurePolicy::Block
}
async fn load(&self, request: &ContextRequest) -> Result<Vec<ContextItem>, ContextError>;
}
pub async fn load_context(
providers: &[Arc<dyn ContextProvider>],
request: &ContextRequest,
) -> Result<Vec<ContextItem>, ContextError> {
let mut items = Vec::new();
for provider in providers {
match provider.load(request).await {
Ok(loaded) => items.extend(loaded),
Err(error) => match provider.failure_policy() {
ContextFailurePolicy::Ignore => {}
ContextFailurePolicy::Warn => {
tracing::warn!(
provider = provider.id(),
?error,
"context provider failed; skipping (Warn policy)"
);
}
ContextFailurePolicy::Block => {
return Err(ContextError::Load(format!(
"provider {} failed: {error}",
provider.id()
)))
}
},
}
}
Ok(items)
}
pub fn truncate_to_budget(text: &str, budget_tokens: Option<u64>) -> String {
let Some(budget) = budget_tokens else {
return text.to_string();
};
let char_budget = budget.saturating_mul(4) as usize;
let count = text.chars().count();
if count <= char_budget {
return text.to_string();
}
let truncated: String = text.chars().take(char_budget).collect();
format!("{truncated}\n…(truncated)")
}
#[cfg(test)]
mod tests {
use super::*;
struct FixedProvider {
id: &'static str,
items: Vec<ContextItem>,
policy: ContextFailurePolicy,
}
#[async_trait]
impl ContextProvider for FixedProvider {
fn id(&self) -> &str {
self.id
}
fn failure_policy(&self) -> ContextFailurePolicy {
self.policy
}
async fn load(&self, _request: &ContextRequest) -> Result<Vec<ContextItem>, ContextError> {
if self.items.is_empty() {
return Err(ContextError::Load("boom".into()));
}
Ok(self.items.clone())
}
}
fn request() -> ContextRequest {
ContextRequest {
session_id: SessionId::from("s"),
run_id: RunId::from("r"),
turn: 1,
}
}
fn item(source: &str, text: &str) -> ContextItem {
ContextItem {
source: source.into(),
version: "1".into(),
text: text.into(),
budget_tokens: None,
persistable: true,
allowed_providers: None,
failure: ContextFailurePolicy::Block,
}
}
#[tokio::test]
async fn block_failure_propagates_error() {
let providers: Vec<Arc<dyn ContextProvider>> = vec![Arc::new(FixedProvider {
id: "strict",
items: Vec::new(),
policy: ContextFailurePolicy::Block,
})];
let result = load_context(&providers, &request()).await;
assert!(matches!(result, Err(ContextError::Load(_))));
}
#[tokio::test]
async fn ignore_and_warn_failures_skip_provider() {
let ok: Arc<dyn ContextProvider> = Arc::new(FixedProvider {
id: "ok",
items: vec![item("ok-source", "text")],
policy: ContextFailurePolicy::Block,
});
for policy in [ContextFailurePolicy::Ignore, ContextFailurePolicy::Warn] {
let failing: Arc<dyn ContextProvider> = Arc::new(FixedProvider {
id: "flaky",
items: Vec::new(),
policy,
});
let items = load_context(&[failing, ok.clone()], &request())
.await
.unwrap();
assert_eq!(items.len(), 1);
assert_eq!(items[0].source, "ok-source");
}
}
#[test]
fn budget_truncation_is_exact() {
let text = "a".repeat(100);
assert_eq!(truncate_to_budget(&text, None), text);
let truncated = truncate_to_budget(&text, Some(10));
assert_eq!(truncated, format!("{}\n…(truncated)", "a".repeat(40)));
assert_eq!(truncate_to_budget("short", Some(10)), "short");
}
#[tokio::test]
async fn default_failure_policy_is_block() {
struct BareProvider;
#[async_trait]
impl ContextProvider for BareProvider {
fn id(&self) -> &str {
"bare"
}
async fn load(
&self,
_request: &ContextRequest,
) -> Result<Vec<ContextItem>, ContextError> {
Err(ContextError::Load("boom".into()))
}
}
let providers: Vec<Arc<dyn ContextProvider>> = vec![Arc::new(BareProvider)];
assert!(load_context(&providers, &request()).await.is_err());
}
}