diff --git a/src/SIL.Machine/FiniteState/DeterministicFsaTraversalMethod.cs b/src/SIL.Machine/FiniteState/DeterministicFsaTraversalMethod.cs index 5470fa689..e06a6ad0a 100644 --- a/src/SIL.Machine/FiniteState/DeterministicFsaTraversalMethod.cs +++ b/src/SIL.Machine/FiniteState/DeterministicFsaTraversalMethod.cs @@ -22,7 +22,8 @@ public override IEnumerable> Traverse( ref int annIndex, Register[,] initRegisters, IList initCmds, - ISet initAnns + ISet initAnns, + bool allMatches ) { Stack> instStack = InitializeStack( @@ -33,15 +34,21 @@ ISet initAnns ); var curResults = new List>(); + var lattice = !allMatches ? CreateFstLattice() : null; + while (instStack.Count != 0) { DeterministicFsaTraversalInstance inst = instStack.Pop(); + DeterministicFsaTraversalInstance origInst = !allMatches + ? CopyInstanceAndBindings(inst) + : null; bool releaseInstance = true; foreach (Arc arc in inst.State.Arcs) { if (CheckInputMatch(arc, inst.AnnotationIndex, inst.VariableBindings)) { + int resultCount = !allMatches ? curResults.Count : 0; foreach ( DeterministicFsaTraversalInstance ni in Advance( inst, @@ -51,6 +58,13 @@ DeterministicFsaTraversalInstance ni in Advance( ) ) { + if (!allMatches) + { + if (curResults.Count > resultCount) + RecordFinalArc(lattice, origInst, arc); + if (RecordedInstance(lattice, ni, origInst, arc)) + continue; + } instStack.Push(ni); } @@ -63,6 +77,12 @@ DeterministicFsaTraversalInstance ni in Advance( ReleaseInstance(inst); } + if (!allMatches) + { + var newResults = ExtractResults(lattice, allMatches); + curResults = newResults; + } + CheckAcceptingStartState(initAnns, initRegisters, curResults); return curResults; diff --git a/src/SIL.Machine/FiniteState/DeterministicFstTraversalMethod.cs b/src/SIL.Machine/FiniteState/DeterministicFstTraversalMethod.cs index 534a2dcd1..946811c37 100644 --- a/src/SIL.Machine/FiniteState/DeterministicFstTraversalMethod.cs +++ b/src/SIL.Machine/FiniteState/DeterministicFstTraversalMethod.cs @@ -26,7 +26,8 @@ public override IEnumerable> Traverse( ref int annIndex, Register[,] initRegisters, IList initCmds, - ISet initAnns + ISet initAnns, + bool allMatches ) { Stack> instStack = InitializeStack( diff --git a/src/SIL.Machine/FiniteState/Fst.cs b/src/SIL.Machine/FiniteState/Fst.cs index 04fb681c4..b996b3255 100644 --- a/src/SIL.Machine/FiniteState/Fst.cs +++ b/src/SIL.Machine/FiniteState/Fst.cs @@ -387,7 +387,7 @@ out IEnumerable> results } List> curResults = traversalMethod - .Traverse(ref annIndex, initRegisters, cmds, initAnns) + .Traverse(ref annIndex, initRegisters, cmds, initAnns, allMatches) .ToList(); if (curResults.Count > 0) { @@ -424,7 +424,7 @@ private int ResultCompare(FstResult x, FstResult compare = -compare; if (IsDeterministic) { - compare = x.IsLazy ? -compare : compare; + compare = (x.IsLazy || y.IsLazy) ? -compare : compare; } else if (compare == 0) { diff --git a/src/SIL.Machine/FiniteState/ITraversalMethod.cs b/src/SIL.Machine/FiniteState/ITraversalMethod.cs index d11dfab46..590e636b3 100644 --- a/src/SIL.Machine/FiniteState/ITraversalMethod.cs +++ b/src/SIL.Machine/FiniteState/ITraversalMethod.cs @@ -11,7 +11,8 @@ IEnumerable> Traverse( ref int annIndex, Register[,] initRegisters, IList initCmds, - ISet initAnns + ISet initAnns, + bool allMatches ); } } diff --git a/src/SIL.Machine/FiniteState/NondeterministicFsaTraversalMethod.cs b/src/SIL.Machine/FiniteState/NondeterministicFsaTraversalMethod.cs index b5d3b3d5e..f1488c802 100644 --- a/src/SIL.Machine/FiniteState/NondeterministicFsaTraversalMethod.cs +++ b/src/SIL.Machine/FiniteState/NondeterministicFsaTraversalMethod.cs @@ -24,7 +24,8 @@ public override IEnumerable> Traverse( ref int annIndex, Register[,] initRegisters, IList initCmds, - ISet initAnns + ISet initAnns, + bool allMatches ) { Stack> instStack = InitializeStack( @@ -35,6 +36,7 @@ ISet initAnns ); var curResults = new List>(); + var lattice = !allMatches ? CreateFstLattice() : null; var traversed = new HashSet, int, Register[,]>>( AnonymousEqualityComparer.Create, int, Register[,]>>( KeyEquals, @@ -44,6 +46,9 @@ ISet initAnns while (instStack.Count != 0) { NondeterministicFsaTraversalInstance inst = instStack.Pop(); + NondeterministicFsaTraversalInstance origInst = !allMatches + ? CopyInstanceAndBindings(inst) + : null; bool releaseInstance = true; VariableBindings varBindings = null; @@ -69,24 +74,36 @@ ISet initAnns } ti.Visited.Add(arc.Target); + int resultCount = !allMatches ? curResults.Count : 0; NondeterministicFsaTraversalInstance newInst = EpsilonAdvance( ti, arc, curResults ); - Tuple, int, Register[,]> key = Tuple.Create( - newInst.State, - newInst.AnnotationIndex, - newInst.Registers - ); - if (!traversed.Contains(key)) + bool skip = false; + if (!allMatches) { - instStack.Push(newInst); - traversed.Add(key); + if (curResults.Count > resultCount) + RecordFinalArc(lattice, origInst, arc); + if (RecordedInstance(lattice, newInst, origInst, arc)) + skip = true; + } + if (!skip) + { + Tuple, int, Register[,]> key = Tuple.Create( + newInst.State, + newInst.AnnotationIndex, + newInst.Registers + ); + if (!traversed.Contains(key)) + { + instStack.Push(newInst); + traversed.Add(key); + } + if (isInstReusable) + releaseInstance = false; + varBindings = null; } - if (isInstReusable) - releaseInstance = false; - varBindings = null; } } else @@ -99,6 +116,7 @@ ISet initAnns ? inst : CopyInstance(inst); + int resultCount = !allMatches ? curResults.Count : 0; foreach ( NondeterministicFsaTraversalInstance newInst in Advance( ti, @@ -109,6 +127,13 @@ NondeterministicFsaTraversalInstance newInst in Advance( ) { newInst.Visited.Clear(); + if (!allMatches) + { + if (curResults.Count > resultCount) + RecordFinalArc(lattice, origInst, arc); + if (RecordedInstance(lattice, newInst, origInst, arc)) + continue; + } Tuple, int, Register[,]> key = Tuple.Create( newInst.State, newInst.AnnotationIndex, @@ -132,6 +157,12 @@ NondeterministicFsaTraversalInstance newInst in Advance( ReleaseInstance(inst); } + if (!allMatches) + { + var newResults = ExtractResults(lattice, allMatches); + curResults = newResults; + } + CheckAcceptingStartState(initAnns, initRegisters, curResults); return curResults; diff --git a/src/SIL.Machine/FiniteState/NondeterministicFstTraversalMethod.cs b/src/SIL.Machine/FiniteState/NondeterministicFstTraversalMethod.cs index e171f4410..aeb45bab9 100644 --- a/src/SIL.Machine/FiniteState/NondeterministicFstTraversalMethod.cs +++ b/src/SIL.Machine/FiniteState/NondeterministicFstTraversalMethod.cs @@ -27,7 +27,8 @@ public override IEnumerable> Traverse( ref int annIndex, Register[,] initRegisters, IList initCmds, - ISet initAnns + ISet initAnns, + bool allMatches ) { Stack> instStack = InitializeStack( diff --git a/src/SIL.Machine/FiniteState/TraversalMethodBase.cs b/src/SIL.Machine/FiniteState/TraversalMethodBase.cs index c5934d991..bee254144 100644 --- a/src/SIL.Machine/FiniteState/TraversalMethodBase.cs +++ b/src/SIL.Machine/FiniteState/TraversalMethodBase.cs @@ -1,6 +1,7 @@ using System; using System.Collections.Generic; using System.Linq; +using SIL.Extensions; using SIL.Machine.Annotations; using SIL.Machine.DataStructures; using SIL.Machine.FeatureModel; @@ -91,7 +92,8 @@ public abstract IEnumerable> Traverse( ref int annIndex, Register[,] initRegisters, IList initCmds, - ISet initAnns + ISet initAnns, + bool allMatches ); protected static void ExecuteCommands( @@ -452,11 +454,256 @@ protected TInst CopyInstance(TInst inst) return ni; } + protected TInst CopyInstanceAndBindings(TInst inst) + { + TInst ni = CopyInstance(inst); + if (inst.VariableBindings != null) + { + ni.VariableBindings = inst.VariableBindings.Clone(); + } + return ni; + } + protected abstract TInst CreateInstance(); protected void ReleaseInstance(TInst inst) { _cachedInstances.Enqueue(inst); } + + protected class LatticeNode + { + public State State { get; set; } + public int AnnotationIndex { get; set; } + public VariableBindings VariableBindings { get; set; } + } + + protected class LatticeArc + { + public TInst Instance { get; set; } + public Arc Arc { get; set; } + public IList Instances { get; set; } + public bool Visited { get; set; } + } + + private readonly LatticeNode _finalState = new LatticeNode(); + + /// + /// Creates a lattice. + /// A lattice is a graph that represents the space of traversals as a packed forest. + /// The nodes are [State, AnnotationIndex] pairs. + /// Each node has a list of incoming arcs that are [Instance, Arc] pairs. + /// The Instance encodes the previous node. + /// + protected IDictionary> CreateFstLattice() + { + return new Dictionary>( + AnonymousEqualityComparer.Create(LatticeNodeKeyEquals, LatticeNodeKeyGetHashCode) + ); + } + + /// + /// Check whether instance is already recorded in lattice. + /// If not, adds instance to lattice. + /// Also adds [origInstance, arc] to instance's incoming arcs. + /// + protected bool RecordedInstance( + IDictionary> lattice, + TInst instance, + TInst origInstance, + Arc arc + ) + { + var nodeKey = new LatticeNode() + { + State = instance.State, + AnnotationIndex = instance.AnnotationIndex, + VariableBindings = instance.VariableBindings?.Clone(), + }; + bool recorded = lattice.TryGetValue(nodeKey, out IList incoming); + if (!recorded) + { + // Add nodeKey to lattice. + incoming = new List(); + lattice[nodeKey] = incoming; + } + // Add [origInstance, arc] to incoming. + incoming.Add(new LatticeArc() { Instance = origInstance, Arc = arc }); + return recorded; + } + + protected void RecordFinalArc( + IDictionary> lattice, + TInst origInstance, + Arc arc + ) + { + bool recorded = lattice.TryGetValue(_finalState, out IList incoming); + if (!recorded) + { + // Add _finalState to lattice. + incoming = new List(); + lattice[_finalState] = incoming; + } + // Add [origInstance, arc] to incoming. + incoming.Add(new LatticeArc() { Instance = origInstance, Arc = arc }); + } + + /// + /// Extract the results encoded in lattice under the final state. + /// + protected List> ExtractResults( + IDictionary> lattice, + bool allMatches + ) + { + List> newResults = new List>(); + IList incoming; + if (!lattice.TryGetValue(_finalState, out incoming)) + return newResults; + foreach (LatticeArc latticeArc in incoming) + { + foreach (TInst instance in ExpandArcInstances(latticeArc, lattice, allMatches)) + { + AdvanceInstance(instance, latticeArc.Arc, null, newResults, null, 0); + } + } + return newResults; + } + + private IList ExpandArcInstances( + LatticeArc latticeArc, + IDictionary> lattice, + bool allMatches + ) + { + if (latticeArc.Visited) + return new List(); + if (latticeArc.Instances == null) + { + try + { + latticeArc.Visited = true; + latticeArc.Instances = ExpandInstances(latticeArc.Instance, lattice, allMatches); + } + finally + { + latticeArc.Visited = false; + } + } + IList instances = new List(); + foreach (TInst instance in latticeArc.Instances) + { + instances.Add(CopyInstanceAndBindings(instance)); + } + return instances; + } + + private IList ExpandInstances( + TInst instance, + IDictionary> lattice, + bool allMatches + ) + { + IList instances = new List(); + IList> curResults = new List>(); + LatticeNode nodeKey = new LatticeNode() + { + State = instance.State, + AnnotationIndex = instance.AnnotationIndex, + VariableBindings = instance.VariableBindings?.Clone(), + }; + bool recorded = lattice.TryGetValue(nodeKey, out IList incoming); + if (!recorded) + { + // The starting instance. + instances.Add(CopyInstanceAndBindings(instance)); + return instances; + } + foreach (LatticeArc latticeArc in incoming) + { + foreach (TInst source in ExpandArcInstances(latticeArc, lattice, allMatches)) + { + AdvanceInstance( + source, + latticeArc.Arc, + instances, + curResults, + instance.State, + instance.AnnotationIndex + ); + } + } + if (!allMatches && instances.Count > 1) + { + instances.Sort(InstanceCompare); + TInst first = instances.First(); + instances.Clear(); + instances.Add(first); + } + return instances; + } + + private void AdvanceInstance( + TInst instance, + Arc arc, + IList instances, + IList> curResults, + State state, + int annotationIndex + ) + { + if (arc.Input.IsEpsilon) + { + TInst ni = EpsilonAdvance(instance, arc, curResults); + if (instances != null && ni.State == state && ni.AnnotationIndex == annotationIndex) + { + instances.Add(ni); + } + } + else if (CheckInputMatch(arc, instance.AnnotationIndex, instance.VariableBindings)) + { + foreach (TInst ni in Advance(instance, instance.VariableBindings, arc, curResults)) + { + if (instances != null && ni.State == state && ni.AnnotationIndex == annotationIndex) + instances.Add(ni); + } + } + } + + private int InstanceCompare(TInst x, TInst y) + { + int compare = 0; + if (x.Priorities != null) + { + foreach (Tuple priorityPair in x.Priorities.Zip(y.Priorities)) + { + compare = priorityPair.Item1.CompareTo(priorityPair.Item2); + if (compare != 0) + break; + } + } + return compare; + } + + private bool LatticeNodeKeyEquals(LatticeNode x, LatticeNode y) + { + if (x.State == null || y.State == null) + return x.State == y.State; + if (x.VariableBindings == null) + return x.State.Equals(y.State) && x.AnnotationIndex.Equals(y.AnnotationIndex); + return x.State.Equals(y.State) + && x.AnnotationIndex.Equals(y.AnnotationIndex) + && x.VariableBindings.Equals(y.VariableBindings); + } + + private int LatticeNodeKeyGetHashCode(LatticeNode m) + { + int code = 23; + code = code * 31 + (m.State != null ? m.State.GetHashCode() : 0); + code = code * 31 + m.AnnotationIndex.GetHashCode(); + code = code * 31 + (m.VariableBindings != null ? m.VariableBindings.GetHashCode() : 0); + return code; + } } } diff --git a/tests/SIL.Machine.Tests/Matching/TraversalDedupMinimalCasesTests.cs b/tests/SIL.Machine.Tests/Matching/TraversalDedupMinimalCasesTests.cs new file mode 100644 index 000000000..74e96272d --- /dev/null +++ b/tests/SIL.Machine.Tests/Matching/TraversalDedupMinimalCasesTests.cs @@ -0,0 +1,126 @@ +using NUnit.Framework; +using SIL.Machine.Annotations; +using SIL.Machine.DataStructures; +using SIL.Machine.FeatureModel; + +namespace SIL.Machine.Matching; + +// Minimal, hand-built reproductions of two cases found by a differential fuzz +// (TraversalDedupDifferentialFuzzTests) comparing master against the +// add-allMatches-to-Traverse branch's traversal-dedup change. Both settings +// use AllSubmatches = false and Nondeterministic = false, matching the fuzz. +public class TraversalDedupMinimalCasesTests : PhoneticTestsBase +{ + private FeatureStruct Ann(string voice, string high, string back) + { + return FeatureStruct + .New(PhoneticFeatSys) + .Feature("voice") + .EqualTo(voice) + .Feature("high") + .EqualTo(high) + .Feature("back") + .EqualTo(back) + .Value; + } + + [Test] + public void NondeterministicTraversal_DedupOnVariableBindingLosesMatch() + { + // Pattern: high=$v0+, anchored to both ends (runs NondeterministicFsaTraversalMethod + // because of the variable). The only way to cover the whole input [0,5) with a + // single consistent value of v0 is the run of high- annotations: [0,2)+[2,4)+[4,5). + // Reaching it requires abandoning, at annotation index 0, the parallel instance that + // consumed the high+ annotation [0,1) (which binds v0=+ but then dead-ends, since no + // annotation starts at offset 1). Both instances reach the same (State, AnnotationIndex) + // with different VariableBindings (v0=+ vs v0=-); the traversal-dedup change keys on + // (State, AnnotationIndex) alone and can keep the v0=+ instance, which can never + // complete the anchored match, discarding the one that would have succeeded. + Pattern pattern = Pattern + .New() + .Annotation(FeatureStruct.New(PhoneticFeatSys).Feature("high").EqualToVariable("v0").Value) + .OneOrMore.Value; + + var data = new AnnotatedStringData(new string('a', 5)); + data.Annotations.Add(0, 2, Ann("voice-", "high-", "back-"), false); + data.Annotations.Add(0, 2, Ann("voice-", "high-", "back-"), false); // duplicate span+values, as the fuzz produced + data.Annotations.Add(0, 1, Ann("voice-", "high+", "back+"), false); + data.Annotations.Add(2, 4, Ann("voice-", "high-", "back+"), false); + data.Annotations.Add(4, 5, Ann("voice+", "high-", "back+"), false); + + var matcher = new Matcher( + pattern, + new MatcherSettings + { + AnchoredToStart = true, + AnchoredToEnd = true, + Direction = Direction.LeftToRight, + AllSubmatches = false, + Nondeterministic = false, + } + ); + + Match match = matcher.Match(data); + + Assert.That(match.Success, Is.True, $"success={match.Success};range={DescribeRange(match.Range)}"); + Assert.That(match.Range, Is.EqualTo(Range.Create(0, 5))); + Assert.That(((FeatureSymbol)match.VariableBindings["v0"]).ID, Is.EqualTo("high-")); + } + + [Test] + public void DeterministicTraversal_DedupOnRegistersShortensMatch() + { + // This verifies that the bug where ResultCompare is asymmetric has been fixed. + // When a pattern gets determinized and there is an alternative, the shortest path is preferred. + Pattern pattern = Pattern + .New() + .Group("g0", g0 => g0.Annotation(FeatureStruct.New(PhoneticFeatSys).Feature("back").EqualTo("back+").Value)) + .Or.Group( + "g1", + g1 => + g1.Annotation(FeatureStruct.New(PhoneticFeatSys).Feature("high").EqualTo("high+").Value) + .Annotation(FeatureStruct.New(PhoneticFeatSys).Feature("back").EqualTo("back+").Value) + ) + .Value; + + var data = new AnnotatedStringData(new string('a', 4)); + data.Annotations.Add(0, 1, Ann("voice-", "high+", "back+"), false); + data.Annotations.Add(0, 2, Ann("voice-", "high+", "back+"), false); + data.Annotations.Add(2, 4, Ann("voice+", "high-", "back+"), false); + + var matcher = new Matcher( + pattern, + new MatcherSettings + { + AnchoredToStart = true, + AnchoredToEnd = false, + Direction = Direction.LeftToRight, + AllSubmatches = false, + Nondeterministic = false, + } + ); + + Match match = matcher.Match(data); + + Assert.That(match.Success, Is.True); + Assert.That( + match.Range, + Is.EqualTo(Range.Create(0, 1)), + $"success={match.Success};range={DescribeRange(match.Range)};" + + $"g0={DescribeCapture(match.GroupCaptures["g0"])};g1={DescribeCapture(match.GroupCaptures["g1"])}" + ); + Assert.That(match.GroupCaptures["g0"].Success, Is.True); + Assert.That(match.GroupCaptures["g1"].Success, Is.False); + Assert.That(match.GroupCaptures["g1"].Range, Is.EqualTo(Range.Create(-1, -1))); + } + + private static string DescribeRange(Range range) + { + return range == Range.Null ? "" : $"[{range.Start},{range.End})"; + } + + private static string DescribeCapture(GroupCapture capture) + { + return capture.Success ? DescribeRange(capture.Range) : ""; + } +}