1#![allow(clippy::single_range_in_vec_init)]
2use std::sync::Arc;
13
14use crate::cache::GeneratorCache;
15use crate::candidate::{RawCandidate, RenderedCandidate, ScoredCandidate};
16use crate::registry::CompletionRegistry;
17use crate::traits::{
18 CandidateAnnotator, CandidateGenerator, CandidateMatcher, CandidateRanker, GenerateContext,
19};
20
21pub struct CompletionPipeline {
33 pub generators: Vec<Arc<dyn CandidateGenerator>>,
34 pub matcher: Arc<dyn CandidateMatcher>,
35 pub rankers: Vec<Arc<dyn CandidateRanker>>,
36 pub annotators: Vec<Arc<dyn CandidateAnnotator>>,
37}
38
39impl CompletionPipeline {
40 pub fn run(
44 &self,
45 ctx: &GenerateContext<'_>,
46 query: &str,
47 cache: &GeneratorCache,
48 ) -> Vec<RenderedCandidate> {
49 let mut raw: Vec<RawCandidate> = Vec::new();
51 for g in &self.generators {
52 let from_cache = match g.cache_key(ctx) {
53 Some(key) => cache.get(&key).map(|cached| (key, cached, g.cache_ttl())),
54 None => None,
55 };
56 match from_cache {
57 Some((_, cached, _)) => {
58 raw.extend(cached);
59 }
60 None => {
61 let produced = g.generate(ctx);
62 if let Some(key) = g.cache_key(ctx) {
63 cache.put(key, produced.clone(), g.cache_ttl());
64 }
65 raw.extend(produced);
66 }
67 }
68 }
69
70 let mut scored: Vec<ScoredCandidate> = raw
72 .into_iter()
73 .filter_map(|c| {
74 self.matcher
75 .matches(query, &c)
76 .map(|(score, ranges)| ScoredCandidate {
77 raw: c,
78 score,
79 match_ranges: ranges,
80 })
81 })
82 .collect();
83
84 for r in &self.rankers {
90 r.rank(&mut scored);
91 }
92
93 let mut rendered: Vec<RenderedCandidate> = scored
95 .into_iter()
96 .map(RenderedCandidate::from_scored)
97 .collect();
98 for a in &self.annotators {
99 for c in rendered.iter_mut() {
100 a.annotate(c);
101 }
102 }
103 rendered
104 }
105}
106
107impl CompletionPipeline {
108 pub fn match_and_rank(&self, query: &str, raw: &[RawCandidate]) -> Vec<RenderedCandidate> {
129 let mut scored: Vec<ScoredCandidate> = raw
130 .iter()
131 .filter_map(|c| {
132 self.matcher
133 .matches(query, c)
134 .map(|(score, ranges)| ScoredCandidate {
135 raw: c.clone(),
136 score,
137 match_ranges: ranges,
138 })
139 })
140 .collect();
141 for r in &self.rankers {
142 r.rank(&mut scored);
143 }
144 scored
145 .into_iter()
146 .map(RenderedCandidate::from_scored)
147 .collect()
148 }
149}
150
151impl CompletionPipeline {
157 pub fn for_generator(
158 registry: &CompletionRegistry,
159 generator: crate::registry::GeneratorId,
160 ) -> Option<Self> {
161 let g = registry.generator(generator)?;
162 let m = registry.matcher(registry.default_matcher?)?;
163 if registry.default_rankers.is_empty() {
170 return None;
171 }
172 let rankers: Vec<_> = registry
173 .default_rankers
174 .iter()
175 .filter_map(|id| registry.ranker(*id))
176 .map(|r| r.inner.clone())
177 .collect();
178 if rankers.is_empty() {
179 return None;
180 }
181 let annotators: Vec<_> = registry
182 .default_annotators
183 .iter()
184 .filter_map(|id| registry.annotator(*id))
185 .map(|a| a.inner.clone())
186 .collect();
187 Some(Self {
188 generators: vec![g.inner.clone()],
189 matcher: m.inner.clone(),
190 rankers,
191 annotators,
192 })
193 }
194}
195
196#[cfg(test)]
197mod tests {
198 #![allow(clippy::unwrap_used, clippy::panic)]
199 use super::*;
200 use crate::candidate::{CacheKey, CandidateKind, MatchScore, RawCandidate};
201 use lattice_core::{Buffer, Document};
202 use lattice_grammar::CommandRegistry;
203 use std::ops::Range;
204 use std::sync::atomic::{AtomicUsize, Ordering};
205
206 struct CountingGen {
209 items: Vec<String>,
210 calls: AtomicUsize,
211 cache_key: Option<CacheKey>,
212 }
213
214 impl CountingGen {
215 fn new(items: Vec<&str>, cache_key: Option<&str>) -> Self {
216 Self {
217 items: items.into_iter().map(String::from).collect(),
218 calls: AtomicUsize::new(0),
219 cache_key: cache_key.map(CacheKey::new),
220 }
221 }
222 }
223
224 impl CandidateGenerator for CountingGen {
225 fn generate(&self, _: &GenerateContext<'_>) -> Vec<RawCandidate> {
226 self.calls.fetch_add(1, Ordering::Relaxed);
227 self.items
228 .iter()
229 .map(|s| RawCandidate::plain(s, CandidateKind::Plain))
230 .collect()
231 }
232 fn cache_key(&self, _: &GenerateContext<'_>) -> Option<CacheKey> {
233 self.cache_key.clone()
234 }
235 }
236
237 struct PrefixMatch;
239 impl CandidateMatcher for PrefixMatch {
240 fn matches(
241 &self,
242 query: &str,
243 c: &RawCandidate,
244 ) -> Option<(MatchScore, Vec<Range<usize>>)> {
245 if c.text.starts_with(query) {
246 Some((MatchScore::PREFIX, vec![0..query.len()]))
247 } else {
248 None
249 }
250 }
251 }
252
253 struct ScoreRank;
255 impl CandidateRanker for ScoreRank {
256 fn rank(&self, scored: &mut Vec<ScoredCandidate>) {
257 scored.sort_by(|a, b| b.score.cmp(&a.score));
258 }
259 }
260
261 struct LengthAnno;
266 impl CandidateAnnotator for LengthAnno {
267 fn annotate(&self, c: &mut RenderedCandidate) {
268 c.annotations.push(crate::candidate::Annotation::Custom {
269 text: std::sync::Arc::from(format!("{} chars", c.raw.text.len())),
270 slot: std::sync::Arc::from("annotation_test_length"),
271 });
272 }
273 }
274
275 fn ctx<'a>(
276 prefix: &'a str,
277 buffer: &'a Buffer,
278 registry: &'a CommandRegistry,
279 ) -> GenerateContext<'a> {
280 GenerateContext {
281 prefix,
282 buffer,
283 registry,
284 case_sensitive: false,
285 }
286 }
287
288 #[test]
289 fn pipeline_runs_all_four_stages() {
290 let document = Document::empty();
291 let buffer = document.buffer().clone();
292 let registry = CommandRegistry::new();
293 let cache = GeneratorCache::new();
294 let p = CompletionPipeline {
295 generators: vec![Arc::new(CountingGen::new(
296 vec!["alpha", "beta", "alphabet"],
297 None,
298 ))],
299 matcher: Arc::new(PrefixMatch),
300 rankers: vec![Arc::new(ScoreRank)],
301 annotators: vec![Arc::new(LengthAnno)],
302 };
303 let result = p.run(&ctx("alph", &buffer, ®istry), "alph", &cache);
304 assert_eq!(result.len(), 2);
306 assert!(result[0].annotations[0].display_text().contains("chars"));
308 }
309
310 #[test]
311 fn caching_prevents_regeneration_on_second_run() {
312 let document = Document::empty();
313 let buffer = document.buffer().clone();
314 let registry = CommandRegistry::new();
315 let cache = GeneratorCache::new();
316 let generator = Arc::new(CountingGen::new(vec!["x", "y"], Some("k1")));
317 let p = CompletionPipeline {
318 generators: vec![generator.clone()],
319 matcher: Arc::new(PrefixMatch),
320 rankers: vec![Arc::new(ScoreRank)],
321 annotators: vec![],
322 };
323 let _ = p.run(&ctx("", &buffer, ®istry), "", &cache);
324 let _ = p.run(&ctx("", &buffer, ®istry), "", &cache);
325 let _ = p.run(&ctx("", &buffer, ®istry), "", &cache);
326 assert_eq!(generator.calls.load(Ordering::Relaxed), 1);
328 }
329
330 #[test]
331 fn no_cache_key_means_no_caching() {
332 let document = Document::empty();
333 let buffer = document.buffer().clone();
334 let registry = CommandRegistry::new();
335 let cache = GeneratorCache::new();
336 let generator = Arc::new(CountingGen::new(vec!["x"], None));
337 let p = CompletionPipeline {
338 generators: vec![generator.clone()],
339 matcher: Arc::new(PrefixMatch),
340 rankers: vec![Arc::new(ScoreRank)],
341 annotators: vec![],
342 };
343 let _ = p.run(&ctx("", &buffer, ®istry), "", &cache);
344 let _ = p.run(&ctx("", &buffer, ®istry), "", &cache);
345 assert_eq!(generator.calls.load(Ordering::Relaxed), 2);
346 }
347
348 #[test]
349 fn matcher_filters_non_matches() {
350 let document = Document::empty();
351 let buffer = document.buffer().clone();
352 let registry = CommandRegistry::new();
353 let cache = GeneratorCache::new();
354 let p = CompletionPipeline {
355 generators: vec![Arc::new(CountingGen::new(vec!["foo", "bar", "baz"], None))],
356 matcher: Arc::new(PrefixMatch),
357 rankers: vec![Arc::new(ScoreRank)],
358 annotators: vec![],
359 };
360 let result = p.run(&ctx("ba", &buffer, ®istry), "ba", &cache);
361 assert_eq!(result.len(), 2);
362 assert!(result.iter().all(|r| r.raw.text.starts_with("ba")));
363 }
364
365 #[test]
366 fn match_ranges_propagate_to_rendered_candidates() {
367 let document = Document::empty();
368 let buffer = document.buffer().clone();
369 let registry = CommandRegistry::new();
370 let cache = GeneratorCache::new();
371 let p = CompletionPipeline {
372 generators: vec![Arc::new(CountingGen::new(vec!["alpha"], None))],
373 matcher: Arc::new(PrefixMatch),
374 rankers: vec![Arc::new(ScoreRank)],
375 annotators: vec![],
376 };
377 let result = p.run(&ctx("alp", &buffer, ®istry), "alp", &cache);
378 assert_eq!(result[0].match_ranges, vec![0..3]);
379 }
380
381 #[test]
382 fn empty_query_with_prefix_matcher_matches_all() {
383 let document = Document::empty();
384 let buffer = document.buffer().clone();
385 let registry = CommandRegistry::new();
386 let cache = GeneratorCache::new();
387 let p = CompletionPipeline {
388 generators: vec![Arc::new(CountingGen::new(vec!["a", "b", "c"], None))],
389 matcher: Arc::new(PrefixMatch),
390 rankers: vec![Arc::new(ScoreRank)],
391 annotators: vec![],
392 };
393 let result = p.run(&ctx("", &buffer, ®istry), "", &cache);
394 assert_eq!(result.len(), 3);
395 }
396
397 #[test]
398 fn for_generator_returns_none_when_default_matcher_unset() {
399 let registry = CompletionRegistry::new();
400 assert!(registry.default_matcher.is_none());
406 }
407}