Files
notarius/internal/modules/dnd/normalize/itemregistry/normalizer.go

358 lines
15 KiB
Go

// Package itemregistry normalizes merged D&D item-registry candidates conservatively.
package itemregistry
import (
"context"
"fmt"
"reflect"
"sort"
"strconv"
"strings"
"gitea.maximumdirect.net/eric/notarius/internal/core/source"
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
"gitea.maximumdirect.net/eric/notarius/internal/framework/pipeline"
"gitea.maximumdirect.net/eric/notarius/internal/framework/semanticreconcile"
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd"
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/items/identity"
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/shared"
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/shared/diagnostics"
)
const (
Key = "dnd/item-registry"
PromptID = "dnd.item_registry.normalize"
PromptVersion = "v1"
normalizationPolicy = "dnd.item_registry.normalize.v3"
NormalizationPolicy = normalizationPolicy
ReasonCodeItemFieldsNormalized = "item_fields_normalized"
ReasonCodeItemIDRecomputed = "item_id_recomputed"
ReasonCodeSourceReferencesNormalized = "source_references_normalized"
ReasonCodeDuplicateItemCollapsed = "duplicate_item_collapsed"
ReasonCodeItemSemanticProposalInvalid = "item_semantic_proposal_invalid"
ReasonCodeItemSemanticReconciliationExhausted = "item_semantic_reconciliation_exhausted"
ReasonCodeItemNormalizationWarningsOmitted = "item_normalization_warnings_omitted"
)
var requiredCapabilities = []string{"merged"}
var providedCapabilities = []string{"normalized"}
var _ contracts.Normalizer[dnd.ItemRegistry] = (*Normalizer)(nil)
var _ contracts.ManifestMetadataProvider = (*Normalizer)(nil)
var _ pipeline.CheckpointFingerprintProvider = (*Normalizer)(nil)
type Options struct{}
type Normalizer struct {
engine *semanticreconcile.Engine
}
func New(llmClient contracts.StructuredLLMClient, _ Options) (*Normalizer, error) {
if llmClient == nil {
return nil, normalizerErrorf("LLM client must not be nil")
}
promptSHA, err := promptAssetMetadata()
if err != nil {
return nil, normalizerErrorf("load prompt metadata: %w", err)
}
engine, err := semanticreconcile.NewEngine(llmClient, semanticreconcile.PromptSpec{
ID: PromptID, Version: PromptVersion, SHA256: promptSHA,
}, semanticreconcile.DefaultLimits())
if err != nil {
return nil, normalizerErrorf("construct semantic reconciliation engine: %w", err)
}
return &Normalizer{engine: engine}, nil
}
func (n *Normalizer) Key() string { return Key }
func (n *Normalizer) ReferenceSlots() []contracts.ReferenceSlot { return nil }
func (n *Normalizer) ManifestMetadata() map[string]any {
if n == nil || n.engine == nil {
return nil
}
metadata := n.engine.ManifestMetadata()
metadata["identity_policy"] = identity.Policy
metadata["normalization_policy"] = normalizationPolicy
return metadata
}
func (n *Normalizer) CheckpointFingerprints() []pipeline.CheckpointFingerprint {
if n == nil || n.engine == nil {
return nil
}
fingerprints := n.engine.CheckpointFingerprints()
return append(fingerprints,
pipeline.CheckpointFingerprint{Name: "identity_policy", Value: identity.Policy},
pipeline.CheckpointFingerprint{Name: "normalization_policy", Value: normalizationPolicy},
)
}
func (n *Normalizer) Normalize(ctx context.Context, req contracts.TypedNormalizeRequest[dnd.ItemRegistry]) (contracts.TypedNormalizeResult[dnd.ItemRegistry], error) {
if n == nil {
return contracts.TypedNormalizeResult[dnd.ItemRegistry]{}, normalizerErrorf("normalizer must not be nil")
}
if n.engine == nil {
return contracts.TypedNormalizeResult[dnd.ItemRegistry]{}, normalizerErrorf("semantic reconciliation engine must not be nil")
}
if ctx == nil {
return contracts.TypedNormalizeResult[dnd.ItemRegistry]{}, normalizerErrorf("context must not be nil")
}
if err := ctx.Err(); err != nil {
return contracts.TypedNormalizeResult[dnd.ItemRegistry]{}, normalizerErrorf("context error before normalize: %w", err)
}
if req.Source == nil {
return contracts.TypedNormalizeResult[dnd.ItemRegistry]{}, normalizerErrorf("source document must not be nil")
}
order := shared.NewSourceRefOrder(req.Source)
records, warnings := preprocessRecords(req.MergeOutput.Value, order)
deterministic := recordList(records)
if len(records) < 2 {
return contracts.TypedNormalizeResult[dnd.ItemRegistry]{Value: deterministic, Warnings: limitWarnings(warnings)}, nil
}
candidates, envelopes, err := reconciliationInputs(records)
if err != nil {
return contracts.TypedNormalizeResult[dnd.ItemRegistry]{}, normalizerErrorf("prepare semantic reconciliation inputs: %w", err)
}
reconciliation, err := n.engine.Reconcile(ctx, semanticreconcile.Request{
StageName: Key, Source: req.Source, Candidates: candidates,
ProfileID: req.LLMProfile, SessionID: req.SessionID,
})
if err != nil {
return contracts.TypedNormalizeResult[dnd.ItemRegistry]{}, normalizerErrorf("reconcile semantic duplicates: %w", err)
}
switch reconciliation.Disposition() {
case semanticreconcile.SkippedInsufficientCandidates:
return contracts.TypedNormalizeResult[dnd.ItemRegistry]{Value: deterministic, Warnings: limitWarnings(warnings)}, nil
case semanticreconcile.SkippedLimitExceeded:
return contracts.TypedNormalizeResult[dnd.ItemRegistry]{Value: deterministic, Warnings: limitWarningsWithSemanticFallback(warnings)}, nil
case semanticreconcile.RetryableInvalidStructuredOutput:
return n.invalidStructuredResult(deterministic, warnings), nil
case semanticreconcile.Complete, semanticreconcile.RetryableDiscardedProposalGroups:
default:
return contracts.TypedNormalizeResult[dnd.ItemRegistry]{}, normalizerErrorf("unknown semantic reconciliation disposition %d", reconciliation.Disposition())
}
applied, semanticWarnings, rejectedGroups, err := applyReconciliationPlan(reconciliation.Plan(), records, envelopes, order)
if err != nil {
return contracts.TypedNormalizeResult[dnd.ItemRegistry]{}, normalizerErrorf("apply semantic reconciliation plan: %w", err)
}
warnings = append(warnings, semanticWarnings...)
discardedGroups := reconciliation.DiscardedGroupCount() + rejectedGroups
if discardedGroups == 0 {
return contracts.TypedNormalizeResult[dnd.ItemRegistry]{Value: recordList(applied), Warnings: limitWarnings(warnings)}, nil
}
return retryResult(recordList(applied), warnings, reconciliation, rejectedGroups), nil
}
func (n *Normalizer) invalidStructuredResult(value dnd.ItemRegistry, warnings []contracts.Warning) contracts.TypedNormalizeResult[dnd.ItemRegistry] {
return contracts.TypedNormalizeResult[dnd.ItemRegistry]{Value: value, Warnings: limitWarningsForRetry(warnings), Retry: &contracts.NormalizeRetry{
ReasonCode: ReasonCodeItemSemanticProposalInvalid, Message: "semantic proposal requires retry: invalid structured output",
FallbackWarnings: []contracts.Warning{semanticFallbackWarning(-1)},
}}
}
func retryResult(value dnd.ItemRegistry, warnings []contracts.Warning, reconciliation semanticreconcile.Result, rejectedGroups int) contracts.TypedNormalizeResult[dnd.ItemRegistry] {
details := semanticreconcile.IssueDetails(reconciliation.Issues())
if rejectedGroups > 0 {
details = append(details, "currency may only be consolidated with aliases of one denomination")
}
discardedGroups := reconciliation.DiscardedGroupCount() + rejectedGroups
return contracts.TypedNormalizeResult[dnd.ItemRegistry]{Value: value, Warnings: limitWarningsForRetry(warnings), Retry: &contracts.NormalizeRetry{
ReasonCode: ReasonCodeItemSemanticProposalInvalid, Message: diagnostics.Aggregate("semantic proposal requires retry", details),
FallbackWarnings: []contracts.Warning{semanticFallbackWarning(discardedGroups)},
}}
}
func semanticFallbackWarning(discarded int) contracts.Warning {
message := "semantic proposal could not be applied"
if discarded >= 0 {
message = fmt.Sprintf("%d proposal group(s) omitted after semantic proposal retry exhaustion", discarded)
}
return contracts.Warning{Scope: "items", ReasonCode: ReasonCodeItemSemanticReconciliationExhausted, Message: message}
}
func limitWarnings(warnings []contracts.Warning) []contracts.Warning {
return diagnostics.LimitWarnings(warnings, "items", ReasonCodeItemNormalizationWarningsOmitted)
}
func limitWarningsForRetry(warnings []contracts.Warning) []contracts.Warning {
if warnings == nil {
return nil
}
if len(warnings) < diagnostics.MaxWarnings {
return append([]contracts.Warning(nil), warnings...)
}
displayed := diagnostics.MaxWarnings - 2
bounded := append([]contracts.Warning(nil), warnings[:displayed]...)
return append(bounded, contracts.Warning{Scope: "items", ReasonCode: ReasonCodeItemNormalizationWarningsOmitted, Message: fmt.Sprintf("%d additional warning(s) omitted", len(warnings)-displayed)})
}
func limitWarningsWithSemanticFallback(warnings []contracts.Warning) []contracts.Warning {
return append(limitWarningsForRetry(warnings), semanticFallbackWarning(-1))
}
type normalizedRecord struct {
item dnd.Item
inputIndexes []int
earliest int
}
func preprocessRecords(input dnd.ItemRegistry, order shared.SourceRefOrder) ([]normalizedRecord, []contracts.Warning) {
if input.Items == nil {
return nil, nil
}
records := make([]normalizedRecord, len(input.Items))
warnings := make([]contracts.Warning, 0)
for index, inputItem := range input.Items {
item, fieldsChanged, refsChanged := normalizeRecord(inputItem, order)
records[index] = normalizedRecord{item: item, inputIndexes: []int{index}, earliest: index}
if fieldsChanged {
warnings = append(warnings, contracts.Warning{Scope: itemScope(index), ReasonCode: ReasonCodeItemFieldsNormalized, Message: fmt.Sprintf("input index %d: item name normalized for %s", index, diagnostics.Quote(inputItem.Name))})
}
if refsChanged {
warnings = append(warnings, contracts.Warning{Scope: itemScope(index), ReasonCode: ReasonCodeSourceReferencesNormalized, Message: fmt.Sprintf("input index %d: source references normalized (original count %d, final count %d)", index, len(inputItem.SourceRefs), len(item.SourceRefs))})
}
if inputItem.ID != item.ID {
warnings = append(warnings, contracts.Warning{Scope: itemScope(index), ReasonCode: ReasonCodeItemIDRecomputed, Message: fmt.Sprintf("input index %d: item ID recomputed from %s", index, diagnostics.Quote(item.Name))})
}
}
groups := comparisonNameGroups(records)
output := make([]normalizedRecord, 0, len(groups))
for _, members := range groups {
retained := cloneRecord(records[members[0]])
for _, member := range members[1:] {
retained.item.SourceRefs = append(retained.item.SourceRefs, records[member].item.SourceRefs...)
retained.inputIndexes = append(retained.inputIndexes, records[member].inputIndexes...)
}
retained.item.SourceRefs = order.Canonicalize(retained.item.SourceRefs)
retained.item.ID = identity.DeriveID(retained.item.Name)
retained.inputIndexes = sortedUniqueIndexes(retained.inputIndexes)
output = append(output, retained)
if len(members) > 1 {
warnings = append(warnings, duplicateWarning(retained.earliest, memberInputIndexes(records, members[1:])))
}
}
return output, warnings
}
func normalizeRecord(input dnd.Item, order shared.SourceRefOrder) (dnd.Item, bool, bool) {
output := cloneItem(input)
output.Name = identity.NormalizeDisplay(input.Name)
output.SourceRefs = order.Canonicalize(input.SourceRefs)
output.ID = identity.DeriveID(output.Name)
return output, input.Name != output.Name, !reflect.DeepEqual(input.SourceRefs, output.SourceRefs)
}
func comparisonNameGroups(records []normalizedRecord) [][]int {
groups := make([][]int, 0, len(records))
groupPositions := make(map[string]int, len(records))
for index, record := range records {
key := identity.ComparisonKey(record.item.Name)
if groupIndex, found := groupPositions[key]; found {
groups[groupIndex] = append(groups[groupIndex], index)
continue
}
groupPositions[key] = len(groups)
groups = append(groups, []int{index})
}
return groups
}
func cloneItem(input dnd.Item) dnd.Item {
input.SourceRefs = cloneSourceRefs(input.SourceRefs)
return input
}
func cloneRecord(input normalizedRecord) normalizedRecord {
input.item = cloneItem(input.item)
input.inputIndexes = append([]int(nil), input.inputIndexes...)
return input
}
func cloneSourceRefs(input []source.SourceRef) []source.SourceRef {
return append([]source.SourceRef(nil), input...)
}
func sortedUniqueIndexes(indexes []int) []int {
if len(indexes) == 0 {
return nil
}
out := append([]int(nil), indexes...)
sort.Ints(out)
write := 1
for _, index := range out[1:] {
if index != out[write-1] {
out[write] = index
write++
}
}
return out[:write]
}
func memberInputIndexes(records []normalizedRecord, members []int) []int {
indexes := make([]int, 0, len(members))
for _, member := range members {
indexes = append(indexes, records[member].inputIndexes...)
}
return sortedUniqueIndexes(indexes)
}
func recordValues(records []normalizedRecord) []dnd.Item {
if records == nil {
return nil
}
values := make([]dnd.Item, len(records))
for index, record := range records {
values[index] = cloneItem(record.item)
}
return values
}
func recordList(records []normalizedRecord) dnd.ItemRegistry {
if records == nil {
return dnd.ItemRegistry{}
}
return dnd.ItemRegistry{Items: recordValues(records)}
}
func duplicateWarning(retainedIndex int, removed []int) contracts.Warning {
const maxDisplayedIndices = 20
displayed := removed
if len(displayed) > maxDisplayedIndices {
displayed = displayed[:maxDisplayedIndices]
}
indices := make([]string, len(displayed))
for index, removedIndex := range displayed {
indices[index] = strconv.Itoa(removedIndex)
}
message := fmt.Sprintf("retained input index %d; removed input indices [%s]", retainedIndex, strings.Join(indices, ", "))
if omitted := len(removed) - len(displayed); omitted > 0 {
message += fmt.Sprintf("; %d additional removed input indices omitted", omitted)
}
return contracts.Warning{Scope: itemScope(retainedIndex), ReasonCode: ReasonCodeDuplicateItemCollapsed, Message: message}
}
func itemScope(index int) string { return fmt.Sprintf("items[%d]", index) }
func ModuleSpec() pipeline.ModuleSpec {
return pipeline.ModuleSpec{Key: Key, Stage: pipeline.StageNormalize, ExecutionClass: contracts.ExecutionClassLLMBacked, Requires: append([]string(nil), requiredCapabilities...), Provides: append([]string(nil), providedCapabilities...), ArtifactKind: dnd.ItemRegistryKind}
}
func Register(registry *pipeline.NormalizerRegistry) error {
return pipeline.RegisterNormalizerBuilder(registry, ModuleSpec(), validateOptions, func(request pipeline.BuildRequest) (contracts.Normalizer[dnd.ItemRegistry], error) {
options, err := DecodeOptions(request.Options)
if err != nil {
return nil, err
}
return New(request.Dependencies.LLM, options)
})
}
func validateOptions(options map[string]any) error { _, err := DecodeOptions(options); return err }
func DecodeOptions(options map[string]any) (Options, error) {
if err := pipeline.RejectUnknownOptions(options); err != nil {
return Options{}, normalizerErrorf("%w", err)
}
return Options{}, nil
}
func normalizerErrorf(format string, args ...any) error {
return fmt.Errorf("dnd item registry normalizer: "+format, args...)
}