Tighten Seriatim input validation
This commit is contained in:
@@ -6,6 +6,7 @@ import (
|
|||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
"fmt"
|
"fmt"
|
||||||
"math"
|
"math"
|
||||||
|
"math/big"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
@@ -124,7 +125,7 @@ func sourceUnit(segment segment, index int, seen map[string]struct{}) (source.So
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return source.SourceUnit{}, err
|
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)
|
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
|
}, 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())
|
raw := strings.TrimSpace(value.String())
|
||||||
if raw == "" {
|
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)
|
parsed, err := strconv.ParseFloat(raw, 64)
|
||||||
if err != nil {
|
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) {
|
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 {
|
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 {
|
func documentID(requestedID string, metadata map[string]any, rawDigest string) string {
|
||||||
|
|||||||
@@ -209,7 +209,7 @@ func TestParseRejectsInvalidInput(t *testing.T) {
|
|||||||
{
|
{
|
||||||
name: "non-numeric end",
|
name: "non-numeric end",
|
||||||
raw: validJSONWithSegment(`"id":"s1","start":0,"end":"late","speaker":"Narrator","text":"Synthetic text."`),
|
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",
|
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."`),
|
raw: validJSONWithSegment(`"id":"s1","start":2,"end":1,"speaker":"Narrator","text":"Synthetic text."`),
|
||||||
wantErr: []string{"end", "start"},
|
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 {
|
for _, tt := range tests {
|
||||||
|
|||||||
@@ -45,13 +45,21 @@ func decodeTranscript(raw []byte) (transcript, error) {
|
|||||||
if !ok {
|
if !ok {
|
||||||
return transcript{}, fmt.Errorf("segments are required")
|
return transcript{}, fmt.Errorf("segments are required")
|
||||||
}
|
}
|
||||||
var segments []segment
|
var segmentValues []json.RawMessage
|
||||||
if err := decodeJSON(segmentsRaw, &segments); err != nil {
|
if err := decodeJSON(segmentsRaw, &segmentValues); err != nil {
|
||||||
return transcript{}, fmt.Errorf("segments must be an array: %w", err)
|
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")
|
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{
|
return transcript{
|
||||||
Metadata: metadata,
|
Metadata: metadata,
|
||||||
@@ -59,6 +67,50 @@ func decodeTranscript(raw []byte) (transcript, error) {
|
|||||||
}, nil
|
}, 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 {
|
func decodeJSON(raw []byte, out any) error {
|
||||||
decoder := json.NewDecoder(bytes.NewReader(raw))
|
decoder := json.NewDecoder(bytes.NewReader(raw))
|
||||||
decoder.UseNumber()
|
decoder.UseNumber()
|
||||||
|
|||||||
Reference in New Issue
Block a user