diff --git a/internal/trim/apply.go b/internal/trim/apply.go index 7c71c97..f9b8680 100644 --- a/internal/trim/apply.go +++ b/internal/trim/apply.go @@ -44,54 +44,29 @@ type MinimalResult struct { RemovedIDs []int } +type projection struct { + retainedIndexes []int + oldToNewID map[int]int + removedIDs []int +} + // Apply trims a full seriatim output transcript by segment ID. func Apply(input schema.Transcript, opts Options) (Result, error) { - if err := validateMode(opts.Mode); err != nil { - return Result{}, err - } - - selected := opts.Selector.IDs() - if len(selected) == 0 { - return Result{}, fmt.Errorf("selector cannot be empty") - } - inputIDs := make([]int, len(input.Segments)) for index, segment := range input.Segments { inputIDs[index] = segment.ID } - - idIndex, err := validateInputIDs(inputIDs) + proj, err := projectSegmentIDs(inputIDs, opts) if err != nil { return Result{}, err } - if err := validateSelectedIDsExist(selected, idIndex); err != nil { - return Result{}, err - } - - kept := make([]schema.Segment, 0, len(input.Segments)) - removed := make([]int, 0, len(input.Segments)) - oldToNew := make(map[int]int, len(input.Segments)) - for _, segment := range input.Segments { - keep := opts.Mode == ModeKeep && opts.Selector.Contains(segment.ID) - if opts.Mode == ModeRemove { - keep = !opts.Selector.Contains(segment.ID) - } - - if !keep { - removed = append(removed, segment.ID) - continue - } - - rewritten := copySegment(segment) - rewritten.ID = len(kept) + 1 + kept := make([]schema.Segment, len(proj.retainedIndexes)) + for outputIndex, inputIndex := range proj.retainedIndexes { + rewritten := copySegment(input.Segments[inputIndex]) + rewritten.ID = outputIndex + 1 rewritten.OverlapGroupID = 0 - kept = append(kept, rewritten) - oldToNew[segment.ID] = rewritten.ID - } - - if len(kept) == 0 && !opts.AllowEmpty { - return Result{}, fmt.Errorf("trim operation produced an empty transcript; set AllowEmpty to true to permit this") + kept[outputIndex] = rewritten } kept, groups := recomputeOverlapGroups(kept) @@ -104,62 +79,35 @@ func Apply(input schema.Transcript, opts Options) (Result, error) { out.OverlapGroups = groups return Result{ Transcript: out, - OldToNewID: oldToNew, - RemovedIDs: removed, + OldToNewID: proj.oldToNewID, + RemovedIDs: proj.removedIDs, }, nil } // ApplyIntermediate trims an intermediate seriatim output transcript by // segment ID. func ApplyIntermediate(input schema.IntermediateTranscript, opts Options) (IntermediateResult, error) { - if err := validateMode(opts.Mode); err != nil { - return IntermediateResult{}, err - } - - selected := opts.Selector.IDs() - if len(selected) == 0 { - return IntermediateResult{}, fmt.Errorf("selector cannot be empty") - } - inputIDs := make([]int, len(input.Segments)) for index, segment := range input.Segments { inputIDs[index] = segment.ID } - idIndex, err := validateInputIDs(inputIDs) + proj, err := projectSegmentIDs(inputIDs, opts) if err != nil { return IntermediateResult{}, err } - if err := validateSelectedIDsExist(selected, idIndex); err != nil { - return IntermediateResult{}, err - } - - kept := make([]schema.IntermediateSegment, 0, len(input.Segments)) - removed := make([]int, 0, len(input.Segments)) - oldToNew := make(map[int]int, len(input.Segments)) - for _, segment := range input.Segments { - keep := opts.Mode == ModeKeep && opts.Selector.Contains(segment.ID) - if opts.Mode == ModeRemove { - keep = !opts.Selector.Contains(segment.ID) - } - if !keep { - removed = append(removed, segment.ID) - continue - } + kept := make([]schema.IntermediateSegment, len(proj.retainedIndexes)) + for outputIndex, inputIndex := range proj.retainedIndexes { + segment := input.Segments[inputIndex] rewritten := schema.IntermediateSegment{ - ID: len(kept) + 1, + ID: outputIndex + 1, Start: segment.Start, End: segment.End, Speaker: segment.Speaker, Text: segment.Text, Categories: append([]string(nil), segment.Categories...), } - kept = append(kept, rewritten) - oldToNew[segment.ID] = rewritten.ID - } - - if len(kept) == 0 && !opts.AllowEmpty { - return IntermediateResult{}, fmt.Errorf("trim operation produced an empty transcript; set AllowEmpty to true to permit this") + kept[outputIndex] = rewritten } return IntermediateResult{ @@ -171,60 +119,33 @@ func ApplyIntermediate(input schema.IntermediateTranscript, opts Options) (Inter }, Segments: kept, }, - OldToNewID: oldToNew, - RemovedIDs: removed, + OldToNewID: proj.oldToNewID, + RemovedIDs: proj.removedIDs, }, nil } // ApplyMinimal trims a minimal seriatim output transcript by segment ID. func ApplyMinimal(input schema.MinimalTranscript, opts Options) (MinimalResult, error) { - if err := validateMode(opts.Mode); err != nil { - return MinimalResult{}, err - } - - selected := opts.Selector.IDs() - if len(selected) == 0 { - return MinimalResult{}, fmt.Errorf("selector cannot be empty") - } - inputIDs := make([]int, len(input.Segments)) for index, segment := range input.Segments { inputIDs[index] = segment.ID } - idIndex, err := validateInputIDs(inputIDs) + proj, err := projectSegmentIDs(inputIDs, opts) if err != nil { return MinimalResult{}, err } - if err := validateSelectedIDsExist(selected, idIndex); err != nil { - return MinimalResult{}, err - } - - kept := make([]schema.MinimalSegment, 0, len(input.Segments)) - removed := make([]int, 0, len(input.Segments)) - oldToNew := make(map[int]int, len(input.Segments)) - for _, segment := range input.Segments { - keep := opts.Mode == ModeKeep && opts.Selector.Contains(segment.ID) - if opts.Mode == ModeRemove { - keep = !opts.Selector.Contains(segment.ID) - } - if !keep { - removed = append(removed, segment.ID) - continue - } + kept := make([]schema.MinimalSegment, len(proj.retainedIndexes)) + for outputIndex, inputIndex := range proj.retainedIndexes { + segment := input.Segments[inputIndex] rewritten := schema.MinimalSegment{ - ID: len(kept) + 1, + ID: outputIndex + 1, Start: segment.Start, End: segment.End, Speaker: segment.Speaker, Text: segment.Text, } - kept = append(kept, rewritten) - oldToNew[segment.ID] = rewritten.ID - } - - if len(kept) == 0 && !opts.AllowEmpty { - return MinimalResult{}, fmt.Errorf("trim operation produced an empty transcript; set AllowEmpty to true to permit this") + kept[outputIndex] = rewritten } return MinimalResult{ @@ -236,11 +157,53 @@ func ApplyMinimal(input schema.MinimalTranscript, opts Options) (MinimalResult, }, Segments: kept, }, - OldToNewID: oldToNew, - RemovedIDs: removed, + OldToNewID: proj.oldToNewID, + RemovedIDs: proj.removedIDs, }, nil } +func projectSegmentIDs(ids []int, opts Options) (projection, error) { + if err := validateMode(opts.Mode); err != nil { + return projection{}, err + } + + selected := opts.Selector.IDs() + if len(selected) == 0 { + return projection{}, fmt.Errorf("selector cannot be empty") + } + + idIndex, err := validateInputIDs(ids) + if err != nil { + return projection{}, err + } + if err := validateSelectedIDsExist(selected, idIndex); err != nil { + return projection{}, err + } + + result := projection{ + retainedIndexes: make([]int, 0, len(ids)), + oldToNewID: make(map[int]int, len(ids)), + removedIDs: make([]int, 0, len(ids)), + } + for index, id := range ids { + keep := opts.Mode == ModeKeep && opts.Selector.Contains(id) + if opts.Mode == ModeRemove { + keep = !opts.Selector.Contains(id) + } + if !keep { + result.removedIDs = append(result.removedIDs, id) + continue + } + result.retainedIndexes = append(result.retainedIndexes, index) + result.oldToNewID[id] = len(result.retainedIndexes) + } + + if len(result.retainedIndexes) == 0 && !opts.AllowEmpty { + return projection{}, fmt.Errorf("trim operation produced an empty transcript; set AllowEmpty to true to permit this") + } + return result, nil +} + func validateMode(mode Mode) error { switch mode { case ModeKeep, ModeRemove: diff --git a/internal/trim/apply_test.go b/internal/trim/apply_test.go index cabdfbd..559f039 100644 --- a/internal/trim/apply_test.go +++ b/internal/trim/apply_test.go @@ -399,6 +399,106 @@ func TestApplyMinimalDoesNotIncludeOverlapGroups(t *testing.T) { } } +func TestApplySelectorPolicyIsSharedAcrossSchemas(t *testing.T) { + type testCase struct { + name string + opts Options + wantTexts []string + wantOldToNew map[int]int + wantRemoved []int + wantSegmentCount int + wantErrorSubstring string + } + + cases := []testCase{ + { + name: "keep preserves input order regardless of selector order", + opts: Options{Mode: ModeKeep, Selector: mustParseSelector(t, "4,1,3")}, + wantTexts: []string{"alpha", "gamma", "delta"}, + wantOldToNew: map[int]int{1: 1, 3: 2, 4: 3}, + wantRemoved: []int{2}, + wantSegmentCount: 3, + }, + { + name: "remove reports deterministic renumbering metadata", + opts: Options{Mode: ModeRemove, Selector: mustParseSelector(t, "2,4")}, + wantTexts: []string{"alpha", "gamma"}, + wantOldToNew: map[int]int{1: 1, 3: 2}, + wantRemoved: []int{2, 4}, + wantSegmentCount: 2, + }, + { + name: "missing selected id returns error", + opts: Options{Mode: ModeKeep, Selector: mustParseSelector(t, "9")}, + wantErrorSubstring: "does not exist", + }, + { + name: "empty selector returns error", + opts: Options{Mode: ModeKeep, Selector: Selector{}}, + wantErrorSubstring: "selector cannot be empty", + }, + { + name: "invalid mode returns error", + opts: Options{Mode: Mode("bad"), Selector: mustParseSelector(t, "1")}, + wantErrorSubstring: `invalid trim mode "bad"`, + }, + { + name: "empty output blocked when allow empty is false", + opts: Options{Mode: ModeRemove, Selector: mustParseSelector(t, "1-4")}, + wantErrorSubstring: "empty transcript", + }, + { + name: "empty output allowed when allow empty is true", + opts: Options{Mode: ModeRemove, Selector: mustParseSelector(t, "1-4"), AllowEmpty: true}, + wantTexts: []string{}, + wantOldToNew: map[int]int{}, + wantRemoved: []int{1, 2, 3, 4}, + wantSegmentCount: 0, + }, + } + + for _, test := range cases { + t.Run(test.name, func(t *testing.T) { + fullInput := fullTranscriptFixture() + intermediateInput := intermediateFixture() + minimalInput := minimalFixture() + + fullResult, fullErr := Apply(fullInput, test.opts) + intermediateResult, intermediateErr := ApplyIntermediate(intermediateInput, test.opts) + minimalResult, minimalErr := ApplyMinimal(minimalInput, test.opts) + + if test.wantErrorSubstring != "" { + assertErrorContains(t, fullErr, test.wantErrorSubstring) + assertErrorContains(t, intermediateErr, test.wantErrorSubstring) + assertErrorContains(t, minimalErr, test.wantErrorSubstring) + return + } + if fullErr != nil { + t.Fatalf("apply full failed: %v", fullErr) + } + if intermediateErr != nil { + t.Fatalf("apply intermediate failed: %v", intermediateErr) + } + if minimalErr != nil { + t.Fatalf("apply minimal failed: %v", minimalErr) + } + + assertIntSlice(t, extractFullIDs(fullResult.Transcript.Segments), extractSequentialIDs(test.wantSegmentCount)) + assertIntSlice(t, extractIntermediateIDs(intermediateResult.Transcript.Segments), extractSequentialIDs(test.wantSegmentCount)) + assertIntSlice(t, extractMinimalIDs(minimalResult.Transcript.Segments), extractSequentialIDs(test.wantSegmentCount)) + assertStringSlice(t, extractFullTexts(fullResult.Transcript.Segments), test.wantTexts) + assertStringSlice(t, extractIntermediateTexts(intermediateResult.Transcript.Segments), test.wantTexts) + assertStringSlice(t, extractMinimalTexts(minimalResult.Transcript.Segments), test.wantTexts) + assertIntMap(t, fullResult.OldToNewID, test.wantOldToNew) + assertIntMap(t, intermediateResult.OldToNewID, test.wantOldToNew) + assertIntMap(t, minimalResult.OldToNewID, test.wantOldToNew) + assertIntSlice(t, fullResult.RemovedIDs, test.wantRemoved) + assertIntSlice(t, intermediateResult.RemovedIDs, test.wantRemoved) + assertIntSlice(t, minimalResult.RemovedIDs, test.wantRemoved) + }) + } +} + func TestApplyOutputInvariantsValidAfterRenumberAndOverlapRecompute(t *testing.T) { input := overlapTranscriptFixture() selector := mustParseSelector(t, "2,1") @@ -666,3 +766,108 @@ func equalStringSlices(got []string, want []string) bool { } return true } + +func assertErrorContains(t *testing.T, err error, substring string) { + t.Helper() + if err == nil { + t.Fatalf("expected error containing %q", substring) + } + if !strings.Contains(err.Error(), substring) { + t.Fatalf("error %q does not contain %q", err.Error(), substring) + } +} + +func assertStringSlice(t *testing.T, got []string, want []string) { + t.Helper() + if !equalStringSlices(got, want) { + t.Fatalf("slice = %v, want %v", got, want) + } +} + +func extractSequentialIDs(count int) []int { + ids := make([]int, count) + for index := range ids { + ids[index] = index + 1 + } + return ids +} + +func extractFullIDs(segments []schema.Segment) []int { + ids := make([]int, len(segments)) + for index, segment := range segments { + ids[index] = segment.ID + } + return ids +} + +func extractIntermediateIDs(segments []schema.IntermediateSegment) []int { + ids := make([]int, len(segments)) + for index, segment := range segments { + ids[index] = segment.ID + } + return ids +} + +func extractMinimalIDs(segments []schema.MinimalSegment) []int { + ids := make([]int, len(segments)) + for index, segment := range segments { + ids[index] = segment.ID + } + return ids +} + +func extractFullTexts(segments []schema.Segment) []string { + texts := make([]string, len(segments)) + for index, segment := range segments { + texts[index] = segment.Text + } + return texts +} + +func extractIntermediateTexts(segments []schema.IntermediateSegment) []string { + texts := make([]string, len(segments)) + for index, segment := range segments { + texts[index] = segment.Text + } + return texts +} + +func extractMinimalTexts(segments []schema.MinimalSegment) []string { + texts := make([]string, len(segments)) + for index, segment := range segments { + texts[index] = segment.Text + } + return texts +} + +func intermediateFixture() schema.IntermediateTranscript { + return schema.IntermediateTranscript{ + Metadata: schema.IntermediateMetadata{ + Application: "seriatim", + Version: "v-test", + OutputSchema: schema.OutputSchemaIntermediate, + }, + Segments: []schema.IntermediateSegment{ + {ID: 1, Start: 1, End: 2, Speaker: "Alice", Text: "alpha", Categories: []string{"word-run"}}, + {ID: 2, Start: 2, End: 3, Speaker: "Bob", Text: "beta", Categories: []string{"filler", "backchannel"}}, + {ID: 3, Start: 3, End: 4, Speaker: "Carol", Text: "gamma", Categories: []string{"normal"}}, + {ID: 4, Start: 4, End: 5, Speaker: "Dan", Text: "delta", Categories: []string{"normal"}}, + }, + } +} + +func minimalFixture() schema.MinimalTranscript { + return schema.MinimalTranscript{ + Metadata: schema.MinimalMetadata{ + Application: "seriatim", + Version: "v-test", + OutputSchema: schema.OutputSchemaMinimal, + }, + Segments: []schema.MinimalSegment{ + {ID: 1, Start: 1, End: 2, Speaker: "Alice", Text: "alpha"}, + {ID: 2, Start: 2, End: 3, Speaker: "Bob", Text: "beta"}, + {ID: 3, Start: 3, End: 4, Speaker: "Carol", Text: "gamma"}, + {ID: 4, Start: 4, End: 5, Speaker: "Dan", Text: "delta"}, + }, + } +}