Skip to main content

a3s_code_core/context/
static_provider.rs

1//! Static context provider for session-local context items.
2
3use super::{ContextItem, ContextProvider, ContextQuery, ContextResult};
4use crate::text::truncate_utf8;
5
6/// Provides pre-built context items through the same retrieval pipeline as
7/// external providers.
8#[derive(Debug, Clone)]
9pub struct StaticContextProvider {
10    name: String,
11    items: Vec<ContextItem>,
12}
13
14impl StaticContextProvider {
15    pub fn new(name: impl Into<String>) -> Self {
16        Self {
17            name: name.into(),
18            items: Vec::new(),
19        }
20    }
21
22    pub fn from_items(
23        name: impl Into<String>,
24        items: impl IntoIterator<Item = ContextItem>,
25    ) -> Self {
26        Self {
27            name: name.into(),
28            items: items.into_iter().collect(),
29        }
30    }
31
32    pub fn with_item(mut self, item: ContextItem) -> Self {
33        self.items.push(item);
34        self
35    }
36}
37
38#[async_trait::async_trait]
39impl ContextProvider for StaticContextProvider {
40    fn name(&self) -> &str {
41        &self.name
42    }
43
44    async fn query(&self, query: &ContextQuery) -> anyhow::Result<ContextResult> {
45        let mut result = ContextResult::new(&self.name);
46        let max_results = query.max_results.max(1);
47        let max_tokens = query.max_tokens.max(1);
48        let mut total_tokens = 0usize;
49
50        for item in self.items.iter().take(max_results) {
51            let item_tokens = estimated_tokens(item);
52            if item.is_required() {
53                total_tokens = total_tokens.saturating_add(item_tokens);
54                result.add_item(item.clone());
55                continue;
56            }
57            if total_tokens + item_tokens > max_tokens {
58                result.truncated = true;
59
60                if result.items.is_empty() {
61                    result.add_item(truncate_item(item, max_tokens));
62                }
63
64                break;
65            }
66
67            let mut item = item.clone();
68            if item.token_count == 0 {
69                item.token_count = item_tokens;
70            }
71            total_tokens += item_tokens;
72            result.add_item(item);
73        }
74
75        if self.items.len() > max_results {
76            result.truncated = true;
77        }
78
79        Ok(result)
80    }
81}
82
83fn estimated_tokens(item: &ContextItem) -> usize {
84    if item.token_count > 0 {
85        item.token_count
86    } else {
87        item.content.split_whitespace().count().max(1)
88    }
89}
90
91fn truncate_item(item: &ContextItem, max_tokens: usize) -> ContextItem {
92    let max_bytes = max_tokens.saturating_mul(4).max(1);
93    let mut truncated = item.clone();
94    let shown = truncate_utf8(&item.content, max_bytes).trim_end();
95    truncated.content = format!("{shown}\n\n[context truncated]");
96    truncated.token_count = max_tokens;
97    truncated
98}
99
100#[cfg(test)]
101mod tests {
102    use super::*;
103    use crate::context::ContextType;
104
105    #[tokio::test]
106    async fn returns_static_items() {
107        let provider = StaticContextProvider::new("static").with_item(
108            ContextItem::new("agents_md", ContextType::Resource, "Follow local rules")
109                .with_relevance(0.95),
110        );
111
112        let result = provider.query(&ContextQuery::new("prompt")).await.unwrap();
113
114        assert_eq!(result.provider, "static");
115        assert_eq!(result.items.len(), 1);
116        assert_eq!(result.items[0].id, "agents_md");
117        assert!(result.items[0].token_count > 0);
118        assert!(!result.truncated);
119    }
120
121    #[tokio::test]
122    async fn respects_result_limit() {
123        let provider = StaticContextProvider::from_items(
124            "static",
125            [
126                ContextItem::new("a", ContextType::Resource, "a").with_token_count(1),
127                ContextItem::new("b", ContextType::Resource, "b").with_token_count(1),
128            ],
129        );
130
131        let result = provider
132            .query(&ContextQuery::new("prompt").with_max_results(1))
133            .await
134            .unwrap();
135
136        assert_eq!(result.items.len(), 1);
137        assert!(result.truncated);
138    }
139
140    #[tokio::test]
141    async fn truncates_oversized_single_item() {
142        let provider = StaticContextProvider::new("static").with_item(
143            ContextItem::new(
144                "large",
145                ContextType::Resource,
146                "one two three four five six",
147            )
148            .with_token_count(6),
149        );
150
151        let result = provider
152            .query(&ContextQuery::new("prompt").with_max_tokens(2))
153            .await
154            .unwrap();
155
156        assert_eq!(result.items.len(), 1);
157        assert_eq!(result.items[0].token_count, 2);
158        assert!(result.items[0].content.contains("[context truncated]"));
159        assert!(result.truncated);
160    }
161
162    #[tokio::test]
163    async fn required_static_item_bypasses_retrieval_token_budget() {
164        let provider = StaticContextProvider::new("static").with_item(
165            ContextItem::new(
166                "instructions",
167                ContextType::Resource,
168                "one two three four five six",
169            )
170            .with_token_count(6)
171            .with_required(),
172        );
173
174        let result = provider
175            .query(&ContextQuery::new("prompt").with_max_tokens(2))
176            .await
177            .unwrap();
178
179        assert_eq!(result.items.len(), 1);
180        assert_eq!(result.items[0].token_count, 6);
181        assert!(!result.truncated);
182    }
183}