From 3db2911072ed454ba02fb2dc40d9e9c62cf4b794 Mon Sep 17 00:00:00 2001 From: Eli Mallon Date: Sat, 9 May 2026 13:08:57 -0700 Subject: [PATCH] vod: track s3 uploads in statedb --- pkg/director/s3_upload.go | 2 +- pkg/s3/s3.go | 54 +++++++++++++++++++++++------ pkg/statedb/s3_segment.go | 73 +++++++++++++++++++++++++++++++++++++++ pkg/statedb/statedb.go | 1 + 4 files changed, 118 insertions(+), 12 deletions(-) create mode 100644 pkg/statedb/s3_segment.go diff --git a/pkg/director/s3_upload.go b/pkg/director/s3_upload.go index 15fa2a0c..70a8de1d 100644 --- a/pkg/director/s3_upload.go +++ b/pkg/director/s3_upload.go @@ -21,7 +21,7 @@ func (ss *StreamSession) maybeStartS3Upload(ctx context.Context, repoDID string) Region: ss.cli.S3Region, } keyPrefix := repoDID + "/" - ss.s3Uploader = s3.NewS3Uploader(cfg, keyPrefix, time.Minute) + ss.s3Uploader = s3.NewS3Uploader(cfg, repoDID, keyPrefix, time.Minute, ss.statefulDB) log.Log(ctx, "S3 upload enabled", "bucket", ss.cli.S3Bucket, "endpoint", ss.cli.S3Endpoint) } diff --git a/pkg/s3/s3.go b/pkg/s3/s3.go index 2bd32551..41507211 100644 --- a/pkg/s3/s3.go +++ b/pkg/s3/s3.go @@ -23,6 +23,15 @@ type Config struct { Region string } +// Recorder is an optional persistence hook for S3Uploader. RecordStart is +// called when a new multipart upload begins; the returned id is passed back +// to RecordComplete when the upload is finalized. Implementations should +// tolerate nil contexts being passed. +type Recorder interface { + RecordStart(ctx context.Context, userDID, bucket, key string, started time.Time) (id string, err error) + RecordComplete(ctx context.Context, id string, parts int32, size int64) error +} + // S3Uploader manages streaming multipart uploads to an S3-compatible endpoint. // Full fMP4 archives are fed via AddSegment. They are run through a muxl // Concatenator to strip duplicate init segments, then uploaded as a @@ -33,26 +42,34 @@ type S3Uploader struct { bucket string cutoverEvery time.Duration keyPrefix string // e.g. "did:plc:abc123/" + userDID string concat *muxl.Concatenator done chan error + recorder Recorder } // S3 requires each part except the last to be at least 5MB. const minPartSize = 5 * 1024 * 1024 type activeUpload struct { - key string - uploadID string - parts []types.CompletedPart - partNum int32 - started time.Time - buf []byte // accumulates segments until we hit minPartSize + key string + uploadID string + recordID string // set by Recorder.RecordStart, used for RecordComplete + parts []types.CompletedPart + partNum int32 + started time.Time + buf []byte // accumulates segments until we hit minPartSize + totalSize int64 // running total of bytes flushed across all parts } +var DefaultCutoverEvery = 10 * time.Minute + // NewS3Uploader creates a new S3Uploader. keyPrefix is prepended to every -// object key (typically the streamer DID + "/"). Starts the muxl Concatenator -// and a background goroutine that reads processed segments and uploads them. -func NewS3Uploader(cfg Config, keyPrefix string, cutoverEvery time.Duration) *S3Uploader { +// object key (typically the streamer DID + "/"). userDID is passed through +// to the Recorder so uploads can be attributed to a user. recorder may be +// nil to disable persistence. Starts the muxl Concatenator and a background +// goroutine that reads processed segments and uploads them. +func NewS3Uploader(cfg Config, userDID, keyPrefix string, cutoverEvery time.Duration, recorder Recorder) *S3Uploader { ctx := context.Background() client := s3.New(s3.Options{ Region: cfg.Region, @@ -65,7 +82,7 @@ func NewS3Uploader(cfg Config, keyPrefix string, cutoverEvery time.Duration) *S3 UsePathStyle: true, }) if cutoverEvery == 0 { - cutoverEvery = time.Minute + cutoverEvery = DefaultCutoverEvery } concat := muxl.NewConcatenator(ctx) u := &S3Uploader{ @@ -73,8 +90,10 @@ func NewS3Uploader(cfg Config, keyPrefix string, cutoverEvery time.Duration) *S3 bucket: cfg.Bucket, cutoverEvery: cutoverEvery, keyPrefix: keyPrefix, + userDID: userDID, concat: concat, done: make(chan error, 1), + recorder: recorder, } go u.uploadLoop(ctx) return u @@ -123,6 +142,13 @@ func (u *S3Uploader) uploadLoop(ctx context.Context) { uploadID: *resp.UploadId, started: now, } + if u.recorder != nil { + id, recErr := u.recorder.RecordStart(ctx, u.userDID, u.bucket, key, now) + if recErr != nil { + log.Error(ctx, "recording S3 upload start", "key", key, "error", recErr) + } + current.recordID = id + } // Prepend init segment to the buffer so the file starts valid if initSeg != nil { current.buf = append(current.buf, initSeg...) @@ -153,6 +179,7 @@ func (u *S3Uploader) uploadLoop(ctx context.Context) { ETag: resp.ETag, PartNumber: aws.Int32(partNum), }) + current.totalSize += int64(len(current.buf)) current.buf = current.buf[:0] return nil } @@ -187,7 +214,12 @@ func (u *S3Uploader) uploadLoop(ctx context.Context) { if err != nil { return fmt.Errorf("completing multipart upload %s: %w", current.key, err) } - log.Log(ctx, "completed S3 multipart upload", "key", current.key, "parts", len(current.parts)) + log.Log(ctx, "completed S3 multipart upload", "key", current.key, "parts", len(current.parts), "size", current.totalSize) + if u.recorder != nil && current.recordID != "" { + if recErr := u.recorder.RecordComplete(ctx, current.recordID, int32(len(current.parts)), current.totalSize); recErr != nil { + log.Error(ctx, "recording S3 upload completion", "key", current.key, "error", recErr) + } + } current = nil return nil } diff --git a/pkg/statedb/s3_segment.go b/pkg/statedb/s3_segment.go new file mode 100644 index 00000000..8b18a166 --- /dev/null +++ b/pkg/statedb/s3_segment.go @@ -0,0 +1,73 @@ +package statedb + +import ( + "context" + "time" + + "github.com/google/uuid" + "gorm.io/gorm" +) + +type S3Segment struct { + ID string `gorm:"column:id;primarykey"` + RepoDID string `gorm:"column:user_did;index;not null"` + Bucket string `gorm:"column:bucket;not null"` + Key string `gorm:"column:key;not null"` + URL string `gorm:"column:url"` + StartedAt time.Time `gorm:"column:started_at"` + CompletedAt *time.Time `gorm:"column:completed_at"` + Size int64 `gorm:"column:size"` + PartCount int32 `gorm:"column:part_count"` + CreatedAt time.Time `gorm:"column:created_at"` + UpdatedAt time.Time `gorm:"column:updated_at"` +} + +func (s *S3Segment) TableName() string { + return "s3_segments" +} + +// RecordStart inserts a new S3Segment row at the start of a multipart upload +// and returns its ID. Implements s3.Recorder. +func (state *StatefulDB) RecordStart(ctx context.Context, repoDID, bucket, key string, started time.Time) (string, error) { + uu, err := uuid.NewV7() + if err != nil { + return "", err + } + seg := &S3Segment{ + ID: uu.String(), + RepoDID: repoDID, + Bucket: bucket, + Key: key, + StartedAt: started, + } + if err := state.DB.WithContext(ctx).Create(seg).Error; err != nil { + return "", err + } + return seg.ID, nil +} + +// RecordComplete marks an S3 multipart upload as completed and records the +// final part count and size. Implements s3.Recorder. +func (state *StatefulDB) RecordComplete(ctx context.Context, id string, parts int32, size int64) error { + now := time.Now().UTC() + return state.DB.WithContext(ctx).Model(&S3Segment{}). + Where("id = ?", id). + Updates(map[string]any{ + "completed_at": &now, + "size": size, + "part_count": parts, + }).Error +} + +// GetS3Segment fetches an S3Segment by ID. Returns (nil, nil) if not found. +func (state *StatefulDB) GetS3Segment(ctx context.Context, id string) (*S3Segment, error) { + var seg S3Segment + err := state.DB.WithContext(ctx).Where("id = ?", id).First(&seg).Error + if err != nil { + if err == gorm.ErrRecordNotFound { + return nil, nil + } + return nil, err + } + return &seg, nil +} diff --git a/pkg/statedb/statedb.go b/pkg/statedb/statedb.go index 14fbcba1..187455cc 100644 --- a/pkg/statedb/statedb.go +++ b/pkg/statedb/statedb.go @@ -55,6 +55,7 @@ var StatefulDBModels = []any{ ModerationAuditLog{}, Storage{}, BroadcastOrigin{}, + S3Segment{}, } var NoPostgresDatabaseCode = "3D000" -- 2.51.2