using FluentAssertions; using PiiRedaction.Core.Detection; using PiiRedaction.Core.Models; using PiiRedaction.Infrastructure.Onnx; namespace PiiRedaction.Infrastructure.Tests.Onnx; [TestFixture] public sealed class RoutingOnnxNerModelRunnerTests { [Test] public void PredictEntities_LatinOnly_UsesEnglishRunnerOnly() { var english = new FakeLanguageNerRunner("Ravi Kumar", NerModelOrigin.English); var tamil = new FakeLanguageNerRunner("தமிழ் பெயர்", NerModelOrigin.Tamil); var router = new RoutingOnnxNerModelRunner(english, tamil, enableTamilNer: true); var result = router.PredictEntities("Customer Ravi Kumar called."); result.Entities.Should().ContainSingle(entity => entity.Value == "Ravi Kumar" && entity.ModelOrigin == NerModelOrigin.English); result.InvokedModels.Should().Equal(NerModelOrigin.English); english.CallCount.Should().Be(1); tamil.CallCount.Should().Be(0); } [Test] public void PredictEntities_TamilOnly_UsesTamilRunnerOnly() { var english = new FakeLanguageNerRunner("Ravi Kumar", NerModelOrigin.English); var tamil = new FakeLanguageNerRunner("ராஜேஷ்", NerModelOrigin.Tamil); var router = new RoutingOnnxNerModelRunner(english, tamil, enableTamilNer: true); var result = router.PredictEntities("வாடிக்கையாளர் ராஜேஷ்"); result.Entities.Should().ContainSingle(entity => entity.Value == "ராஜேஷ்" && entity.ModelOrigin == NerModelOrigin.Tamil); result.InvokedModels.Should().Equal(NerModelOrigin.Tamil); english.CallCount.Should().Be(0); tamil.CallCount.Should().Be(1); } [Test] public void PredictEntities_Mixed_InvokesBothRunnersAndTagsOrigins() { var english = new FakeLanguageNerRunner("Priya", NerModelOrigin.English); var tamil = new FakeLanguageNerRunner("மற்றும்", NerModelOrigin.Tamil); var router = new RoutingOnnxNerModelRunner(english, tamil, enableTamilNer: true); var result = router.PredictEntities("Rajesh மற்றும் Priya"); english.CallCount.Should().Be(1); tamil.CallCount.Should().Be(1); result.InvokedModels.Should().Equal(NerModelOrigin.English, NerModelOrigin.Tamil); result.Entities.Should().Contain(entity => entity.ModelOrigin == NerModelOrigin.English); result.Entities.Should().Contain(entity => entity.ModelOrigin == NerModelOrigin.Tamil); } [Test] public void PredictEntities_NoLetters_InvokesNeither() { var english = new FakeLanguageNerRunner("ignored", NerModelOrigin.English); var tamil = new FakeLanguageNerRunner("ignored", NerModelOrigin.Tamil); var router = new RoutingOnnxNerModelRunner(english, tamil, enableTamilNer: true); var result = router.PredictEntities("9876543210"); result.Entities.Should().BeEmpty(); result.InvokedModels.Should().BeEmpty(); english.CallCount.Should().Be(0); tamil.CallCount.Should().Be(0); } [Test] public void PredictEntities_TamilDisabled_SkipsTamilRunnerForMixedText() { var english = new FakeLanguageNerRunner("EnglishName", NerModelOrigin.English); var tamil = new FakeLanguageNerRunner("தமிழ்", NerModelOrigin.Tamil); var router = new RoutingOnnxNerModelRunner(english, tamil, enableTamilNer: false); var result = router.PredictEntities("Rajesh மற்றும் Priya"); english.CallCount.Should().Be(1); tamil.CallCount.Should().Be(0); result.InvokedModels.Should().Equal(NerModelOrigin.English); } [Test] public void MergePersonSpans_PrefersLongerOverlappingSpan() { var entities = new[] { CreatePerson("Raj", 0, 3, NerModelOrigin.English), CreatePerson("Rajesh", 0, 6, NerModelOrigin.Tamil) }; var merged = RoutingOnnxNerModelRunner.MergePersonSpans(entities); merged.Should().ContainSingle(entity => entity.Value == "Rajesh" && entity.ModelOrigin == NerModelOrigin.Tamil); } private static PiiEntity CreatePerson(string value, int start, int length, NerModelOrigin origin) => new(PiiEntityType.Person, value, start, length, PiiDetectionSource.Ner, ModelOrigin: origin); private sealed class FakeLanguageNerRunner : IOnnxNerModelRunner { private readonly string _personValue; private readonly NerModelOrigin _origin; public FakeLanguageNerRunner(string personValue, NerModelOrigin origin) { _personValue = personValue; _origin = origin; } public int CallCount { get; private set; } public bool IsModelAvailable => true; public NerPredictionResult PredictEntities(string text) { CallCount++; var start = text.IndexOf(_personValue, StringComparison.Ordinal); if (start < 0) { start = 0; } return new NerPredictionResult( [ new PiiEntity( PiiEntityType.Person, _personValue, start, _personValue.Length, PiiDetectionSource.Ner, ModelOrigin: _origin) ], [_origin]); } } }