Skip to content

Commit 997b96e

Browse files
enaboappsOwen McGirr
andauthored
Keep native-context predictions stable without fallback tracking (#806)
Separate fallback tracking availability from native context caching and stop unchanged polls from resetting prediction batches. Cover scan timing, pass limits, selected choices and stale-context rejection. Co-authored-by: Owen McGirr <owenmcgirr@Owens-Mac-Studio-2.local>
1 parent eac20ae commit 997b96e

3 files changed

Lines changed: 189 additions & 8 deletions

File tree

src-tauri/src/prediction/mod.rs

Lines changed: 13 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -152,6 +152,18 @@ struct Service {
152152
tracking: bool,
153153
}
154154
impl Service {
155+
fn suggestions(
156+
&mut self,
157+
keyboard: &mut Keyboard,
158+
batch: Option<worker::Batch>,
159+
tracking: bool,
160+
) {
161+
self.tracking = tracking;
162+
if !tracking {
163+
self.edit = None;
164+
}
165+
keyboard.predictions(batch, false);
166+
}
155167
fn fail(&mut self, keyboard: &mut Keyboard) {
156168
self.client = None;
157169
self.failed = true;
@@ -270,12 +282,7 @@ pub fn poll(app: &AppHandle, keyboard: Option<&mut Keyboard>, enabled: bool, ign
270282
batch,
271283
tracking,
272284
} if generation == s.generation => {
273-
s.tracking = tracking;
274-
if !tracking {
275-
s.edit = None;
276-
s.reset = true;
277-
}
278-
keyboard.predictions(batch, false);
285+
s.suggestions(keyboard, batch, tracking);
279286
}
280287
Response::Insert { generation, text } => {
281288
s.accepting = false;

src-tauri/src/prediction/worker.rs

Lines changed: 105 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -139,9 +139,16 @@ impl<A: Adapter> Engine<A> {
139139
.as_ref()
140140
.is_some_and(|old| self.adapter.same(old, &target).unwrap_or(false));
141141
let uninterrupted = epoch == self.activity;
142-
if !same || !uninterrupted || reset || !healthy {
142+
if !same || !uninterrupted || reset {
143143
self.clear();
144144
}
145+
if !healthy {
146+
self.fallback.clear();
147+
self.boundary = false;
148+
if self.snapshot.as_ref().is_some_and(|raw| raw.position == -1) {
149+
self.clear();
150+
}
151+
}
145152
self.activity = epoch;
146153
if same && uninterrupted && !reset && self.tracked && healthy {
147154
if let Some(text) = edit {
@@ -369,6 +376,7 @@ mod tests {
369376
target: usize,
370377
protected: bool,
371378
unsupported: bool,
379+
editable: bool,
372380
reads: usize,
373381
change_on_read: bool,
374382
}
@@ -385,6 +393,7 @@ mod tests {
385393
target: 1,
386394
protected: false,
387395
unsupported: false,
396+
editable: true,
388397
reads: 0,
389398
change_on_read: false,
390399
}
@@ -399,7 +408,7 @@ mod tests {
399408
Ok(self.protected)
400409
}
401410
fn editable(&mut self, _: &usize) -> bool {
402-
true
411+
self.editable
403412
}
404413
fn same(&mut self, a: &usize, b: &usize) -> Result<bool, Status> {
405414
Ok(a == b)
@@ -423,6 +432,100 @@ mod tests {
423432
e
424433
}
425434
#[test]
435+
fn native_predictions_survive_polls_without_fallback_tracking() {
436+
for unavailable in 0..3 {
437+
let mut e = engine();
438+
match unavailable {
439+
0 => e.tracked = false,
440+
1 => e.adapter.editable = false,
441+
_ => e.observe = || (0, false),
442+
}
443+
let mut service = crate::prediction::Service::default();
444+
let mut keyboard = crate::scan_keyboard::Keyboard::new(false);
445+
keyboard.enable_predictions(true);
446+
let mut original = None;
447+
for _ in 0..40 {
448+
let response = e.respond(Request::Query {
449+
generation: 0,
450+
edit: service.edit.take(),
451+
reset: std::mem::take(&mut service.reset),
452+
shift: false,
453+
caps: false,
454+
});
455+
let Response::Suggestions {
456+
batch, tracking, ..
457+
} = response
458+
else {
459+
panic!("expected suggestions");
460+
};
461+
assert!(!tracking);
462+
let batch = batch.unwrap();
463+
let expected = original.get_or_insert_with(|| batch.clone());
464+
assert_eq!(batch.token, expected.token);
465+
assert_eq!(batch.words, expected.words);
466+
service.suggestions(&mut keyboard, Some(batch), tracking);
467+
keyboard.advance(250, 1000);
468+
}
469+
let token = original.unwrap().token;
470+
assert_eq!(e.accept(token, 0), Some("ter ".into()));
471+
assert!(e.accept(token, 0).is_none());
472+
}
473+
}
474+
#[test]
475+
fn genuine_context_changes_invalidate_identical_native_candidates() {
476+
for change in 0..4 {
477+
let mut e = engine();
478+
e.adapter.editable = false;
479+
let original = e.query(None, false, false, false).unwrap();
480+
match change {
481+
0 => e.adapter.target += 1,
482+
1 => e.adapter.raw.position += 1,
483+
2 => e.activity = u64::MAX,
484+
_ => {}
485+
}
486+
let replacement = e.query(None, change == 3, false, false).unwrap();
487+
assert_eq!(replacement.words, original.words);
488+
assert_ne!(replacement.token, original.token);
489+
assert!(e.accept(original.token, 0).is_none());
490+
assert_eq!(e.accept(replacement.token, 0), Some("ter ".into()));
491+
}
492+
}
493+
#[test]
494+
fn native_edits_selection_and_casing_refresh_without_tracking() {
495+
let mut e = engine();
496+
e.tracked = false;
497+
let original = e.query(None, false, false, false).unwrap();
498+
e.adapter.raw.before.push('t');
499+
let typed = e.query(None, false, false, false).unwrap();
500+
assert_ne!(typed.token, original.token);
501+
assert!(e.accept(original.token, 0).is_none());
502+
e.adapter.raw.before.pop();
503+
let deleted = e.query(None, false, false, false).unwrap();
504+
assert_ne!(deleted.token, typed.token);
505+
let shifted = e.query(None, false, true, false).unwrap();
506+
assert_ne!(shifted.token, deleted.token);
507+
assert!(e.accept(deleted.token, 0).is_none());
508+
e.adapter.raw.has_selection = true;
509+
assert!(e.query(None, false, false, false).is_none());
510+
assert!(e.accept(shifted.token, 0).is_none());
511+
e.adapter.raw.has_selection = false;
512+
let restored = e.query(None, false, false, false).unwrap();
513+
e.adapter.protected = true;
514+
assert!(e.query(None, false, false, false).is_none());
515+
assert!(e.accept(restored.token, 0).is_none());
516+
}
517+
#[test]
518+
fn losing_tracking_does_not_erase_a_pending_context_reset() {
519+
let mut service = crate::prediction::Service {
520+
reset: true,
521+
edit: Some("queued".into()),
522+
..Default::default()
523+
};
524+
service.suggestions(&mut crate::scan_keyboard::Keyboard::new(false), None, false);
525+
assert!(service.reset);
526+
assert!(service.edit.is_none());
527+
}
528+
#[test]
426529
fn only_suffix_and_space_are_returned_once() {
427530
let mut e = engine();
428531
let b = e.query(None, false, false, false).unwrap();

src-tauri/src/scan_keyboard.rs

Lines changed: 71 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -619,6 +619,77 @@ mod tests {
619619
assert_eq!(keyboard.scan.nav.index(), 1);
620620
}
621621

622+
#[test]
623+
fn unchanged_predictions_preserve_navigation_timing_and_pass_limits() {
624+
use crate::scan_preferences::{Direction, Pattern, Resolved};
625+
for direction in [Direction::Forward, Direction::Reverse] {
626+
for pattern in [Pattern::Grouped, Pattern::Linear] {
627+
for automatic in [false, true] {
628+
let options = Resolved {
629+
direction,
630+
pattern,
631+
pass_limit: 2,
632+
..Default::default()
633+
};
634+
let mut baseline = Keyboard::configured(false, options);
635+
let mut polled = Keyboard::configured(false, options);
636+
let batch = crate::prediction::worker::Batch {
637+
token: 7,
638+
words: vec!["water".into(), "walk".into()],
639+
};
640+
for k in [&mut baseline, &mut polled] {
641+
k.enable_predictions(true);
642+
k.predictions(Some(batch.clone()), false);
643+
k.restart();
644+
}
645+
let ticks = 12 * (polled.rows.iter().map(Vec::len).sum::<usize>() + 1);
646+
for tick in 0..ticks {
647+
polled.predictions(Some(batch.clone()), false);
648+
if automatic {
649+
baseline.advance(250, 1000);
650+
polled.advance(250, 1000);
651+
} else if tick % 4 == 3 {
652+
baseline.handle(Action::Next);
653+
polled.handle(Action::Next);
654+
}
655+
assert_eq!(
656+
polled.scan.position(&polled.rows),
657+
baseline.scan.position(&baseline.rows)
658+
);
659+
assert_eq!(polled.suspended(), baseline.suspended());
660+
assert_eq!(polled.disabled(), baseline.disabled());
661+
assert_eq!(polled.predictions.as_ref().unwrap().token, 7);
662+
assert!(polled.queued_predictions.is_none());
663+
}
664+
if automatic {
665+
assert!(polled.suspended());
666+
}
667+
}
668+
}
669+
}
670+
}
671+
#[test]
672+
fn unchanged_predictions_preserve_selected_key_and_execute_once() {
673+
let mut k = Keyboard::new(false);
674+
k.enable_predictions(true);
675+
let batch = crate::prediction::worker::Batch {
676+
token: 7,
677+
words: vec!["water".into(), "walk".into()],
678+
};
679+
k.predictions(Some(batch.clone()), false);
680+
k.handle(Action::Select);
681+
k.advance(750, 1000);
682+
k.predictions(Some(batch.clone()), false);
683+
assert_eq!(k.scan.position(&k.rows), (0, Some(0)));
684+
k.advance(250, 1000);
685+
assert_eq!(k.scan.position(&k.rows), (0, Some(1)));
686+
k.predictions(Some(batch), false);
687+
assert_eq!(
688+
k.handle(Action::Select),
689+
Some(Output::Prediction { token: 7, index: 1 })
690+
);
691+
assert!(k.handle(Action::Select).is_none());
692+
}
622693
#[test]
623694
fn predictions_skip_empty_slots_and_defer_acceptance() {
624695
let mut k = Keyboard::new(false);

0 commit comments

Comments
 (0)