From 39fcfba605ed9a0c017d80404eb9e52976e2b852 Mon Sep 17 00:00:00 2001 From: Eric Rakestraw Date: Fri, 3 Jul 2026 23:14:33 +0000 Subject: [PATCH] Tighten Seriatim input validation --- internal/modules/input/seriatim/adapter.go | 19 +++--- .../modules/input/seriatim/adapter_test.go | 7 ++- internal/modules/input/seriatim/model.go | 58 ++++++++++++++++++- 3 files changed, 73 insertions(+), 11 deletions(-) diff --git a/internal/modules/input/seriatim/adapter.go b/internal/modules/input/seriatim/adapter.go index 7738541..5da7c58 100644 --- a/internal/modules/input/seriatim/adapter.go +++ b/internal/modules/input/seriatim/adapter.go @@ -6,6 +6,7 @@ import ( "encoding/hex" "fmt" "math" + "math/big" "strconv" "strings" @@ -124,7 +125,7 @@ func sourceUnit(segment segment, index int, seen map[string]struct{}) (source.So if err != nil { return source.SourceUnit{}, err } - if end < start { + if end.Cmp(start) < 0 { return source.SourceUnit{}, inputErrorf("segment %q end must be greater than or equal to start", segment.ID) } @@ -144,22 +145,26 @@ func sourceUnit(segment segment, index int, seen map[string]struct{}) (source.So }, nil } -func validTimestamp(value fmt.Stringer, label string) (float64, error) { +func validTimestamp(value fmt.Stringer, label string) (*big.Rat, error) { raw := strings.TrimSpace(value.String()) if raw == "" { - return 0, inputErrorf("%s must not be empty", label) + return nil, inputErrorf("%s must not be empty", label) } parsed, err := strconv.ParseFloat(raw, 64) if err != nil { - return 0, inputErrorf("%s must be a valid number: %w", label, err) + return nil, inputErrorf("%s must be a valid number: %w", label, err) } if math.IsInf(parsed, 0) || math.IsNaN(parsed) { - return 0, inputErrorf("%s must be finite", label) + return nil, inputErrorf("%s must be finite", label) } if parsed < 0 { - return 0, inputErrorf("%s must not be negative", label) + return nil, inputErrorf("%s must not be negative", label) } - return parsed, nil + rat, ok := new(big.Rat).SetString(raw) + if !ok { + return nil, inputErrorf("%s must be a valid number", label) + } + return rat, nil } func documentID(requestedID string, metadata map[string]any, rawDigest string) string { diff --git a/internal/modules/input/seriatim/adapter_test.go b/internal/modules/input/seriatim/adapter_test.go index ff58f2c..baa630a 100644 --- a/internal/modules/input/seriatim/adapter_test.go +++ b/internal/modules/input/seriatim/adapter_test.go @@ -209,7 +209,7 @@ func TestParseRejectsInvalidInput(t *testing.T) { { name: "non-numeric end", raw: validJSONWithSegment(`"id":"s1","start":0,"end":"late","speaker":"Narrator","text":"Synthetic text."`), - wantErr: []string{"segments", "json"}, + wantErr: []string{"segment[0]", "end", "number"}, }, { name: "non-finite timestamp", @@ -221,6 +221,11 @@ func TestParseRejectsInvalidInput(t *testing.T) { raw: validJSONWithSegment(`"id":"s1","start":2,"end":1,"speaker":"Narrator","text":"Synthetic text."`), wantErr: []string{"end", "start"}, }, + { + name: "end before start beyond float precision", + raw: validJSONWithSegment(`"id":"s1","start":9007199254740993,"end":9007199254740992,"speaker":"Narrator","text":"Synthetic text."`), + wantErr: []string{"end", "start"}, + }, } for _, tt := range tests { diff --git a/internal/modules/input/seriatim/model.go b/internal/modules/input/seriatim/model.go index ac1f8ac..2cdd123 100644 --- a/internal/modules/input/seriatim/model.go +++ b/internal/modules/input/seriatim/model.go @@ -45,13 +45,21 @@ func decodeTranscript(raw []byte) (transcript, error) { if !ok { return transcript{}, fmt.Errorf("segments are required") } - var segments []segment - if err := decodeJSON(segmentsRaw, &segments); err != nil { + var segmentValues []json.RawMessage + if err := decodeJSON(segmentsRaw, &segmentValues); err != nil { return transcript{}, fmt.Errorf("segments must be an array: %w", err) } - if segments == nil { + if segmentValues == nil { return transcript{}, fmt.Errorf("segments must be an array") } + segments := make([]segment, 0, len(segmentValues)) + for i, rawSegment := range segmentValues { + segment, err := decodeSegment(rawSegment, i) + if err != nil { + return transcript{}, err + } + segments = append(segments, segment) + } return transcript{ Metadata: metadata, @@ -59,6 +67,50 @@ func decodeTranscript(raw []byte) (transcript, error) { }, nil } +func decodeSegment(raw []byte, index int) (segment, error) { + var fields map[string]json.RawMessage + if err := decodeJSON(raw, &fields); err != nil { + return segment{}, fmt.Errorf("segment[%d] must be an object: %w", index, err) + } + if fields == nil { + return segment{}, fmt.Errorf("segment[%d] must be an object", index) + } + + var decoded segment + if err := decodeOptionalString(fields, "id", &decoded.ID); err != nil { + return segment{}, fmt.Errorf("segment[%d] id must be a string: %w", index, err) + } + if err := decodeOptionalNumber(fields, "start", &decoded.Start); err != nil { + return segment{}, fmt.Errorf("segment[%d] start must be a number: %w", index, err) + } + if err := decodeOptionalNumber(fields, "end", &decoded.End); err != nil { + return segment{}, fmt.Errorf("segment[%d] end must be a number: %w", index, err) + } + if err := decodeOptionalString(fields, "speaker", &decoded.Speaker); err != nil { + return segment{}, fmt.Errorf("segment[%d] speaker must be a string: %w", index, err) + } + if err := decodeOptionalString(fields, "text", &decoded.Text); err != nil { + return segment{}, fmt.Errorf("segment[%d] text must be a string: %w", index, err) + } + return decoded, nil +} + +func decodeOptionalString(fields map[string]json.RawMessage, key string, out *string) error { + raw, ok := fields[key] + if !ok { + return nil + } + return decodeJSON(raw, out) +} + +func decodeOptionalNumber(fields map[string]json.RawMessage, key string, out *json.Number) error { + raw, ok := fields[key] + if !ok { + return nil + } + return decodeJSON(raw, out) +} + func decodeJSON(raw []byte, out any) error { decoder := json.NewDecoder(bytes.NewReader(raw)) decoder.UseNumber()