1use std::collections::HashMap;
15use std::sync::atomic::{AtomicU64, Ordering};
16
17use lattice_grammar::CommandId;
18use lattice_grammar::source::SourceLocation;
19
20use crate::cache::GeneratorCache;
21use crate::traits::{CandidateAnnotator, CandidateGenerator, CandidateMatcher, CandidateRanker};
22
23#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
25pub struct GeneratorId(pub CommandId);
26
27#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
28pub struct MatcherId(pub CommandId);
29
30#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
31pub struct RankerId(pub CommandId);
32
33#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
34pub struct AnnotatorId(pub CommandId);
35
36pub struct RegisteredGenerator {
40 pub id: GeneratorId,
41 pub name: String,
42 pub doc: String,
43 pub source: SourceLocation,
44 pub inner: std::sync::Arc<dyn CandidateGenerator>,
45}
46
47pub struct RegisteredMatcher {
48 pub id: MatcherId,
49 pub name: String,
50 pub doc: String,
51 pub source: SourceLocation,
52 pub inner: std::sync::Arc<dyn CandidateMatcher>,
53}
54
55pub struct RegisteredRanker {
56 pub id: RankerId,
57 pub name: String,
58 pub doc: String,
59 pub source: SourceLocation,
60 pub inner: std::sync::Arc<dyn CandidateRanker>,
61}
62
63pub struct RegisteredAnnotator {
64 pub id: AnnotatorId,
65 pub name: String,
66 pub doc: String,
67 pub source: SourceLocation,
68 pub inner: std::sync::Arc<dyn CandidateAnnotator>,
69}
70
71#[derive(Default)]
72pub struct CompletionRegistry {
73 generators: HashMap<GeneratorId, RegisteredGenerator>,
74 matchers: HashMap<MatcherId, RegisteredMatcher>,
75 rankers: HashMap<RankerId, RegisteredRanker>,
76 annotators: HashMap<AnnotatorId, RegisteredAnnotator>,
77
78 sources: HashMap<String, crate::source_registration::SourceRegistration>,
91
92 pub default_matcher: Option<MatcherId>,
96 pub default_rankers: Vec<RankerId>,
104 pub default_annotators: Vec<AnnotatorId>,
109
110 pub cache: GeneratorCache,
112}
113
114impl std::fmt::Debug for CompletionRegistry {
115 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
116 f.debug_struct("CompletionRegistry")
117 .field("generators", &self.generators.len())
118 .field("matchers", &self.matchers.len())
119 .field("rankers", &self.rankers.len())
120 .field("annotators", &self.annotators.len())
121 .field("sources", &self.sources.len())
122 .field("default_matcher", &self.default_matcher)
123 .field("default_rankers", &self.default_rankers)
124 .field("default_annotators", &self.default_annotators)
125 .finish_non_exhaustive()
126 }
127}
128
129impl CompletionRegistry {
130 pub fn new() -> Self {
131 Self::default()
132 }
133
134 #[track_caller]
137 pub fn register_generator(
138 &mut self,
139 name: &str,
140 doc: &str,
141 generator: impl CandidateGenerator + 'static,
142 ) -> GeneratorId {
143 let source = capture_builtin_source();
144 self.insert_generator(name, doc, std::sync::Arc::new(generator), source)
145 }
146
147 #[track_caller]
148 pub fn register_matcher(
149 &mut self,
150 name: &str,
151 doc: &str,
152 m: impl CandidateMatcher + 'static,
153 ) -> MatcherId {
154 let source = capture_builtin_source();
155 self.insert_matcher(name, doc, std::sync::Arc::new(m), source)
156 }
157
158 #[track_caller]
159 pub fn register_ranker(
160 &mut self,
161 name: &str,
162 doc: &str,
163 r: impl CandidateRanker + 'static,
164 ) -> RankerId {
165 let source = capture_builtin_source();
166 self.insert_ranker(name, doc, std::sync::Arc::new(r), source)
167 }
168
169 #[track_caller]
170 pub fn register_annotator(
171 &mut self,
172 name: &str,
173 doc: &str,
174 a: impl CandidateAnnotator + 'static,
175 ) -> AnnotatorId {
176 let source = capture_builtin_source();
177 self.insert_annotator(name, doc, std::sync::Arc::new(a), source)
178 }
179
180 pub fn register_source(&mut self, reg: crate::source_registration::SourceRegistration) {
194 let id = reg.spec.id.clone();
195 self.sources.insert(id, reg);
196 }
197
198 pub fn source_by_id(
200 &self,
201 id: &str,
202 ) -> Option<&crate::source_registration::SourceRegistration> {
203 self.sources.get(id)
204 }
205
206 pub fn sources(&self) -> impl Iterator<Item = &crate::source_registration::SourceRegistration> {
211 self.sources.values()
212 }
213
214 pub fn source_count(&self) -> usize {
216 self.sources.len()
217 }
218
219 pub(crate) fn insert_generator(
224 &mut self,
225 name: &str,
226 doc: &str,
227 inner: std::sync::Arc<dyn CandidateGenerator>,
228 source: SourceLocation,
229 ) -> GeneratorId {
230 let id = GeneratorId(next_id());
231 self.generators.insert(
232 id,
233 RegisteredGenerator {
234 id,
235 name: name.to_string(),
236 doc: doc.to_string(),
237 source,
238 inner,
239 },
240 );
241 id
242 }
243
244 pub(crate) fn insert_matcher(
245 &mut self,
246 name: &str,
247 doc: &str,
248 inner: std::sync::Arc<dyn CandidateMatcher>,
249 source: SourceLocation,
250 ) -> MatcherId {
251 let id = MatcherId(next_id());
252 self.matchers.insert(
253 id,
254 RegisteredMatcher {
255 id,
256 name: name.to_string(),
257 doc: doc.to_string(),
258 source,
259 inner,
260 },
261 );
262 id
263 }
264
265 pub(crate) fn insert_ranker(
266 &mut self,
267 name: &str,
268 doc: &str,
269 inner: std::sync::Arc<dyn CandidateRanker>,
270 source: SourceLocation,
271 ) -> RankerId {
272 let id = RankerId(next_id());
273 self.rankers.insert(
274 id,
275 RegisteredRanker {
276 id,
277 name: name.to_string(),
278 doc: doc.to_string(),
279 source,
280 inner,
281 },
282 );
283 id
284 }
285
286 pub(crate) fn insert_annotator(
287 &mut self,
288 name: &str,
289 doc: &str,
290 inner: std::sync::Arc<dyn CandidateAnnotator>,
291 source: SourceLocation,
292 ) -> AnnotatorId {
293 let id = AnnotatorId(next_id());
294 self.annotators.insert(
295 id,
296 RegisteredAnnotator {
297 id,
298 name: name.to_string(),
299 doc: doc.to_string(),
300 source,
301 inner,
302 },
303 );
304 id
305 }
306
307 pub fn generator(&self, id: GeneratorId) -> Option<&RegisteredGenerator> {
310 self.generators.get(&id)
311 }
312 pub fn matcher(&self, id: MatcherId) -> Option<&RegisteredMatcher> {
313 self.matchers.get(&id)
314 }
315 pub fn ranker(&self, id: RankerId) -> Option<&RegisteredRanker> {
316 self.rankers.get(&id)
317 }
318 pub fn annotator(&self, id: AnnotatorId) -> Option<&RegisteredAnnotator> {
319 self.annotators.get(&id)
320 }
321
322 pub fn generator_by_name(&self, name: &str) -> Option<&RegisteredGenerator> {
323 self.generators.values().find(|g| g.name == name)
324 }
325 pub fn matcher_by_name(&self, name: &str) -> Option<&RegisteredMatcher> {
326 self.matchers.values().find(|m| m.name == name)
327 }
328 pub fn ranker_by_name(&self, name: &str) -> Option<&RegisteredRanker> {
329 self.rankers.values().find(|r| r.name == name)
330 }
331 pub fn annotator_by_name(&self, name: &str) -> Option<&RegisteredAnnotator> {
332 self.annotators.values().find(|a| a.name == name)
333 }
334
335 pub fn generator_count(&self) -> usize {
336 self.generators.len()
337 }
338}
339
340fn next_id() -> CommandId {
341 static NEXT: AtomicU64 = AtomicU64::new(1);
342 CommandId::new(NEXT.fetch_add(1, Ordering::Relaxed))
343}
344
345#[track_caller]
346fn capture_builtin_source() -> SourceLocation {
347 let loc = std::panic::Location::caller();
348 SourceLocation {
349 layer: lattice_grammar::SourceLayer::Builtin,
350 kind: lattice_grammar::SourceKind::File {
351 path: std::path::PathBuf::from(loc.file()),
352 line: Some(loc.line()),
353 },
354 }
355}
356
357#[cfg(test)]
358mod tests {
359 #![allow(clippy::unwrap_used, clippy::panic)]
360 use super::*;
361 use crate::candidate::{MatchScore, RawCandidate, RenderedCandidate, ScoredCandidate};
362 use crate::traits::GenerateContext;
363
364 struct StubGen;
365 impl CandidateGenerator for StubGen {
366 fn generate(&self, _: &GenerateContext<'_>) -> Vec<RawCandidate> {
367 Vec::new()
368 }
369 }
370
371 struct StubMatch;
372 impl CandidateMatcher for StubMatch {
373 fn matches(
374 &self,
375 _: &str,
376 _: &RawCandidate,
377 ) -> Option<(MatchScore, Vec<std::ops::Range<usize>>)> {
378 None
379 }
380 }
381
382 struct StubRank;
383 impl CandidateRanker for StubRank {
384 fn rank(&self, _: &mut Vec<ScoredCandidate>) {}
385 }
386
387 struct StubAnno;
388 impl CandidateAnnotator for StubAnno {
389 fn annotate(&self, _: &mut RenderedCandidate) {}
390 }
391
392 #[test]
393 fn empty_registry() {
394 let r = CompletionRegistry::new();
395 assert_eq!(r.generator_count(), 0);
396 assert!(r.default_matcher.is_none());
397 }
398
399 #[test]
400 fn register_returns_id_and_finds_by_name() {
401 let mut r = CompletionRegistry::new();
402 let id = r.register_generator("gen:test", "doc", StubGen);
403 assert!(r.generator(id).is_some());
404 assert_eq!(r.generator_by_name("gen:test").map(|g| g.id), Some(id));
405 }
406
407 #[test]
408 fn distinct_ids_for_distinct_registrations() {
409 let mut r = CompletionRegistry::new();
410 let a = r.register_generator("a", "", StubGen);
411 let b = r.register_generator("b", "", StubGen);
412 assert_ne!(a, b);
413 }
414
415 #[test]
416 fn each_kind_has_independent_namespace() {
417 let mut r = CompletionRegistry::new();
418 let _g = r.register_generator("x", "", StubGen);
419 let _m = r.register_matcher("x", "", StubMatch);
420 let _rk = r.register_ranker("x", "", StubRank);
421 let _a = r.register_annotator("x", "", StubAnno);
422 assert!(r.generator_by_name("x").is_some());
423 assert!(r.matcher_by_name("x").is_some());
424 assert!(r.ranker_by_name("x").is_some());
425 assert!(r.annotator_by_name("x").is_some());
426 }
427
428 #[test]
429 fn track_caller_records_registration_site() {
430 let mut r = CompletionRegistry::new();
431 let expected = line!() + 1;
432 let id = r.register_generator("gen:caller-test", "", StubGen);
433 let g = r.generator(id).unwrap();
434 match &g.source.kind {
435 lattice_grammar::SourceKind::File {
436 path,
437 line: Some(line),
438 } => {
439 assert!(path.to_string_lossy().ends_with("registry.rs"));
440 assert_eq!(*line, expected);
441 }
442 other => panic!("expected File source, got {other:?}"),
443 }
444 }
445
446 #[test]
447 fn default_slots_start_unset() {
448 let r = CompletionRegistry::new();
449 assert!(r.default_matcher.is_none());
450 assert!(r.default_rankers.is_empty());
451 assert!(r.default_annotators.is_empty());
452 }
453
454 #[test]
455 fn default_annotators_can_be_appended() {
456 let mut r = CompletionRegistry::new();
457 let a1 = r.register_annotator("a1", "", StubAnno);
458 let a2 = r.register_annotator("a2", "", StubAnno);
459 r.default_annotators.push(a1);
460 r.default_annotators.push(a2);
461 assert_eq!(r.default_annotators, vec![a1, a2]);
462 }
463
464 #[test]
470 fn register_source_round_trips_by_id() {
471 use crate::candidate::{CandidateKind, RawCandidate};
472 use crate::source_registration::{CandidateSourceKind, SourceRegistration, SourceSpec};
473
474 let mut r = CompletionRegistry::new();
475 assert_eq!(r.source_count(), 0);
476
477 let rows = vec![RawCandidate::plain("hello", CandidateKind::Plain)];
478 let reg = SourceRegistration {
479 spec: SourceSpec {
480 id: "test:smoke".to_string(),
481 doc: "smoke test source".to_string(),
482 args_schema: None,
483 live: false,
484 },
485 kind: CandidateSourceKind::PreSupplied(std::sync::Arc::new(rows)),
486 accept: None,
487 matcher_override: None,
488 ranker_overrides: Vec::new(),
489 annotator_extras: Vec::new(),
490 };
491 r.register_source(reg);
492
493 assert_eq!(r.source_count(), 1);
494 let looked_up = r.source_by_id("test:smoke").expect("must be registered");
495 assert_eq!(looked_up.spec.id, "test:smoke");
496 assert_eq!(looked_up.spec.doc, "smoke test source");
497 assert!(matches!(
498 looked_up.kind,
499 CandidateSourceKind::PreSupplied(_)
500 ));
501 assert!(r.source_by_id("nope").is_none());
502 }
503
504 #[test]
508 fn register_source_last_write_wins_on_duplicate_id() {
509 use crate::source_registration::{CandidateSourceKind, SourceRegistration, SourceSpec};
510
511 let mut r = CompletionRegistry::new();
512 let first = SourceRegistration {
513 spec: SourceSpec {
514 id: "test:dup".to_string(),
515 doc: "first".to_string(),
516 args_schema: None,
517 live: false,
518 },
519 kind: CandidateSourceKind::PreSupplied(std::sync::Arc::new(Vec::new())),
520 accept: None,
521 matcher_override: None,
522 ranker_overrides: Vec::new(),
523 annotator_extras: Vec::new(),
524 };
525 let second = SourceRegistration {
526 spec: SourceSpec {
527 id: "test:dup".to_string(),
528 doc: "second".to_string(),
529 args_schema: None,
530 live: true,
531 },
532 kind: CandidateSourceKind::PreSupplied(std::sync::Arc::new(Vec::new())),
533 accept: None,
534 matcher_override: None,
535 ranker_overrides: Vec::new(),
536 annotator_extras: Vec::new(),
537 };
538 r.register_source(first);
539 r.register_source(second);
540 assert_eq!(r.source_count(), 1);
541 let looked_up = r.source_by_id("test:dup").expect("must be registered");
542 assert_eq!(looked_up.spec.doc, "second");
543 assert!(looked_up.spec.live);
544 }
545}