diff --git a/internal/modules/dnd/shared/citations.go b/internal/modules/dnd/shared/citations.go index 97e80e6..f8b129f 100644 --- a/internal/modules/dnd/shared/citations.go +++ b/internal/modules/dnd/shared/citations.go @@ -7,27 +7,37 @@ import ( "gitea.maximumdirect.net/eric/notarius/internal/core/source" ) +// CitationResolver resolves cited ranges against one source document. +type CitationResolver struct { + doc *source.SourceDocument + index source.DocumentIndex +} + +// NewCitationResolver prepares citation resolution for doc. +func NewCitationResolver(doc *source.SourceDocument) (*CitationResolver, error) { + if doc == nil { + return nil, fmt.Errorf("source document must not be nil") + } + return &CitationResolver{doc: doc, index: source.NewDocumentIndex(doc)}, nil +} + // CitedText resolves cited source ranges in document order, including each // source unit once, and joins the resulting text with newlines. -func CitedText(doc *source.SourceDocument, refs []source.SourceRef) (string, error) { - if doc == nil { - return "", fmt.Errorf("source document must not be nil") - } - - included := make([]bool, len(doc.Units)) +func (r *CitationResolver) CitedText(refs []source.SourceRef) (string, error) { + included := make([]bool, len(r.doc.Units)) for _, ref := range refs { - if err := source.ValidateRef(doc, ref); err != nil { + if err := r.index.ValidateRef(ref); err != nil { return "", fmt.Errorf("resolve cited source range: %w", err) } - start, _ := source.UnitIndex(doc, ref.StartUnitID) - end, _ := source.UnitIndex(doc, ref.EndUnitID) + start, _ := r.index.Position(ref.StartUnitID) + end, _ := r.index.Position(ref.EndUnitID) for index := start; index <= end; index++ { included[index] = true } } - parts := make([]string, 0, len(doc.Units)) - for index, unit := range doc.Units { + parts := make([]string, 0, len(r.doc.Units)) + for index, unit := range r.doc.Units { if included[index] { parts = append(parts, unit.Text) } diff --git a/internal/modules/dnd/shared/citations_test.go b/internal/modules/dnd/shared/citations_test.go index cf41b6a..7a090de 100644 --- a/internal/modules/dnd/shared/citations_test.go +++ b/internal/modules/dnd/shared/citations_test.go @@ -7,8 +7,12 @@ import ( "gitea.maximumdirect.net/eric/notarius/internal/core/source" ) -func TestCitedText(t *testing.T) { +func TestCitationResolver(t *testing.T) { doc := citationDocument() + resolver, err := NewCitationResolver(doc) + if err != nil { + t.Fatalf("NewCitationResolver() error = %v", err) + } tests := []struct { name string refs []source.SourceRef @@ -51,7 +55,7 @@ func TestCitedText(t *testing.T) { t.Run(test.name, func(t *testing.T) { beforeUnits := append([]source.SourceUnit(nil), doc.Units...) refs := append([]source.SourceRef(nil), test.refs...) - got, err := CitedText(doc, refs) + got, err := resolver.CitedText(refs) if (err != nil) != test.wantErr { t.Fatalf("CitedText() error = %v, want error = %t", err, test.wantErr) } @@ -65,9 +69,9 @@ func TestCitedText(t *testing.T) { } } -func TestCitedTextRejectsNilDocument(t *testing.T) { - if _, err := CitedText(nil, nil); err == nil { - t.Fatal("CitedText() error = nil, want nil-document error") +func TestNewCitationResolverRejectsNilDocument(t *testing.T) { + if _, err := NewCitationResolver(nil); err == nil { + t.Fatal("NewCitationResolver() error = nil, want nil-document error") } } diff --git a/internal/modules/dnd/validate/combatturns/source_relatedness/validator.go b/internal/modules/dnd/validate/combatturns/source_relatedness/validator.go index e9848a9..8f4a461 100644 --- a/internal/modules/dnd/validate/combatturns/source_relatedness/validator.go +++ b/internal/modules/dnd/validate/combatturns/source_relatedness/validator.go @@ -37,10 +37,14 @@ func (v *Validator) Validate(_ context.Context, req contracts.TypedValidationReq if combatshape.Validate(req.Value) != nil { return contracts.ValidationResult{Approved: true}, nil } + resolver, err := shared.NewCitationResolver(req.Source) + if err != nil { + return contracts.ValidationResult{Approved: true}, nil + } citedTexts := make([]string, len(req.Value.CombatTurns)) for turnIndex, turn := range req.Value.CombatTurns { - citedText, err := shared.CitedText(req.Source, turn.SourceRefs) + citedText, err := resolver.CitedText(turn.SourceRefs) if err != nil { return contracts.ValidationResult{Approved: true}, nil } diff --git a/internal/modules/dnd/validate/npcinteractions/source_relatedness/validator.go b/internal/modules/dnd/validate/npcinteractions/source_relatedness/validator.go index 02d6fdf..a5b8465 100644 --- a/internal/modules/dnd/validate/npcinteractions/source_relatedness/validator.go +++ b/internal/modules/dnd/validate/npcinteractions/source_relatedness/validator.go @@ -39,9 +39,13 @@ func (v *Validator) Validate(_ context.Context, req contracts.TypedValidationReq if interactionshape.Validate(req.Value) != nil { return contracts.ValidationResult{Approved: true}, nil } + resolver, err := shared.NewCitationResolver(req.Source) + if err != nil { + return contracts.ValidationResult{Approved: true}, nil + } citedTexts := make([]string, len(req.Value.Interactions)) for index, interaction := range req.Value.Interactions { - citedText, err := shared.CitedText(req.Source, interaction.SourceRefs) + citedText, err := resolver.CitedText(interaction.SourceRefs) if err != nil { return contracts.ValidationResult{Approved: true}, nil } diff --git a/internal/modules/dnd/validate/npcs/source_relatedness/validator.go b/internal/modules/dnd/validate/npcs/source_relatedness/validator.go index 339d195..6ceb258 100644 --- a/internal/modules/dnd/validate/npcs/source_relatedness/validator.go +++ b/internal/modules/dnd/validate/npcs/source_relatedness/validator.go @@ -37,9 +37,13 @@ func (v *Validator) Validate(_ context.Context, req contracts.TypedValidationReq if err := npcshape.Validate(req.Value); err != nil { return contracts.ValidationResult{Approved: true}, nil } + resolver, err := shared.NewCitationResolver(req.Source) + if err != nil { + return contracts.ValidationResult{Approved: true}, nil + } citedTexts := make([]string, len(req.Value.NPCs)) for npcIndex, npc := range req.Value.NPCs { - citedText, err := shared.CitedText(req.Source, npc.SourceRefs) + citedText, err := resolver.CitedText(npc.SourceRefs) if err != nil { return contracts.ValidationResult{Approved: true}, nil } diff --git a/internal/modules/dnd/validate/scenedescriptions/source_relatedness/validator.go b/internal/modules/dnd/validate/scenedescriptions/source_relatedness/validator.go index e2d23ac..29b28b9 100644 --- a/internal/modules/dnd/validate/scenedescriptions/source_relatedness/validator.go +++ b/internal/modules/dnd/validate/scenedescriptions/source_relatedness/validator.go @@ -48,9 +48,13 @@ func (v *Validator) Validate(_ context.Context, req contracts.TypedValidationReq if shape.Validate(req.Value) != nil || !sourceRefsValid(req.Source, req.Value) { return contracts.ValidationResult{Approved: true}, nil } + resolver, err := shared.NewCitationResolver(req.Source) + if err != nil { + return contracts.ValidationResult{Approved: true}, nil + } warnings := make([]contracts.Warning, 0) for index, scene := range req.Value.Scenes { - citedText, err := shared.CitedText(req.Source, []source.SourceRef{scene.SourceRef}) + citedText, err := resolver.CitedText([]source.SourceRef{scene.SourceRef}) if err != nil { return contracts.ValidationResult{Approved: true}, nil } diff --git a/internal/modules/dnd/validate/spells/source_relatedness/validator.go b/internal/modules/dnd/validate/spells/source_relatedness/validator.go index aafa0d9..96eded7 100644 --- a/internal/modules/dnd/validate/spells/source_relatedness/validator.go +++ b/internal/modules/dnd/validate/spells/source_relatedness/validator.go @@ -36,9 +36,13 @@ func (v *Validator) Validate(_ context.Context, req contracts.TypedValidationReq if err := spellshape.Validate(req.Value); err != nil { return contracts.ValidationResult{Approved: true}, nil } + resolver, err := shared.NewCitationResolver(req.Source) + if err != nil { + return contracts.ValidationResult{Approved: true}, nil + } citedTexts := make([]string, len(req.Value.SpellCasts)) for spellIndex, spell := range req.Value.SpellCasts { - citedText, err := shared.CitedText(req.Source, spell.SourceRefs) + citedText, err := resolver.CitedText(spell.SourceRefs) if err != nil { return contracts.ValidationResult{Approved: true}, nil }