diff --git a/pkg/media/random_access_src.go b/pkg/media/random_access_src.go new file mode 100644 index 000000000..d16c00d98 --- /dev/null +++ b/pkg/media/random_access_src.go @@ -0,0 +1,128 @@ +package media + +import ( + "context" + "errors" + "fmt" + "io" + + "github.com/go-gst/go-gst/gst" + "github.com/go-gst/go-gst/gst/app" + "stream.place/streamplace/pkg/log" +) + +// RandomAccessSrcBin wraps an appsrc element in random-access (BYTES) mode +// inside a gst.Bin with a single ghost src pad named "src". The element is +// driven by the supplied io.ReaderAt, whose total length must be known up +// front and passed as size. +// +// The intended use is feeding arbitrary uploaded media (which lives on local +// disk or S3) into parsebin/qtdemux/etc. parsebin needs random access so it +// can find the moov atom near the end of MP4 files. +// +// The context is captured for the bin's lifetime. Callers MUST cancel ctx +// when they're done with the bin so any in-flight ReadAt — particularly the +// S3 case, where ReadAt is an HTTP request — can abort cleanly. +func RandomAccessSrcBin(ctx context.Context, name string, src io.ReaderAt, size int64) (*gst.Bin, error) { + if size < 0 { + return nil, fmt.Errorf("size must be non-negative, got %d", size) + } + bin := gst.NewBin(name + "-bin") + + appSrc, err := gst.NewElementWithProperties("appsrc", map[string]interface{}{ + "name": name, + }) + if err != nil { + return nil, fmt.Errorf("create appsrc: %w", err) + } + if err := bin.Add(appSrc); err != nil { + return nil, fmt.Errorf("add appsrc to bin: %w", err) + } + + source := app.SrcFromElement(appSrc) + // appsrc defaults to GST_FORMAT_BYTES, which is what we want for byte- + // offset seeking. SetSize lets downstream elements (parsebin, qtdemux) + // query a duration in bytes and seek to the end of the stream. + source.SetSize(size) + source.SetStreamType(app.AppStreamTypeRandomAccess) + + // pos tracks where the next NeedDataFunc read should start. eos is a + // local guard against pushing past end-of-stream after we've already + // emitted it; SeekDataFunc clears it so post-EOS seeks (which qtdemux + // performs after parsing moov) can resume reads. + // + // appsrc serializes its own callbacks, so no mutex is needed. + var pos int64 + var eos bool + + source.SetCallbacks(&app.SourceCallbacks{ + NeedDataFunc: func(self *app.Source, length uint) { + if ctx.Err() != nil { + self.EndStream() + return + } + if eos { + return + } + remaining := size - pos + if remaining <= 0 { + self.EndStream() + eos = true + return + } + n := int64(length) + if n <= 0 { + n = 64 * 1024 + } + if n > remaining { + n = remaining + } + buf := make([]byte, n) + read, err := src.ReadAt(buf, pos) + if read > 0 { + gbuf := gst.NewBufferWithSize(int64(read)) + gbuf.Map(gst.MapWrite).WriteData(buf[:read]) + gbuf.Unmap() + if ret := self.PushBuffer(gbuf); ret != gst.FlowOK { + log.Debug(ctx, "RandomAccessSrcBin: push buffer non-OK", "ret", ret.String()) + } + pos += int64(read) + } + if err != nil && !errors.Is(err, io.EOF) { + log.Error(ctx, "RandomAccessSrcBin: read failed", "offset", pos, "error", err) + self.Error("read failed", err) + eos = true + return + } + if errors.Is(err, io.EOF) || pos >= size { + self.EndStream() + eos = true + } + }, + SeekDataFunc: func(self *app.Source, offset uint64) bool { + if ctx.Err() != nil { + return false + } + if int64(offset) > size { + return false + } + pos = int64(offset) + eos = false + return true + }, + }) + + srcPad := appSrc.GetStaticPad("src") + if srcPad == nil { + return nil, fmt.Errorf("appsrc missing src pad") + } + ghost := gst.NewGhostPad("src", srcPad) + if ghost == nil { + return nil, fmt.Errorf("create ghost pad") + } + if !bin.AddPad(ghost.Pad) { + return nil, fmt.Errorf("add ghost pad to bin") + } + + return bin, nil +} diff --git a/pkg/media/random_access_src_test.go b/pkg/media/random_access_src_test.go new file mode 100644 index 000000000..526d9bf2b --- /dev/null +++ b/pkg/media/random_access_src_test.go @@ -0,0 +1,137 @@ +package media + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "os" + "testing" + + "github.com/go-gst/go-gst/gst" + "github.com/go-gst/go-gst/gst/app" + "github.com/stretchr/testify/require" + "stream.place/streamplace/pkg/log" +) + +// TestRandomAccessSrcBin_Passthrough verifies that the bin drains the +// supplied ReaderAt end-to-end and the bytes arriving at an appsink +// exactly match the source data, both byte-count and content. +func TestRandomAccessSrcBin_Passthrough(t *testing.T) { + withNoGSTLeaks(t, func() { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + ctx = log.WithLogValues(ctx, "test", "TestRandomAccessSrcBin_Passthrough") + + fixture, err := os.ReadFile(getFixture("sample-segment.mp4")) + require.NoError(t, err) + require.NotEmpty(t, fixture) + + pipeline, err := gst.NewPipeline("ra-src-passthrough") + require.NoError(t, err) + + srcBin, err := RandomAccessSrcBin(ctx, "ra-src", bytes.NewReader(fixture), int64(len(fixture))) + require.NoError(t, err) + require.NoError(t, pipeline.Add(srcBin.Element)) + + sinkEle, err := gst.NewElementWithProperties("appsink", map[string]interface{}{ + "name": "ra-sink", + "sync": false, + }) + require.NoError(t, err) + require.NoError(t, pipeline.Add(sinkEle)) + + ghostPad := srcBin.GetStaticPad("src") + require.NotNil(t, ghostPad) + sinkPad := sinkEle.GetStaticPad("sink") + require.NotNil(t, sinkPad) + require.Equal(t, gst.PadLinkOK, ghostPad.Link(sinkPad)) + + out := &bytes.Buffer{} + sink := app.SinkFromElement(sinkEle) + sink.SetCallbacks(&app.SinkCallbacks{ + NewSampleFunc: WriterNewSample(ctx, out), + }) + + errCh := make(chan error, 1) + go func() { errCh <- HandleBusMessages(ctx, pipeline) }() + + require.NoError(t, pipeline.SetState(gst.StatePlaying)) + require.NoError(t, <-errCh) + require.NoError(t, pipeline.BlockSetState(gst.StateNull)) + + require.Equal(t, fixture, out.Bytes()) + }) +} + +// TestRandomAccessSrcBin_Seek exercises the random-access path by having a +// downstream identity element issue manual seek events. We hash the bytes +// arriving at the appsink and compare against a known good full-file hash; +// if any seek lands at the wrong offset or returns the wrong slice of +// bytes, this catches it. +func TestRandomAccessSrcBin_Seek(t *testing.T) { + withNoGSTLeaks(t, func() { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + ctx = log.WithLogValues(ctx, "test", "TestRandomAccessSrcBin_Seek") + + fixture, err := os.ReadFile(getFixture("sample-segment.mp4")) + require.NoError(t, err) + require.GreaterOrEqual(t, len(fixture), 1024) + + // Read three disjoint regions via two seeks. Compare to direct + // fixture slices. + reader := bytes.NewReader(fixture) + size := int64(len(fixture)) + + check := func(name string, offset, length int64) { + pipeline, err := gst.NewPipeline("ra-src-seek-" + name) + require.NoError(t, err) + + srcBin, err := RandomAccessSrcBin(ctx, "ra-src", reader, size) + require.NoError(t, err) + require.NoError(t, pipeline.Add(srcBin.Element)) + + sinkEle, err := gst.NewElementWithProperties("appsink", map[string]interface{}{ + "name": "ra-sink", + "sync": false, + }) + require.NoError(t, err) + require.NoError(t, pipeline.Add(sinkEle)) + + require.Equal(t, gst.PadLinkOK, srcBin.GetStaticPad("src").Link(sinkEle.GetStaticPad("sink"))) + + out := &bytes.Buffer{} + sink := app.SinkFromElement(sinkEle) + sink.SetCallbacks(&app.SinkCallbacks{ + NewSampleFunc: WriterNewSample(ctx, out), + }) + + errCh := make(chan error, 1) + go func() { errCh <- HandleBusMessages(ctx, pipeline) }() + require.NoError(t, pipeline.SetState(gst.StatePaused)) + + // Issue a byte-format seek into the desired range. + ok := pipeline.SeekSimple(offset, gst.FormatBytes, gst.SeekFlagFlush) + require.True(t, ok, "%s: seek to %d failed", name, offset) + + require.NoError(t, pipeline.SetState(gst.StatePlaying)) + require.NoError(t, <-errCh) + require.NoError(t, pipeline.BlockSetState(gst.StateNull)) + + got := out.Bytes() + require.GreaterOrEqual(t, len(got), int(length), "%s: expected at least %d bytes, got %d", name, length, len(got)) + + wantH := sha256.Sum256(fixture[offset : offset+length]) + gotH := sha256.Sum256(got[:length]) + require.Equal(t, hex.EncodeToString(wantH[:]), hex.EncodeToString(gotH[:]), "%s: byte mismatch at offset %d len %d", name, offset, length) + } + + // Tail slice + check("tail", size-1024, 1024) + // Middle slice + check("middle", size/2, 1024) + // Head slice + check("head", 0, 1024) + }) +} diff --git a/pkg/s3/readerat.go b/pkg/s3/readerat.go new file mode 100644 index 000000000..3ffed73c6 --- /dev/null +++ b/pkg/s3/readerat.go @@ -0,0 +1,148 @@ +package s3 + +import ( + "context" + "errors" + "fmt" + "io" + "strings" + "sync" + + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/service/s3" +) + +// ReaderAt is an io.ReaderAt against an S3 object. It re-uses a single +// open GetObject body for sequential reads and transparently closes & +// reopens it on a non-sequential read. This makes it efficient for the +// common case of "seek once, then drain forward" that gstreamer demuxers +// produce in pull mode, while still allowing arbitrary jumps. +// +// Concurrent ReadAt calls are serialized by an internal mutex; this matches +// the gstreamer appsrc usage where need-data and seek-data callbacks fire +// on a single streaming thread. +type ReaderAt struct { + ctx context.Context + client *s3.Client + bucket string + key string + size int64 + + mu sync.Mutex + body io.ReadCloser + pos int64 +} + +// NewReaderAt issues a HEAD against the given S3 object to discover its +// size, then returns a ReaderAt. The provided context is used for both +// the HEAD and any subsequent GetObject requests; cancelling it aborts +// in-flight reads and is the supported way to tear the ReaderAt down. +func NewReaderAt(ctx context.Context, client *s3.Client, bucket, key string) (*ReaderAt, error) { + head, err := client.HeadObject(ctx, &s3.HeadObjectInput{ + Bucket: aws.String(bucket), + Key: aws.String(key), + }) + if err != nil { + return nil, fmt.Errorf("head s3://%s/%s: %w", bucket, key, err) + } + if head.ContentLength == nil { + return nil, fmt.Errorf("s3://%s/%s missing content-length", bucket, key) + } + return &ReaderAt{ + ctx: ctx, + client: client, + bucket: bucket, + key: key, + size: *head.ContentLength, + }, nil +} + +// Size returns the object size discovered at construction time. +func (r *ReaderAt) Size() int64 { return r.size } + +// ReadAt implements io.ReaderAt. Sequential reads (off == previous end) +// drain the open body; non-sequential reads close the body and issue a +// fresh ranged GetObject starting at off. +func (r *ReaderAt) ReadAt(p []byte, off int64) (int, error) { + if off < 0 { + return 0, fmt.Errorf("negative offset %d", off) + } + if off >= r.size { + return 0, io.EOF + } + + r.mu.Lock() + defer r.mu.Unlock() + + if r.body == nil || off != r.pos { + _ = r.closeBodyLocked() + if err := r.openLocked(off); err != nil { + return 0, err + } + } + + want := len(p) + remaining := r.size - off + if int64(want) > remaining { + want = int(remaining) + } + n, err := io.ReadFull(r.body, p[:want]) + r.pos += int64(n) + + // Truncated reads at the tail of the object surface as + // ErrUnexpectedEOF from ReadFull; report EOF to the caller and + // drop the now-empty body so the next read reopens. + if errors.Is(err, io.ErrUnexpectedEOF) || (err == nil && r.pos >= r.size) { + _ = r.closeBodyLocked() + if r.pos >= r.size { + err = io.EOF + } + } + return n, err +} + +func (r *ReaderAt) openLocked(off int64) error { + resp, err := r.client.GetObject(r.ctx, &s3.GetObjectInput{ + Bucket: aws.String(r.bucket), + Key: aws.String(r.key), + Range: aws.String(fmt.Sprintf("bytes=%d-", off)), + }) + if err != nil { + return fmt.Errorf("get s3://%s/%s bytes=%d-: %w", r.bucket, r.key, off, err) + } + r.body = resp.Body + r.pos = off + return nil +} + +func (r *ReaderAt) closeBodyLocked() error { + if r.body == nil { + return nil + } + err := r.body.Close() + r.body = nil + return err +} + +// Close releases any open GetObject body. Safe to call concurrently with +// ReadAt; in-flight reads observe the closed body as an error. +func (r *ReaderAt) Close() error { + r.mu.Lock() + defer r.mu.Unlock() + return r.closeBodyLocked() +} + +// ParseURL parses an "s3://bucket/key" URL into its bucket and key parts. +// The key may contain forward slashes. +func ParseURL(u string) (bucket, key string, err error) { + const prefix = "s3://" + if !strings.HasPrefix(u, prefix) { + return "", "", fmt.Errorf("not an s3:// URL: %q", u) + } + rest := strings.TrimPrefix(u, prefix) + slash := strings.IndexByte(rest, '/') + if slash <= 0 || slash == len(rest)-1 { + return "", "", fmt.Errorf("invalid s3:// URL: %q", u) + } + return rest[:slash], rest[slash+1:], nil +} diff --git a/pkg/s3/readerat_test.go b/pkg/s3/readerat_test.go new file mode 100644 index 000000000..8c5837e55 --- /dev/null +++ b/pkg/s3/readerat_test.go @@ -0,0 +1,43 @@ +package s3 + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestParseURL(t *testing.T) { + t.Run("happy path", func(t *testing.T) { + bucket, key, err := ParseURL("s3://my-bucket/path/to/object.mp4") + require.NoError(t, err) + require.Equal(t, "my-bucket", bucket) + require.Equal(t, "path/to/object.mp4", key) + }) + + t.Run("simple key", func(t *testing.T) { + bucket, key, err := ParseURL("s3://b/k") + require.NoError(t, err) + require.Equal(t, "b", bucket) + require.Equal(t, "k", key) + }) + + t.Run("missing scheme", func(t *testing.T) { + _, _, err := ParseURL("/just/a/path") + require.Error(t, err) + }) + + t.Run("missing key", func(t *testing.T) { + _, _, err := ParseURL("s3://only-bucket") + require.Error(t, err) + }) + + t.Run("empty key", func(t *testing.T) { + _, _, err := ParseURL("s3://bucket/") + require.Error(t, err) + }) + + t.Run("missing bucket", func(t *testing.T) { + _, _, err := ParseURL("s3:///key") + require.Error(t, err) + }) +}