diff --git a/pkg/api/api.go b/pkg/api/api.go index 249a9afed..035cbcca7 100644 --- a/pkg/api/api.go +++ b/pkg/api/api.go @@ -35,6 +35,7 @@ import ( "stream.place/streamplace/pkg/director" apierrors "stream.place/streamplace/pkg/errors" "stream.place/streamplace/pkg/linking" + "stream.place/streamplace/pkg/localdb" "stream.place/streamplace/pkg/log" "stream.place/streamplace/pkg/media" "stream.place/streamplace/pkg/mist/mistconfig" @@ -56,6 +57,7 @@ type StreamplaceAPI struct { CLI *config.CLI Model model.Model StatefulDB *statedb.StatefulDB + LocalDB localdb.LocalDB Updater *Updater Signer *eip712.EIP712Signer Mimes map[string]string @@ -93,7 +95,7 @@ type WebsocketTracker struct { mu sync.RWMutex } -func MakeStreamplaceAPI(cli *config.CLI, mod model.Model, statefulDB *statedb.StatefulDB, noter notifications.FirebaseNotifier, mm *media.MediaManager, ms media.MediaSigner, bus *bus.Bus, atsync *atproto.ATProtoSynchronizer, d *director.Director, op *oatproxy.OATProxy) (*StreamplaceAPI, error) { +func MakeStreamplaceAPI(cli *config.CLI, mod model.Model, statefulDB *statedb.StatefulDB, noter notifications.FirebaseNotifier, mm *media.MediaManager, ms media.MediaSigner, bus *bus.Bus, atsync *atproto.ATProtoSynchronizer, d *director.Director, op *oatproxy.OATProxy, ldb localdb.LocalDB) (*StreamplaceAPI, error) { updater, err := PrepareUpdater(cli) if err != nil { return nil, err @@ -117,6 +119,7 @@ func MakeStreamplaceAPI(cli *config.CLI, mod model.Model, statefulDB *statedb.St sessionsLock: sync.RWMutex{}, rtmpSessions: make(map[string]*media.RTMPSession), rtmpSessionsLock: sync.Mutex{}, + LocalDB: ldb, } a.Mimes, err = updater.GetMimes() if err != nil { @@ -152,7 +155,7 @@ func (a *StreamplaceAPI) Handler(ctx context.Context) (http.Handler, error) { Recorder: metrics.NewRecorder(metrics.Config{}), }) var xrpc http.Handler - xrpc, err := spxrpc.NewServer(ctx, a.CLI, a.Model, a.StatefulDB, a.op, mdlw, a.ATSync, a.Bus) + xrpc, err := spxrpc.NewServer(ctx, a.CLI, a.Model, a.StatefulDB, a.op, mdlw, a.ATSync, a.Bus, a.LocalDB) if err != nil { return nil, err } @@ -203,8 +206,6 @@ func (a *StreamplaceAPI) Handler(ctx context.Context) (http.Handler, error) { addHandle(apiRouter, "GET", "/api/chat/:repoDID", a.HandleChat(ctx)) addHandle(apiRouter, "GET", "/api/websocket/:repoDID", a.HandleWebsocket(ctx)) addHandle(apiRouter, "GET", "/api/livestream/:repoDID", a.HandleLivestream(ctx)) - addHandle(apiRouter, "GET", "/api/segment/recent", a.HandleRecentSegments(ctx)) - addHandle(apiRouter, "GET", "/api/segment/recent/:repoDID", a.HandleUserRecentSegments(ctx)) addHandle(apiRouter, "GET", "/api/bluesky/resolve/:handle", a.HandleBlueskyResolve(ctx)) addHandle(apiRouter, "GET", "/api/view-count/:user", a.HandleViewCount(ctx)) addHandle(apiRouter, "GET", "/api/clip/:user/:file", a.HandleClip(ctx)) @@ -561,59 +562,6 @@ func (a *StreamplaceAPI) HandlePlayerEvent(ctx context.Context) httprouter.Handl } } -func (a *StreamplaceAPI) HandleRecentSegments(ctx context.Context) httprouter.Handle { - return func(w http.ResponseWriter, req *http.Request, params httprouter.Params) { - segs, err := a.Model.MostRecentSegments() - if err != nil { - apierrors.WriteHTTPInternalServerError(w, "could not get segments", err) - return - } - bs, err := json.Marshal(segs) - if err != nil { - apierrors.WriteHTTPInternalServerError(w, "could not marshal segments", err) - return - } - w.Header().Add("Content-Type", "application/json") - if _, err := w.Write(bs); err != nil { - log.Error(ctx, "error writing response", "error", err) - } - } -} - -func (a *StreamplaceAPI) HandleUserRecentSegments(ctx context.Context) httprouter.Handle { - return func(w http.ResponseWriter, req *http.Request, params httprouter.Params) { - user := params.ByName("repoDID") - if user == "" { - apierrors.WriteHTTPBadRequest(w, "user required", nil) - return - } - user, err := a.NormalizeUser(ctx, user) - if err != nil { - apierrors.WriteHTTPNotFound(w, "user not found", err) - return - } - seg, err := a.Model.LatestSegmentForUser(user) - if err != nil { - apierrors.WriteHTTPInternalServerError(w, "could not get segments", err) - return - } - streamplaceSeg, err := seg.ToStreamplaceSegment() - if err != nil { - apierrors.WriteHTTPInternalServerError(w, "could not convert segment to streamplace segment", err) - return - } - bs, err := json.Marshal(streamplaceSeg) - if err != nil { - apierrors.WriteHTTPInternalServerError(w, "could not marshal segments", err) - return - } - w.Header().Add("Content-Type", "application/json") - if _, err := w.Write(bs); err != nil { - log.Error(ctx, "error writing response", "error", err) - } - } -} - func (a *StreamplaceAPI) HandleViewCount(ctx context.Context) httprouter.Handle { return func(w http.ResponseWriter, req *http.Request, params httprouter.Params) { user := params.ByName("user") diff --git a/pkg/api/api_internal.go b/pkg/api/api_internal.go index ef33d5809..2c4380187 100644 --- a/pkg/api/api_internal.go +++ b/pkg/api/api_internal.go @@ -298,7 +298,7 @@ func (a *StreamplaceAPI) InternalHandler(ctx context.Context) (http.Handler, err errors.WriteHTTPBadRequest(w, "id required", nil) return } - segment, err := a.Model.GetSegment(id) + segment, err := a.LocalDB.GetSegment(id) if err != nil { errors.WriteHTTPBadRequest(w, err.Error(), err) return @@ -553,7 +553,7 @@ func (a *StreamplaceAPI) InternalHandler(ctx context.Context) (http.Handler, err } after := time.Now().Add(-time.Duration(secs) * time.Second) w.Header().Set("Content-Type", "video/mp4") - err = media.ClipUser(ctx, a.Model, a.CLI, user, w, nil, &after) + err = media.ClipUser(ctx, a.LocalDB, a.CLI, user, w, nil, &after) if err != nil { errors.WriteHTTPInternalServerError(w, "unable to clip user", err) return diff --git a/pkg/api/playback.go b/pkg/api/playback.go index f92c78031..95a30c10d 100644 --- a/pkg/api/playback.go +++ b/pkg/api/playback.go @@ -272,7 +272,7 @@ func (a *StreamplaceAPI) HandleThumbnailPlayback(ctx context.Context) httprouter errors.WriteHTTPNotFound(w, "user not found", err) return } - thumb, err := a.Model.LatestThumbnailForUser(user) + thumb, err := a.LocalDB.LatestThumbnailForUser(user) if err != nil { errors.WriteHTTPInternalServerError(w, "could not query thumbnail", err) return diff --git a/pkg/api/websocket.go b/pkg/api/websocket.go index a3f1a1a86..6373111f3 100644 --- a/pkg/api/websocket.go +++ b/pkg/api/websocket.go @@ -181,7 +181,7 @@ func (a *StreamplaceAPI) HandleWebsocket(ctx context.Context) httprouter.Handle }() go func() { - seg, err := a.Model.LatestSegmentForUser(repoDID) + seg, err := a.LocalDB.LatestSegmentForUser(repoDID) if err != nil { log.Error(ctx, "could not get replies", "error", err) return diff --git a/pkg/cmd/streamplace.go b/pkg/cmd/streamplace.go index 2aa178e3b..64a531aed 100644 --- a/pkg/cmd/streamplace.go +++ b/pkg/cmd/streamplace.go @@ -29,6 +29,7 @@ import ( "stream.place/streamplace/pkg/director" "stream.place/streamplace/pkg/gstinit" "stream.place/streamplace/pkg/iroh/generated/iroh_streamplace" + "stream.place/streamplace/pkg/localdb" "stream.place/streamplace/pkg/log" "stream.place/streamplace/pkg/media" "stream.place/streamplace/pkg/notifications" @@ -237,6 +238,11 @@ func start(build *config.BuildFlags, platformJobs []jobFunc) error { return fmt.Errorf("error creating streamplace dir at %s:%w", cli.DataDir, err) } + ldb, err := localdb.MakeDB(cli.LocalDBURL) + if err != nil { + return err + } + mod, err := model.MakeDB(cli.DataFilePath([]string{"index"})) if err != nil { return err @@ -291,7 +297,7 @@ func start(build *config.BuildFlags, platformJobs []jobFunc) error { return fmt.Errorf("failed to migrate: %w", err) } - mm, err := media.MakeMediaManager(ctx, &cli, signer, mod, b, atsync) + mm, err := media.MakeMediaManager(ctx, &cli, signer, mod, b, atsync, ldb) if err != nil { return err } @@ -380,8 +386,8 @@ func start(build *config.BuildFlags, platformJobs []jobFunc) error { ClientMetadata: clientMetadata, Public: cli.PublicOAuth, }) - d := director.NewDirector(mm, mod, &cli, b, op, state, replicator) - a, err := api.MakeStreamplaceAPI(&cli, mod, state, noter, mm, ms, b, atsync, d, op) + d := director.NewDirector(mm, mod, &cli, b, op, state, replicator, ldb) + a, err := api.MakeStreamplaceAPI(&cli, mod, state, noter, mm, ms, b, atsync, d, op, ldb) if err != nil { return err } @@ -446,11 +452,11 @@ func start(build *config.BuildFlags, platformJobs []jobFunc) error { }) group.Go(func() error { - return storage.StartSegmentCleaner(ctx, mod, &cli) + return storage.StartSegmentCleaner(ctx, ldb, &cli) }) group.Go(func() error { - return mod.StartSegmentCleaner(ctx) + return ldb.StartSegmentCleaner(ctx) }) group.Go(func() error { diff --git a/pkg/config/config.go b/pkg/config/config.go index 66ee59646..b7ece9b6e 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -56,6 +56,7 @@ type CLI struct { Build *BuildFlags DataDir string DBURL string + LocalDBURL string EthAccountAddr string EthKeystorePath string EthPassword string @@ -242,6 +243,8 @@ func (cli *CLI) NewFlagSet(name string) *flag.FlagSet { cli.StringSliceFlag(fs, &cli.AdminDIDs, "admin-dids", []string{}, "comma-separated list of DIDs that are authorized to modify branding and other admin operations") cli.StringSliceFlag(fs, &cli.Syndicate, "syndicate", []string{}, "list of DIDs that we should rebroadcast ('*' for everybody)") fs.BoolVar(&cli.PlayerTelemetry, "player-telemetry", true, "enable player telemetry") + fs.StringVar(&cli.LocalDBURL, "local-db-url", "sqlite://$SP_DATA_DIR/localdb.sqlite", "URL of the local database to use for storing local data") + cli.dataDirFlags = append(cli.dataDirFlags, &cli.LocalDBURL) fs.Bool("external-signing", true, "DEPRECATED, does nothing.") fs.Bool("insecure", false, "DEPRECATED, does nothing.") diff --git a/pkg/director/director.go b/pkg/director/director.go index 3bbf23106..c12d22bf5 100644 --- a/pkg/director/director.go +++ b/pkg/director/director.go @@ -9,6 +9,7 @@ import ( "golang.org/x/sync/errgroup" "stream.place/streamplace/pkg/bus" "stream.place/streamplace/pkg/config" + "stream.place/streamplace/pkg/localdb" "stream.place/streamplace/pkg/log" "stream.place/streamplace/pkg/media" "stream.place/streamplace/pkg/model" @@ -32,9 +33,10 @@ type Director struct { op *oatproxy.OATProxy statefulDB *statedb.StatefulDB replicator replication.Replicator + localDB localdb.LocalDB } -func NewDirector(mm *media.MediaManager, mod model.Model, cli *config.CLI, bus *bus.Bus, op *oatproxy.OATProxy, statefulDB *statedb.StatefulDB, replicator replication.Replicator) *Director { +func NewDirector(mm *media.MediaManager, mod model.Model, cli *config.CLI, bus *bus.Bus, op *oatproxy.OATProxy, statefulDB *statedb.StatefulDB, replicator replication.Replicator, ldb localdb.LocalDB) *Director { return &Director{ mm: mm, mod: mod, @@ -45,6 +47,7 @@ func NewDirector(mm *media.MediaManager, mod model.Model, cli *config.CLI, bus * op: op, statefulDB: statefulDB, replicator: replicator, + localDB: ldb, } } @@ -79,6 +82,7 @@ func (d *Director) Start(ctx context.Context) error { // Initialize notification channels (buffered size 1 for coalescing) statusUpdateChan: make(chan struct{}, 1), originUpdateChan: make(chan struct{}, 1), + localDB: d.localDB, } d.streamSessions[not.Segment.RepoDID] = ss g.Go(func() error { diff --git a/pkg/director/stream_session.go b/pkg/director/stream_session.go index 0f4732ff4..d34ba0991 100644 --- a/pkg/director/stream_session.go +++ b/pkg/director/stream_session.go @@ -20,6 +20,7 @@ import ( "stream.place/streamplace/pkg/bus" "stream.place/streamplace/pkg/config" "stream.place/streamplace/pkg/livepeer" + "stream.place/streamplace/pkg/localdb" "stream.place/streamplace/pkg/log" "stream.place/streamplace/pkg/media" "stream.place/streamplace/pkg/model" @@ -44,6 +45,7 @@ type StreamSession struct { lastStatus time.Time lastStatusCID *string lastOriginTime time.Time + localDB localdb.LocalDB // Channels for background workers statusUpdateChan chan struct{} // Signal to update status @@ -178,7 +180,7 @@ func (ss *StreamSession) NewSegment(ctx context.Context, notif *media.NewSegment aqt := aqtime.FromTime(notif.Segment.StartTime) ctx = log.WithLogValues(ctx, "segID", notif.Segment.ID, "repoDID", notif.Segment.RepoDID, "timestamp", aqt.FileSafeString()) notif.Segment.MediaData.Size = len(notif.Data) - err := ss.mod.CreateSegment(notif.Segment) + err := ss.localDB.CreateSegment(notif.Segment) if err != nil { return fmt.Errorf("could not add segment to database: %w", err) } @@ -292,7 +294,7 @@ func (ss *StreamSession) Thumbnail(ctx context.Context, repoDID string, not *med return nil } defer lock.Unlock() - oldThumb, err := ss.mod.LatestThumbnailForUser(not.Segment.RepoDID) + oldThumb, err := ss.localDB.LatestThumbnailForUser(not.Segment.RepoDID) if err != nil { return err } @@ -311,11 +313,11 @@ func (ss *StreamSession) Thumbnail(ctx context.Context, repoDID string, not *med if err != nil { return err } - thumb := &model.Thumbnail{ + thumb := &localdb.Thumbnail{ Format: "jpeg", SegmentID: not.Segment.ID, } - err = ss.mod.CreateThumbnail(thumb) + err = ss.localDB.CreateThumbnail(thumb) if err != nil { return err } diff --git a/pkg/localdb/localdb.go b/pkg/localdb/localdb.go new file mode 100644 index 000000000..0b186bc08 --- /dev/null +++ b/pkg/localdb/localdb.go @@ -0,0 +1,81 @@ +package localdb + +import ( + "context" + "fmt" + "strings" + "time" + + "gorm.io/driver/sqlite" + "gorm.io/gorm" + "gorm.io/plugin/prometheus" + "stream.place/streamplace/pkg/config" + "stream.place/streamplace/pkg/log" +) + +type LocalDB interface { + CreateSegment(segment *Segment) error + MostRecentSegments() ([]Segment, error) + LatestSegmentForUser(user string) (*Segment, error) + LatestSegmentsForUser(user string, limit int, before *time.Time, after *time.Time) ([]Segment, error) + FilterLiveRepoDIDs(repoDIDs []string) ([]string, error) + CreateThumbnail(thumb *Thumbnail) error + LatestThumbnailForUser(user string) (*Thumbnail, error) + GetSegment(id string) (*Segment, error) + GetExpiredSegments(ctx context.Context) ([]Segment, error) + DeleteSegment(ctx context.Context, id string) error + StartSegmentCleaner(ctx context.Context) error + SegmentCleaner(ctx context.Context) error +} + +type LocalDatabase struct { + DB *gorm.DB +} + +func MakeDB(dbURL string) (LocalDB, error) { + log.Log(context.Background(), "starting database", "dbURL", dbURL) + if strings.HasPrefix(dbURL, "sqlite://") { + dbURL = dbURL[len("sqlite://"):] + } else if dbURL != ":memory:" { + return nil, fmt.Errorf("unsupported database URL (most start with sqlite://): %s", dbURL) + } + dial := sqlite.Open(dbURL) + + db, err := gorm.Open(dial, &gorm.Config{ + SkipDefaultTransaction: true, + TranslateError: true, + Logger: config.GormLogger, + }) + if err != nil { + return nil, fmt.Errorf("error starting database: %w", err) + } + err = db.Exec("PRAGMA journal_mode=WAL;").Error + if err != nil { + return nil, fmt.Errorf("error setting journal mode: %w", err) + } + + err = db.Use(prometheus.New(prometheus.Config{ + DBName: "localdb", + RefreshInterval: 10, + StartServer: false, + })) + if err != nil { + return nil, fmt.Errorf("error using prometheus plugin: %w", err) + } + + sqlDB, err := db.DB() + if err != nil { + return nil, fmt.Errorf("error getting database: %w", err) + } + sqlDB.SetMaxOpenConns(1) + for _, model := range []any{ + Segment{}, + Thumbnail{}, + } { + err = db.AutoMigrate(model) + if err != nil { + return nil, err + } + } + return &LocalDatabase{DB: db}, nil +} diff --git a/pkg/localdb/segment.go b/pkg/localdb/segment.go new file mode 100644 index 000000000..3cfc810cd --- /dev/null +++ b/pkg/localdb/segment.go @@ -0,0 +1,410 @@ +package localdb + +import ( + "context" + "database/sql/driver" + "encoding/json" + "errors" + "fmt" + "time" + + "gorm.io/gorm" + "stream.place/streamplace/pkg/aqtime" + "stream.place/streamplace/pkg/log" + "stream.place/streamplace/pkg/streamplace" +) + +type SegmentMediadataVideo struct { + Width int `json:"width"` + Height int `json:"height"` + FPSNum int `json:"fpsNum"` + FPSDen int `json:"fpsDen"` + BFrames bool `json:"bframes"` +} + +type SegmentMediadataAudio struct { + Rate int `json:"rate"` + Channels int `json:"channels"` +} + +type SegmentMediaData struct { + Video []*SegmentMediadataVideo `json:"video"` + Audio []*SegmentMediadataAudio `json:"audio"` + Duration int64 `json:"duration"` + Size int `json:"size"` +} + +// Scan scan value into Jsonb, implements sql.Scanner interface +func (j *SegmentMediaData) Scan(value any) error { + bytes, ok := value.([]byte) + if !ok { + return errors.New(fmt.Sprint("Failed to unmarshal JSONB value:", value)) + } + + result := SegmentMediaData{} + err := json.Unmarshal(bytes, &result) + *j = SegmentMediaData(result) + return err +} + +// Value return json value, implement driver.Valuer interface +func (j SegmentMediaData) Value() (driver.Value, error) { + return json.Marshal(j) +} + +// ContentRights represents content rights and attribution information +type ContentRights struct { + CopyrightNotice *string `json:"copyrightNotice,omitempty"` + CopyrightYear *int64 `json:"copyrightYear,omitempty"` + Creator *string `json:"creator,omitempty"` + CreditLine *string `json:"creditLine,omitempty"` + License *string `json:"license,omitempty"` +} + +// Scan scan value into ContentRights, implements sql.Scanner interface +func (c *ContentRights) Scan(value any) error { + if value == nil { + *c = ContentRights{} + return nil + } + bytes, ok := value.([]byte) + if !ok { + return errors.New(fmt.Sprint("Failed to unmarshal ContentRights value:", value)) + } + + result := ContentRights{} + err := json.Unmarshal(bytes, &result) + *c = ContentRights(result) + return err +} + +// Value return json value, implement driver.Valuer interface +func (c ContentRights) Value() (driver.Value, error) { + return json.Marshal(c) +} + +// DistributionPolicy represents distribution policy information +type DistributionPolicy struct { + DeleteAfterSeconds *int64 `json:"deleteAfterSeconds,omitempty"` +} + +// Scan scan value into DistributionPolicy, implements sql.Scanner interface +func (d *DistributionPolicy) Scan(value any) error { + if value == nil { + *d = DistributionPolicy{} + return nil + } + bytes, ok := value.([]byte) + if !ok { + return errors.New(fmt.Sprint("Failed to unmarshal DistributionPolicy value:", value)) + } + + result := DistributionPolicy{} + err := json.Unmarshal(bytes, &result) + *d = DistributionPolicy(result) + return err +} + +// Value return json value, implement driver.Valuer interface +func (d DistributionPolicy) Value() (driver.Value, error) { + return json.Marshal(d) +} + +// ContentWarningsSlice is a custom type for storing content warnings as JSON in the database +type ContentWarningsSlice []string + +// Scan scan value into ContentWarningsSlice, implements sql.Scanner interface +func (c *ContentWarningsSlice) Scan(value any) error { + if value == nil { + *c = ContentWarningsSlice{} + return nil + } + bytes, ok := value.([]byte) + if !ok { + return errors.New(fmt.Sprint("Failed to unmarshal ContentWarningsSlice value:", value)) + } + + result := ContentWarningsSlice{} + err := json.Unmarshal(bytes, &result) + *c = ContentWarningsSlice(result) + return err +} + +// Value return json value, implement driver.Valuer interface +func (c ContentWarningsSlice) Value() (driver.Value, error) { + return json.Marshal(c) +} + +type Segment struct { + ID string `json:"id" gorm:"primaryKey"` + SigningKeyDID string `json:"signingKeyDID" gorm:"column:signing_key_did"` + StartTime time.Time `json:"startTime" gorm:"index:latest_segments,priority:2;index:start_time"` + RepoDID string `json:"repoDID" gorm:"index:latest_segments,priority:1;column:repo_did"` + Title string `json:"title"` + Size int `json:"size" gorm:"column:size"` + MediaData *SegmentMediaData `json:"mediaData,omitempty"` + ContentWarnings ContentWarningsSlice `json:"contentWarnings,omitempty"` + ContentRights *ContentRights `json:"contentRights,omitempty"` + DistributionPolicy *DistributionPolicy `json:"distributionPolicy,omitempty"` + DeleteAfter *time.Time `json:"deleteAfter,omitempty" gorm:"column:delete_after;index:delete_after"` +} + +func (s *Segment) ToStreamplaceSegment() (*streamplace.Segment, error) { + aqt := aqtime.FromTime(s.StartTime) + if s.MediaData == nil { + return nil, fmt.Errorf("media data is nil") + } + if len(s.MediaData.Video) == 0 || s.MediaData.Video[0] == nil { + return nil, fmt.Errorf("video data is nil") + } + if len(s.MediaData.Audio) == 0 || s.MediaData.Audio[0] == nil { + return nil, fmt.Errorf("audio data is nil") + } + duration := s.MediaData.Duration + sizei64 := int64(s.Size) + + // Convert model metadata to streamplace metadata + var contentRights *streamplace.MetadataContentRights + if s.ContentRights != nil { + contentRights = &streamplace.MetadataContentRights{ + CopyrightNotice: s.ContentRights.CopyrightNotice, + CopyrightYear: s.ContentRights.CopyrightYear, + Creator: s.ContentRights.Creator, + CreditLine: s.ContentRights.CreditLine, + License: s.ContentRights.License, + } + } + + var contentWarnings *streamplace.MetadataContentWarnings + if len(s.ContentWarnings) > 0 { + contentWarnings = &streamplace.MetadataContentWarnings{ + Warnings: []string(s.ContentWarnings), + } + } + + var distributionPolicy *streamplace.MetadataDistributionPolicy + if s.DistributionPolicy != nil && s.DistributionPolicy.DeleteAfterSeconds != nil { + distributionPolicy = &streamplace.MetadataDistributionPolicy{ + DeleteAfter: s.DistributionPolicy.DeleteAfterSeconds, + } + } + + return &streamplace.Segment{ + LexiconTypeID: "place.stream.segment", + Creator: s.RepoDID, + Id: s.ID, + SigningKey: s.SigningKeyDID, + StartTime: string(aqt), + Duration: &duration, + Size: &sizei64, + ContentRights: contentRights, + ContentWarnings: contentWarnings, + DistributionPolicy: distributionPolicy, + Video: []*streamplace.Segment_Video{ + { + Codec: "h264", + Width: int64(s.MediaData.Video[0].Width), + Height: int64(s.MediaData.Video[0].Height), + Framerate: &streamplace.Segment_Framerate{ + Num: int64(s.MediaData.Video[0].FPSNum), + Den: int64(s.MediaData.Video[0].FPSDen), + }, + Bframes: &s.MediaData.Video[0].BFrames, + }, + }, + Audio: []*streamplace.Segment_Audio{ + { + Codec: "opus", + Rate: int64(s.MediaData.Audio[0].Rate), + Channels: int64(s.MediaData.Audio[0].Channels), + }, + }, + }, nil +} + +func (m *LocalDatabase) CreateSegment(seg *Segment) error { + err := m.DB.Model(Segment{}).Create(seg).Error + if err != nil { + return err + } + return nil +} + +// should return the most recent segment for each user, ordered by most recent first +// only includes segments from the last 30 seconds +func (m *LocalDatabase) MostRecentSegments() ([]Segment, error) { + var segments []Segment + thirtySecondsAgo := time.Now().Add(-30 * time.Second) + + err := m.DB.Table("segments"). + Select("segments.*"). + Where("start_time > ?", thirtySecondsAgo.UTC()). + Order("start_time DESC"). + Find(&segments).Error + if err != nil { + return nil, err + } + if segments == nil { + return []Segment{}, nil + } + + segmentMap := make(map[string]Segment) + for _, seg := range segments { + prev, ok := segmentMap[seg.RepoDID] + if !ok { + segmentMap[seg.RepoDID] = seg + } else { + if seg.StartTime.After(prev.StartTime) { + segmentMap[seg.RepoDID] = seg + } + } + } + + filteredSegments := []Segment{} + for _, seg := range segmentMap { + filteredSegments = append(filteredSegments, seg) + } + + return filteredSegments, nil +} + +func (m *LocalDatabase) LatestSegmentForUser(user string) (*Segment, error) { + var seg Segment + err := m.DB.Model(Segment{}).Where("repo_did = ?", user).Order("start_time DESC").First(&seg).Error + if err != nil { + return nil, err + } + return &seg, nil +} + +func (m *LocalDatabase) FilterLiveRepoDIDs(repoDIDs []string) ([]string, error) { + if len(repoDIDs) == 0 { + return []string{}, nil + } + + thirtySecondsAgo := time.Now().Add(-30 * time.Second) + + var liveDIDs []string + + err := m.DB.Table("segments"). + Select("DISTINCT repo_did"). + Where("repo_did IN ? AND start_time > ?", repoDIDs, thirtySecondsAgo.UTC()). + Pluck("repo_did", &liveDIDs).Error + + if err != nil { + return nil, err + } + + return liveDIDs, nil +} + +func (m *LocalDatabase) LatestSegmentsForUser(user string, limit int, before *time.Time, after *time.Time) ([]Segment, error) { + var segs []Segment + if before == nil { + later := time.Now().Add(1000 * time.Hour) + before = &later + } + if after == nil { + earlier := time.Time{} + after = &earlier + } + err := m.DB.Model(Segment{}).Where("repo_did = ? AND start_time < ? AND start_time > ?", user, before.UTC(), after.UTC()).Order("start_time DESC").Limit(limit).Find(&segs).Error + if err != nil { + return nil, err + } + return segs, nil +} + +func (m *LocalDatabase) GetSegment(id string) (*Segment, error) { + var seg Segment + + err := m.DB.Model(&Segment{}). + Preload("Repo"). + Where("id = ?", id). + First(&seg).Error + + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + if err != nil { + return nil, err + } + + return &seg, nil +} + +func (m *LocalDatabase) GetExpiredSegments(ctx context.Context) ([]Segment, error) { + + var expiredSegments []Segment + now := time.Now() + err := m.DB. + Where("delete_after IS NOT NULL AND delete_after < ?", now.UTC()). + Find(&expiredSegments).Error + if err != nil { + return nil, err + } + + return expiredSegments, nil +} + +func (m *LocalDatabase) DeleteSegment(ctx context.Context, id string) error { + return m.DB.Delete(&Segment{}, "id = ?", id).Error +} + +func (m *LocalDatabase) StartSegmentCleaner(ctx context.Context) error { + err := m.SegmentCleaner(ctx) + if err != nil { + return err + } + ticker := time.NewTicker(1 * time.Minute) + defer ticker.Stop() + + for { + select { + case <-ctx.Done(): + return nil + case <-ticker.C: + err := m.SegmentCleaner(ctx) + if err != nil { + log.Error(ctx, "Failed to clean segments", "error", err) + } + } + } +} + +func (m *LocalDatabase) SegmentCleaner(ctx context.Context) error { + // Calculate the cutoff time (10 minutes ago) + cutoffTime := aqtime.FromTime(time.Now().Add(-10 * time.Minute)).Time() + + // Find all unique repo_did values + var repoDIDs []string + if err := m.DB.Model(&Segment{}).Distinct("repo_did").Pluck("repo_did", &repoDIDs).Error; err != nil { + log.Error(ctx, "Failed to get unique repo_dids for segment cleaning", "error", err) + return err + } + + // For each user, keep their last 10 segments and delete older ones + for _, repoDID := range repoDIDs { + // Get IDs of the last 10 segments for this user + var keepSegmentIDs []string + if err := m.DB.Model(&Segment{}). + Where("repo_did = ?", repoDID). + Order("start_time DESC"). + Limit(10). + Pluck("id", &keepSegmentIDs).Error; err != nil { + log.Error(ctx, "Failed to get segment IDs to keep", "repo_did", repoDID, "error", err) + return err + } + + // Delete old segments except the ones we want to keep + result := m.DB.Where("repo_did = ? AND start_time < ? AND id NOT IN ?", + repoDID, cutoffTime, keepSegmentIDs).Delete(&Segment{}) + + if result.Error != nil { + log.Error(ctx, "Failed to clean old segments", "repo_did", repoDID, "error", result.Error) + } else if result.RowsAffected > 0 { + log.Log(ctx, "Cleaned old segments", "repo_did", repoDID, "count", result.RowsAffected) + } + } + return nil +} diff --git a/pkg/localdb/segment_test.go b/pkg/localdb/segment_test.go new file mode 100644 index 000000000..59a3becb3 --- /dev/null +++ b/pkg/localdb/segment_test.go @@ -0,0 +1,59 @@ +package localdb + +import ( + "fmt" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" + "stream.place/streamplace/pkg/config" +) + +func TestSegmentPerf(t *testing.T) { + config.DisableSQLLogging() + // dburl := filepath.Join(t.TempDir(), "test.db") + db, err := MakeDB(":memory:") + require.NoError(t, err) + // Create a ldb instance + ldb := db.(*LocalDatabase) + t.Cleanup(func() { + // os.Remove(dburl) + }) + + defer config.EnableSQLLogging() + // Create 250000 segments with timestamps 1 hour ago, each one second apart + wg := sync.WaitGroup{} + segCount := 250000 + wg.Add(segCount) + baseTime := time.Now() + for i := 0; i < segCount; i++ { + segment := &Segment{ + ID: fmt.Sprintf("segment-%d", i), + RepoDID: "did:plc:test123", + StartTime: baseTime.Add(-time.Duration(i) * time.Second).UTC(), + } + go func() { + defer wg.Done() + err = ldb.DB.Create(segment).Error + require.NoError(t, err) + }() + } + wg.Wait() + + startTime := time.Now() + wg = sync.WaitGroup{} + runs := 1000 + wg.Add(runs) + for i := 0; i < runs; i++ { + go func() { + defer wg.Done() + _, err := ldb.MostRecentSegments() + require.NoError(t, err) + // require.Len(t, segments, 1) + }() + } + wg.Wait() + fmt.Printf("Time taken: %s\n", time.Since(startTime)) + require.Less(t, time.Since(startTime), 10*time.Second) +} diff --git a/pkg/localdb/thumbnail.go b/pkg/localdb/thumbnail.go new file mode 100644 index 000000000..0aac69b5b --- /dev/null +++ b/pkg/localdb/thumbnail.go @@ -0,0 +1,60 @@ +package localdb + +import ( + "fmt" + + "github.com/google/uuid" +) + +type Thumbnail struct { + ID string `json:"id" gorm:"primaryKey"` + Format string `json:"format"` + SegmentID string `json:"segmentId" gorm:"index"` + Segment Segment `json:"segment,omitempty" gorm:"foreignKey:SegmentID;references:id"` +} + +func (m *LocalDatabase) CreateThumbnail(thumb *Thumbnail) error { + uu, err := uuid.NewV7() + if err != nil { + return err + } + if thumb.SegmentID == "" { + return fmt.Errorf("segmentID is required") + } + thumb.ID = uu.String() + err = m.DB.Model(Thumbnail{}).Create(thumb).Error + if err != nil { + return err + } + return nil +} + +// return the most recent thumbnail for a user +func (m *LocalDatabase) LatestThumbnailForUser(user string) (*Thumbnail, error) { + var thumbnail Thumbnail + + res := m.DB.Table("thumbnails AS t"). + Select("t.*"). + Joins("JOIN segments AS s ON t.segment_id = s.id"). + Where("s.repo_did = ?", user). + Order("s.start_time DESC"). + Limit(1). + Scan(&thumbnail) + + if res.RowsAffected == 0 { + return nil, nil + } + if res.Error != nil { + return nil, res.Error + } + + var seg Segment + err := m.DB.First(&seg, "id = ?", thumbnail.SegmentID).Error + if err != nil { + return nil, fmt.Errorf("could not find segment for thumbnail SegmentID=%s", thumbnail.SegmentID) + } + + thumbnail.Segment = seg + + return &thumbnail, nil +} diff --git a/pkg/media/clip_user.go b/pkg/media/clip_user.go index b1ae8453c..51f34f3ea 100644 --- a/pkg/media/clip_user.go +++ b/pkg/media/clip_user.go @@ -10,11 +10,11 @@ import ( "stream.place/streamplace/pkg/aqtime" "stream.place/streamplace/pkg/config" - "stream.place/streamplace/pkg/model" + "stream.place/streamplace/pkg/localdb" ) -func ClipUser(ctx context.Context, mod model.Model, cli *config.CLI, user string, writer io.Writer, before *time.Time, after *time.Time) error { - segments, err := mod.LatestSegmentsForUser(user, -1, before, after) +func ClipUser(ctx context.Context, localDB localdb.LocalDB, cli *config.CLI, user string, writer io.Writer, before *time.Time, after *time.Time) error { + segments, err := localDB.LatestSegmentsForUser(user, -1, before, after) if err != nil { return fmt.Errorf("unable to get segments: %w", err) } diff --git a/pkg/media/media.go b/pkg/media/media.go index 4e437ac4d..30e45ee8b 100644 --- a/pkg/media/media.go +++ b/pkg/media/media.go @@ -21,6 +21,7 @@ import ( c2patypes "stream.place/streamplace/pkg/c2patypes" "stream.place/streamplace/pkg/config" "stream.place/streamplace/pkg/gstinit" + "stream.place/streamplace/pkg/localdb" "stream.place/streamplace/pkg/model" "stream.place/streamplace/pkg/streamplace" @@ -51,10 +52,11 @@ type MediaManager struct { atsync *atproto.ATProtoSynchronizer webrtcAPI *webrtc.API webrtcConfig webrtc.Configuration + localDB localdb.LocalDB } type NewSegmentNotification struct { - Segment *model.Segment + Segment *localdb.Segment Data []byte Metadata *SegmentMetadata Local bool @@ -65,7 +67,7 @@ func RunSelfTest(ctx context.Context) error { return SelfTest(ctx) } -func MakeMediaManager(ctx context.Context, cli *config.CLI, signer crypto.Signer, mod model.Model, bus *bus.Bus, atsync *atproto.ATProtoSynchronizer) (*MediaManager, error) { +func MakeMediaManager(ctx context.Context, cli *config.CLI, signer crypto.Signer, mod model.Model, bus *bus.Bus, atsync *atproto.ATProtoSynchronizer, ldb localdb.LocalDB) (*MediaManager, error) { gstinit.InitGST() err := SelfTest(ctx) if err != nil { @@ -127,6 +129,7 @@ func MakeMediaManager(ctx context.Context, cli *config.CLI, signer crypto.Signer atsync: atsync, webrtcAPI: api, webrtcConfig: config, + localDB: ldb, }, nil } @@ -190,8 +193,8 @@ type SegmentMetadata struct { Title string Creator string ContentWarnings []string - ContentRights *model.ContentRights - DistributionPolicy *model.DistributionPolicy + ContentRights *localdb.ContentRights + DistributionPolicy *localdb.DistributionPolicy MetadataConfiguration *streamplace.MetadataConfiguration Livestream *streamplace.Livestream } @@ -312,7 +315,7 @@ func extractContentWarnings(mani *c2patypes.Manifest) []string { } // extractContentRights extracts content rights from the C2PA manifest -func extractContentRights(mani *c2patypes.Manifest) *model.ContentRights { +func extractContentRights(mani *c2patypes.Manifest) *localdb.ContentRights { ass := findAssertion(mani, StreamplaceMetadata) if ass == nil { return nil @@ -323,7 +326,7 @@ func extractContentRights(mani *c2patypes.Manifest) *model.ContentRights { return nil } - rights := &model.ContentRights{} + rights := &localdb.ContentRights{} // Extract copyright notice if notice, ok := data["dc:rights"]; ok { @@ -375,7 +378,7 @@ func extractContentRights(mani *c2patypes.Manifest) *model.ContentRights { } // extractDistributionPolicy extracts distribution policy from the C2PA manifest -func extractDistributionPolicy(mani *c2patypes.Manifest, segmentStart aqtime.AQTime) *model.DistributionPolicy { +func extractDistributionPolicy(mani *c2patypes.Manifest, segmentStart aqtime.AQTime) *localdb.DistributionPolicy { metadataConfig := extractMetadataConfiguration(mani) if metadataConfig == nil { return nil @@ -392,7 +395,7 @@ func extractDistributionPolicy(mani *c2patypes.Manifest, segmentStart aqtime.AQT // deleteAfter contains an offset in seconds from creation time deleteAfterSeconds := *metadataConfig.DistributionPolicy.DeleteAfter - return &model.DistributionPolicy{ + return &localdb.DistributionPolicy{ DeleteAfterSeconds: &deleteAfterSeconds, } } diff --git a/pkg/media/media_data_parser.go b/pkg/media/media_data_parser.go index c83243e79..532aac2a0 100644 --- a/pkg/media/media_data_parser.go +++ b/pkg/media/media_data_parser.go @@ -13,15 +13,15 @@ import ( "github.com/go-gst/go-gst/gst" "github.com/go-gst/go-gst/gst/app" "go.opentelemetry.io/otel" + "stream.place/streamplace/pkg/localdb" "stream.place/streamplace/pkg/log" - "stream.place/streamplace/pkg/model" ) func padProbeEmpty(_ *gst.Pad, _ *gst.PadProbeInfo) gst.PadProbeReturn { return gst.PadProbeOK } -func ParseSegmentMediaData(ctx context.Context, mp4bs []byte) (*model.SegmentMediaData, error) { +func ParseSegmentMediaData(ctx context.Context, mp4bs []byte) (*localdb.SegmentMediaData, error) { ctx, span := otel.Tracer("signer").Start(ctx, "ParseSegmentMediaData") defer span.End() ctx = log.WithLogValues(ctx, "GStreamerFunc", "ParseSegmentMediaData") @@ -40,8 +40,8 @@ func ParseSegmentMediaData(ctx context.Context, mp4bs []byte) (*model.SegmentMed return nil, fmt.Errorf("error creating SegmentMetadata pipeline: %w", err) } - var videoMetadata *model.SegmentMediadataVideo - var audioMetadata *model.SegmentMediadataAudio + var videoMetadata *localdb.SegmentMediadataVideo + var audioMetadata *localdb.SegmentMediadataAudio appsrc, err := pipeline.GetElementByName("appsrc") if err != nil { @@ -118,7 +118,7 @@ func ParseSegmentMediaData(ctx context.Context, mp4bs []byte) (*model.SegmentMed name := structure.Name() if name[:5] == "video" { - videoMetadata = &model.SegmentMediadataVideo{} + videoMetadata = &localdb.SegmentMediadataVideo{} // Get some common video properties widthVal, _ := structure.GetValue("width") heightVal, _ := structure.GetValue("height") @@ -147,7 +147,7 @@ func ParseSegmentMediaData(ctx context.Context, mp4bs []byte) (*model.SegmentMed } if name[:5] == "audio" { - audioMetadata = &model.SegmentMediadataAudio{} + audioMetadata = &localdb.SegmentMediadataAudio{} // Get some common audio properties rateVal, _ := structure.GetValue("rate") channelsVal, _ := structure.GetValue("channels") @@ -275,9 +275,9 @@ func ParseSegmentMediaData(ctx context.Context, mp4bs []byte) (*model.SegmentMed videoMetadata.BFrames = hasBFrames - meta := &model.SegmentMediaData{ - Video: []*model.SegmentMediadataVideo{videoMetadata}, - Audio: []*model.SegmentMediadataAudio{audioMetadata}, + meta := &localdb.SegmentMediaData{ + Video: []*localdb.SegmentMediadataVideo{videoMetadata}, + Audio: []*localdb.SegmentMediadataAudio{audioMetadata}, } ok, dur := pipeline.QueryDuration(gst.FormatTime) diff --git a/pkg/media/media_test.go b/pkg/media/media_test.go index 0784a10e4..f5e189c02 100644 --- a/pkg/media/media_test.go +++ b/pkg/media/media_test.go @@ -11,6 +11,7 @@ import ( "stream.place/streamplace/pkg/bus" "stream.place/streamplace/pkg/config" ct "stream.place/streamplace/pkg/config/configtesting" + "stream.place/streamplace/pkg/localdb" "stream.place/streamplace/pkg/model" "stream.place/streamplace/pkg/statedb" ) @@ -24,6 +25,8 @@ func getFixture(name string) string { func getStaticTestMediaManager(t *testing.T) (*MediaManager, MediaSigner) { mod, err := model.MakeDB(":memory:") require.NoError(t, err) + ldb, err := localdb.MakeDB(":memory:") + require.NoError(t, err) // signer, err := c2pa.MakeStaticSigner(eip712test.KeyBytes) require.NoError(t, err) if err != nil { @@ -42,7 +45,7 @@ func getStaticTestMediaManager(t *testing.T) (*MediaManager, MediaSigner) { StatefulDB: statedb, Bus: bus.NewBus(), } - mm, err := MakeMediaManager(context.Background(), cli, nil, mod, bus.NewBus(), atsync) + mm, err := MakeMediaManager(context.Background(), cli, nil, mod, bus.NewBus(), atsync, ldb) require.NoError(t, err) // ms, err := MakeMediaSigner(context.Background(), cli, "test-person", signer) // require.NoError(t, err) diff --git a/pkg/media/validate.go b/pkg/media/validate.go index 61d9066af..6d15447a8 100644 --- a/pkg/media/validate.go +++ b/pkg/media/validate.go @@ -18,8 +18,8 @@ import ( "stream.place/streamplace/pkg/constants" "stream.place/streamplace/pkg/crypto/signers" "stream.place/streamplace/pkg/iroh/generated/iroh_streamplace" + "stream.place/streamplace/pkg/localdb" "stream.place/streamplace/pkg/log" - "stream.place/streamplace/pkg/model" ) type ManifestAndCert struct { @@ -47,7 +47,7 @@ func (mm *MediaManager) ValidateMP4(ctx context.Context, input io.Reader, local label := manifest.Label if label != nil && mm.model != nil { - oldSeg, err := mm.model.GetSegment(*label) + oldSeg, err := mm.localDB.GetSegment(*label) if err != nil { return fmt.Errorf("failed to get old segment: %w", err) } @@ -117,7 +117,7 @@ func (mm *MediaManager) ValidateMP4(ctx context.Context, input io.Reader, local expiryTime := meta.StartTime.Time().Add(time.Duration(*meta.DistributionPolicy.DeleteAfterSeconds) * time.Second) deleteAfter = &expiryTime } - seg := &model.Segment{ + seg := &localdb.Segment{ ID: *label, SigningKeyDID: signingKeyDID, RepoDID: repoDID, @@ -125,7 +125,7 @@ func (mm *MediaManager) ValidateMP4(ctx context.Context, input io.Reader, local Title: meta.Title, Size: len(buf), MediaData: mediaData, - ContentWarnings: model.ContentWarningsSlice(meta.ContentWarnings), + ContentWarnings: localdb.ContentWarningsSlice(meta.ContentWarnings), ContentRights: meta.ContentRights, DistributionPolicy: meta.DistributionPolicy, DeleteAfter: deleteAfter, @@ -205,7 +205,7 @@ func (mm *MediaManager) isWarningBlocked(warning string) bool { type ValidationResult struct { Pub *atcrypto.PublicKeyK256 Meta *SegmentMetadata - MediaData *model.SegmentMediaData + MediaData *localdb.SegmentMediaData Manifest *c2patypes.Manifest Cert string } diff --git a/pkg/model/model.go b/pkg/model/model.go index 88e60d43c..ab1bcd6c5 100644 --- a/pkg/model/model.go +++ b/pkg/model/model.go @@ -28,19 +28,6 @@ type Model interface { PlayerReport(playerID string) (map[string]any, error) ClearPlayerEvents() error - CreateSegment(segment *Segment) error - MostRecentSegments() ([]Segment, error) - LatestSegmentForUser(user string) (*Segment, error) - LatestSegmentsForUser(user string, limit int, before *time.Time, after *time.Time) ([]Segment, error) - FilterLiveRepoDIDs(repoDIDs []string) ([]string, error) - CreateThumbnail(thumb *Thumbnail) error - LatestThumbnailForUser(user string) (*Thumbnail, error) - GetSegment(id string) (*Segment, error) - GetExpiredSegments(ctx context.Context) ([]Segment, error) - DeleteSegment(ctx context.Context, id string) error - StartSegmentCleaner(ctx context.Context) error - SegmentCleaner(ctx context.Context) error - GetIdentity(id string) (*Identity, error) UpdateIdentity(ident *Identity) error @@ -178,8 +165,6 @@ func MakeDB(dbURL string) (Model, error) { sqlDB.SetMaxOpenConns(1) for _, model := range []any{ PlayerEvent{}, - Segment{}, - Thumbnail{}, Identity{}, Repo{}, SigningKey{}, diff --git a/pkg/model/segment.go b/pkg/model/segment.go index 51f2561ad..8b5379070 100644 --- a/pkg/model/segment.go +++ b/pkg/model/segment.go @@ -1,412 +1 @@ package model - -import ( - "context" - "database/sql/driver" - "encoding/json" - "errors" - "fmt" - "time" - - "gorm.io/gorm" - "stream.place/streamplace/pkg/aqtime" - "stream.place/streamplace/pkg/log" - "stream.place/streamplace/pkg/streamplace" -) - -type SegmentMediadataVideo struct { - Width int `json:"width"` - Height int `json:"height"` - FPSNum int `json:"fpsNum"` - FPSDen int `json:"fpsDen"` - BFrames bool `json:"bframes"` -} - -type SegmentMediadataAudio struct { - Rate int `json:"rate"` - Channels int `json:"channels"` -} - -type SegmentMediaData struct { - Video []*SegmentMediadataVideo `json:"video"` - Audio []*SegmentMediadataAudio `json:"audio"` - Duration int64 `json:"duration"` - Size int `json:"size"` -} - -// Scan scan value into Jsonb, implements sql.Scanner interface -func (j *SegmentMediaData) Scan(value any) error { - bytes, ok := value.([]byte) - if !ok { - return errors.New(fmt.Sprint("Failed to unmarshal JSONB value:", value)) - } - - result := SegmentMediaData{} - err := json.Unmarshal(bytes, &result) - *j = SegmentMediaData(result) - return err -} - -// Value return json value, implement driver.Valuer interface -func (j SegmentMediaData) Value() (driver.Value, error) { - return json.Marshal(j) -} - -// ContentRights represents content rights and attribution information -type ContentRights struct { - CopyrightNotice *string `json:"copyrightNotice,omitempty"` - CopyrightYear *int64 `json:"copyrightYear,omitempty"` - Creator *string `json:"creator,omitempty"` - CreditLine *string `json:"creditLine,omitempty"` - License *string `json:"license,omitempty"` -} - -// Scan scan value into ContentRights, implements sql.Scanner interface -func (c *ContentRights) Scan(value any) error { - if value == nil { - *c = ContentRights{} - return nil - } - bytes, ok := value.([]byte) - if !ok { - return errors.New(fmt.Sprint("Failed to unmarshal ContentRights value:", value)) - } - - result := ContentRights{} - err := json.Unmarshal(bytes, &result) - *c = ContentRights(result) - return err -} - -// Value return json value, implement driver.Valuer interface -func (c ContentRights) Value() (driver.Value, error) { - return json.Marshal(c) -} - -// DistributionPolicy represents distribution policy information -type DistributionPolicy struct { - DeleteAfterSeconds *int64 `json:"deleteAfterSeconds,omitempty"` -} - -// Scan scan value into DistributionPolicy, implements sql.Scanner interface -func (d *DistributionPolicy) Scan(value any) error { - if value == nil { - *d = DistributionPolicy{} - return nil - } - bytes, ok := value.([]byte) - if !ok { - return errors.New(fmt.Sprint("Failed to unmarshal DistributionPolicy value:", value)) - } - - result := DistributionPolicy{} - err := json.Unmarshal(bytes, &result) - *d = DistributionPolicy(result) - return err -} - -// Value return json value, implement driver.Valuer interface -func (d DistributionPolicy) Value() (driver.Value, error) { - return json.Marshal(d) -} - -// ContentWarningsSlice is a custom type for storing content warnings as JSON in the database -type ContentWarningsSlice []string - -// Scan scan value into ContentWarningsSlice, implements sql.Scanner interface -func (c *ContentWarningsSlice) Scan(value any) error { - if value == nil { - *c = ContentWarningsSlice{} - return nil - } - bytes, ok := value.([]byte) - if !ok { - return errors.New(fmt.Sprint("Failed to unmarshal ContentWarningsSlice value:", value)) - } - - result := ContentWarningsSlice{} - err := json.Unmarshal(bytes, &result) - *c = ContentWarningsSlice(result) - return err -} - -// Value return json value, implement driver.Valuer interface -func (c ContentWarningsSlice) Value() (driver.Value, error) { - return json.Marshal(c) -} - -type Segment struct { - ID string `json:"id" gorm:"primaryKey"` - SigningKeyDID string `json:"signingKeyDID" gorm:"column:signing_key_did"` - SigningKey *SigningKey `json:"signingKey,omitempty" gorm:"foreignKey:DID;references:SigningKeyDID"` - StartTime time.Time `json:"startTime" gorm:"index:latest_segments,priority:2;index:start_time"` - RepoDID string `json:"repoDID" gorm:"index:latest_segments,priority:1;column:repo_did"` - Repo *Repo `json:"repo,omitempty" gorm:"foreignKey:DID;references:RepoDID"` - Title string `json:"title"` - Size int `json:"size" gorm:"column:size"` - MediaData *SegmentMediaData `json:"mediaData,omitempty"` - ContentWarnings ContentWarningsSlice `json:"contentWarnings,omitempty"` - ContentRights *ContentRights `json:"contentRights,omitempty"` - DistributionPolicy *DistributionPolicy `json:"distributionPolicy,omitempty"` - DeleteAfter *time.Time `json:"deleteAfter,omitempty" gorm:"column:delete_after;index:delete_after"` -} - -func (s *Segment) ToStreamplaceSegment() (*streamplace.Segment, error) { - aqt := aqtime.FromTime(s.StartTime) - if s.MediaData == nil { - return nil, fmt.Errorf("media data is nil") - } - if len(s.MediaData.Video) == 0 || s.MediaData.Video[0] == nil { - return nil, fmt.Errorf("video data is nil") - } - if len(s.MediaData.Audio) == 0 || s.MediaData.Audio[0] == nil { - return nil, fmt.Errorf("audio data is nil") - } - duration := s.MediaData.Duration - sizei64 := int64(s.Size) - - // Convert model metadata to streamplace metadata - var contentRights *streamplace.MetadataContentRights - if s.ContentRights != nil { - contentRights = &streamplace.MetadataContentRights{ - CopyrightNotice: s.ContentRights.CopyrightNotice, - CopyrightYear: s.ContentRights.CopyrightYear, - Creator: s.ContentRights.Creator, - CreditLine: s.ContentRights.CreditLine, - License: s.ContentRights.License, - } - } - - var contentWarnings *streamplace.MetadataContentWarnings - if len(s.ContentWarnings) > 0 { - contentWarnings = &streamplace.MetadataContentWarnings{ - Warnings: []string(s.ContentWarnings), - } - } - - var distributionPolicy *streamplace.MetadataDistributionPolicy - if s.DistributionPolicy != nil && s.DistributionPolicy.DeleteAfterSeconds != nil { - distributionPolicy = &streamplace.MetadataDistributionPolicy{ - DeleteAfter: s.DistributionPolicy.DeleteAfterSeconds, - } - } - - return &streamplace.Segment{ - LexiconTypeID: "place.stream.segment", - Creator: s.RepoDID, - Id: s.ID, - SigningKey: s.SigningKeyDID, - StartTime: string(aqt), - Duration: &duration, - Size: &sizei64, - ContentRights: contentRights, - ContentWarnings: contentWarnings, - DistributionPolicy: distributionPolicy, - Video: []*streamplace.Segment_Video{ - { - Codec: "h264", - Width: int64(s.MediaData.Video[0].Width), - Height: int64(s.MediaData.Video[0].Height), - Framerate: &streamplace.Segment_Framerate{ - Num: int64(s.MediaData.Video[0].FPSNum), - Den: int64(s.MediaData.Video[0].FPSDen), - }, - Bframes: &s.MediaData.Video[0].BFrames, - }, - }, - Audio: []*streamplace.Segment_Audio{ - { - Codec: "opus", - Rate: int64(s.MediaData.Audio[0].Rate), - Channels: int64(s.MediaData.Audio[0].Channels), - }, - }, - }, nil -} - -func (m *DBModel) CreateSegment(seg *Segment) error { - err := m.DB.Model(Segment{}).Create(seg).Error - if err != nil { - return err - } - return nil -} - -// should return the most recent segment for each user, ordered by most recent first -// only includes segments from the last 30 seconds -func (m *DBModel) MostRecentSegments() ([]Segment, error) { - var segments []Segment - thirtySecondsAgo := time.Now().Add(-30 * time.Second) - - err := m.DB.Table("segments"). - Select("segments.*"). - Where("start_time > ?", thirtySecondsAgo.UTC()). - Order("start_time DESC"). - Find(&segments).Error - if err != nil { - return nil, err - } - if segments == nil { - return []Segment{}, nil - } - - segmentMap := make(map[string]Segment) - for _, seg := range segments { - prev, ok := segmentMap[seg.RepoDID] - if !ok { - segmentMap[seg.RepoDID] = seg - } else { - if seg.StartTime.After(prev.StartTime) { - segmentMap[seg.RepoDID] = seg - } - } - } - - filteredSegments := []Segment{} - for _, seg := range segmentMap { - filteredSegments = append(filteredSegments, seg) - } - - return filteredSegments, nil -} - -func (m *DBModel) LatestSegmentForUser(user string) (*Segment, error) { - var seg Segment - err := m.DB.Model(Segment{}).Where("repo_did = ?", user).Order("start_time DESC").First(&seg).Error - if err != nil { - return nil, err - } - return &seg, nil -} - -func (m *DBModel) FilterLiveRepoDIDs(repoDIDs []string) ([]string, error) { - if len(repoDIDs) == 0 { - return []string{}, nil - } - - thirtySecondsAgo := time.Now().Add(-30 * time.Second) - - var liveDIDs []string - - err := m.DB.Table("segments"). - Select("DISTINCT repo_did"). - Where("repo_did IN ? AND start_time > ?", repoDIDs, thirtySecondsAgo.UTC()). - Pluck("repo_did", &liveDIDs).Error - - if err != nil { - return nil, err - } - - return liveDIDs, nil -} - -func (m *DBModel) LatestSegmentsForUser(user string, limit int, before *time.Time, after *time.Time) ([]Segment, error) { - var segs []Segment - if before == nil { - later := time.Now().Add(1000 * time.Hour) - before = &later - } - if after == nil { - earlier := time.Time{} - after = &earlier - } - err := m.DB.Model(Segment{}).Where("repo_did = ? AND start_time < ? AND start_time > ?", user, before.UTC(), after.UTC()).Order("start_time DESC").Limit(limit).Find(&segs).Error - if err != nil { - return nil, err - } - return segs, nil -} - -func (m *DBModel) GetSegment(id string) (*Segment, error) { - var seg Segment - - err := m.DB.Model(&Segment{}). - Preload("Repo"). - Where("id = ?", id). - First(&seg).Error - - if errors.Is(err, gorm.ErrRecordNotFound) { - return nil, nil - } - if err != nil { - return nil, err - } - - return &seg, nil -} - -func (m *DBModel) GetExpiredSegments(ctx context.Context) ([]Segment, error) { - - var expiredSegments []Segment - now := time.Now() - err := m.DB. - Where("delete_after IS NOT NULL AND delete_after < ?", now.UTC()). - Find(&expiredSegments).Error - if err != nil { - return nil, err - } - - return expiredSegments, nil -} - -func (m *DBModel) DeleteSegment(ctx context.Context, id string) error { - return m.DB.Delete(&Segment{}, "id = ?", id).Error -} - -func (m *DBModel) StartSegmentCleaner(ctx context.Context) error { - err := m.SegmentCleaner(ctx) - if err != nil { - return err - } - ticker := time.NewTicker(1 * time.Minute) - defer ticker.Stop() - - for { - select { - case <-ctx.Done(): - return nil - case <-ticker.C: - err := m.SegmentCleaner(ctx) - if err != nil { - log.Error(ctx, "Failed to clean segments", "error", err) - } - } - } -} - -func (m *DBModel) SegmentCleaner(ctx context.Context) error { - // Calculate the cutoff time (10 minutes ago) - cutoffTime := aqtime.FromTime(time.Now().Add(-10 * time.Minute)).Time() - - // Find all unique repo_did values - var repoDIDs []string - if err := m.DB.Model(&Segment{}).Distinct("repo_did").Pluck("repo_did", &repoDIDs).Error; err != nil { - log.Error(ctx, "Failed to get unique repo_dids for segment cleaning", "error", err) - return err - } - - // For each user, keep their last 10 segments and delete older ones - for _, repoDID := range repoDIDs { - // Get IDs of the last 10 segments for this user - var keepSegmentIDs []string - if err := m.DB.Model(&Segment{}). - Where("repo_did = ?", repoDID). - Order("start_time DESC"). - Limit(10). - Pluck("id", &keepSegmentIDs).Error; err != nil { - log.Error(ctx, "Failed to get segment IDs to keep", "repo_did", repoDID, "error", err) - return err - } - - // Delete old segments except the ones we want to keep - result := m.DB.Where("repo_did = ? AND start_time < ? AND id NOT IN ?", - repoDID, cutoffTime, keepSegmentIDs).Delete(&Segment{}) - - if result.Error != nil { - log.Error(ctx, "Failed to clean old segments", "repo_did", repoDID, "error", result.Error) - } else if result.RowsAffected > 0 { - log.Log(ctx, "Cleaned old segments", "repo_did", repoDID, "count", result.RowsAffected) - } - } - return nil -} diff --git a/pkg/model/segment_test.go b/pkg/model/segment_test.go index ac2c047a5..8b5379070 100644 --- a/pkg/model/segment_test.go +++ b/pkg/model/segment_test.go @@ -1,66 +1 @@ package model - -import ( - "fmt" - "sync" - "testing" - "time" - - "github.com/stretchr/testify/require" - "stream.place/streamplace/pkg/config" -) - -func TestSegmentPerf(t *testing.T) { - config.DisableSQLLogging() - // dburl := filepath.Join(t.TempDir(), "test.db") - db, err := MakeDB(":memory:") - require.NoError(t, err) - // Create a model instance - model := db.(*DBModel) - t.Cleanup(func() { - // os.Remove(dburl) - }) - - // Create a repo for testing - repo := &Repo{ - DID: "did:plc:test123", - } - err = model.DB.Create(repo).Error - require.NoError(t, err) - - defer config.EnableSQLLogging() - // Create 250000 segments with timestamps 1 hour ago, each one second apart - wg := sync.WaitGroup{} - segCount := 250000 - wg.Add(segCount) - baseTime := time.Now() - for i := 0; i < segCount; i++ { - segment := &Segment{ - ID: fmt.Sprintf("segment-%d", i), - RepoDID: repo.DID, - StartTime: baseTime.Add(-time.Duration(i) * time.Second).UTC(), - } - go func() { - defer wg.Done() - err = model.DB.Create(segment).Error - require.NoError(t, err) - }() - } - wg.Wait() - - startTime := time.Now() - wg = sync.WaitGroup{} - runs := 1000 - wg.Add(runs) - for i := 0; i < runs; i++ { - go func() { - defer wg.Done() - _, err := model.MostRecentSegments() - require.NoError(t, err) - // require.Len(t, segments, 1) - }() - } - wg.Wait() - fmt.Printf("Time taken: %s\n", time.Since(startTime)) - require.Less(t, time.Since(startTime), 10*time.Second) -} diff --git a/pkg/model/thumbnail.go b/pkg/model/thumbnail.go index 924c9ab9d..8b5379070 100644 --- a/pkg/model/thumbnail.go +++ b/pkg/model/thumbnail.go @@ -1,60 +1 @@ package model - -import ( - "fmt" - - "github.com/google/uuid" -) - -type Thumbnail struct { - ID string `json:"id" gorm:"primaryKey"` - Format string `json:"format"` - SegmentID string `json:"segmentId" gorm:"index"` - Segment Segment `json:"segment,omitempty" gorm:"foreignKey:SegmentID;references:id"` -} - -func (m *DBModel) CreateThumbnail(thumb *Thumbnail) error { - uu, err := uuid.NewV7() - if err != nil { - return err - } - if thumb.SegmentID == "" { - return fmt.Errorf("segmentID is required") - } - thumb.ID = uu.String() - err = m.DB.Model(Thumbnail{}).Create(thumb).Error - if err != nil { - return err - } - return nil -} - -// return the most recent thumbnail for a user -func (m *DBModel) LatestThumbnailForUser(user string) (*Thumbnail, error) { - var thumbnail Thumbnail - - res := m.DB.Table("thumbnails AS t"). - Select("t.*"). - Joins("JOIN segments AS s ON t.segment_id = s.id"). - Where("s.repo_did = ?", user). - Order("s.start_time DESC"). - Limit(1). - Scan(&thumbnail) - - if res.RowsAffected == 0 { - return nil, nil - } - if res.Error != nil { - return nil, res.Error - } - - var seg Segment - err := m.DB.First(&seg, "id = ?", thumbnail.SegmentID).Error - if err != nil { - return nil, fmt.Errorf("could not find segment for thumbnail SegmentID=%s", thumbnail.SegmentID) - } - - thumbnail.Segment = seg - - return &thumbnail, nil -} diff --git a/pkg/spxrpc/app_bsky_feed.go b/pkg/spxrpc/app_bsky_feed.go index 517a6412c..4a3d258f1 100644 --- a/pkg/spxrpc/app_bsky_feed.go +++ b/pkg/spxrpc/app_bsky_feed.go @@ -56,7 +56,7 @@ func (s *Server) handleAppBskyFeedGetFeedSkeleton(ctx context.Context, inCursor outCursor = fmt.Sprintf("%d::%s", ts, last.CID) } } else if name == FeedLiveStreams { - segs, err := s.model.MostRecentSegments() + segs, err := s.localDB.MostRecentSegments() if err != nil { return nil, echo.NewHTTPError(http.StatusInternalServerError, fmt.Sprintf("failed to get recent segments: %v", err)) } diff --git a/pkg/spxrpc/com_atproto_moderation.go b/pkg/spxrpc/com_atproto_moderation.go index f1cab7c7a..f017909f3 100644 --- a/pkg/spxrpc/com_atproto_moderation.go +++ b/pkg/spxrpc/com_atproto_moderation.go @@ -13,9 +13,9 @@ import ( "github.com/labstack/echo/v4" "github.com/streamplace/oatproxy/pkg/oatproxy" "stream.place/streamplace/pkg/config" + "stream.place/streamplace/pkg/localdb" "stream.place/streamplace/pkg/log" "stream.place/streamplace/pkg/media" - "stream.place/streamplace/pkg/model" ) func (s *Server) handleComAtprotoModerationCreateReport(ctx context.Context, body *comatprototypes.ModerationCreateReport_Input) (*comatprototypes.ModerationCreateReport_Output, error) { @@ -76,7 +76,7 @@ func (s *Server) handleComAtprotoModerationCreateReport(ctx context.Context, bod return nil, echo.NewHTTPError(http.StatusBadRequest, "invalid subject") } - clipID, err := makeClip(ctx, s.cli, s.model, did) + clipID, err := makeClip(ctx, s.cli, s.localDB, did) if err != nil { // we still want the report to go through! log.Error(ctx, "failed to make clip for report", "error", err) @@ -99,7 +99,7 @@ func (s *Server) handleComAtprotoModerationCreateReport(ctx context.Context, bod return &output, nil } -func makeClip(ctx context.Context, cli *config.CLI, mod model.Model, did string) (string, error) { +func makeClip(ctx context.Context, cli *config.CLI, localDB localdb.LocalDB, did string) (string, error) { after := time.Now().Add(-time.Duration(60) * time.Second) uu, err := uuid.NewV7() @@ -113,7 +113,7 @@ func makeClip(ctx context.Context, cli *config.CLI, mod model.Model, did string) } defer fd.Close() - err = media.ClipUser(ctx, mod, cli, did, fd, nil, &after) + err = media.ClipUser(ctx, localDB, cli, did, fd, nil, &after) if err != nil { return "", echo.NewHTTPError(http.StatusInternalServerError, "failed to clip user") } diff --git a/pkg/spxrpc/place_stream_live.go b/pkg/spxrpc/place_stream_live.go index f2e2e0a0e..dcc1ee3c7 100644 --- a/pkg/spxrpc/place_stream_live.go +++ b/pkg/spxrpc/place_stream_live.go @@ -82,7 +82,7 @@ func (s *Server) handlePlaceStreamLiveGetSegments(ctx context.Context, before st beforeTime = &parsedTime } - segments, err := s.model.LatestSegmentsForUser(userDID, limit, beforeTime, nil) + segments, err := s.localDB.LatestSegmentsForUser(userDID, limit, beforeTime, nil) if err != nil { return nil, echo.NewHTTPError(http.StatusInternalServerError, "Failed to fetch segments") } @@ -223,7 +223,7 @@ func (s *Server) handlePlaceStreamLiveGetRecommendations(ctx context.Context, us } // Filter for only live streamers - liveStreamers, err := s.model.FilterLiveRepoDIDs(streamers) + liveStreamers, err := s.localDB.FilterLiveRepoDIDs(streamers) if err != nil { return nil, echo.NewHTTPError(http.StatusInternalServerError, "Failed to filter live streamers") } @@ -256,7 +256,7 @@ func (s *Server) handlePlaceStreamLiveGetRecommendations(ctx context.Context, us followDIDs[i] = follow.SubjectDID } - liveFollows, err := s.model.FilterLiveRepoDIDs(followDIDs) + liveFollows, err := s.localDB.FilterLiveRepoDIDs(followDIDs) if err != nil { return nil, echo.NewHTTPError(http.StatusInternalServerError, "Failed to filter live follows") } @@ -281,7 +281,7 @@ func (s *Server) handlePlaceStreamLiveGetRecommendations(ctx context.Context, us // Final fallback: use host's default recommendations defaultStreamers := s.cli.DefaultRecommendedStreamers if len(defaultStreamers) > 0 { - liveDefaults, err := s.model.FilterLiveRepoDIDs(defaultStreamers) + liveDefaults, err := s.localDB.FilterLiveRepoDIDs(defaultStreamers) if err != nil { return nil, echo.NewHTTPError(http.StatusInternalServerError, "Failed to filter default streamers") } diff --git a/pkg/spxrpc/spxrpc.go b/pkg/spxrpc/spxrpc.go index b14125770..68d0cff97 100644 --- a/pkg/spxrpc/spxrpc.go +++ b/pkg/spxrpc/spxrpc.go @@ -18,6 +18,7 @@ import ( "stream.place/streamplace/pkg/atproto" "stream.place/streamplace/pkg/bus" "stream.place/streamplace/pkg/config" + "stream.place/streamplace/pkg/localdb" "stream.place/streamplace/pkg/log" "stream.place/streamplace/pkg/model" "stream.place/streamplace/pkg/statedb" @@ -33,9 +34,10 @@ type Server struct { statefulDB *statedb.StatefulDB bus *bus.Bus op *oatproxy.OATProxy + localDB localdb.LocalDB } -func NewServer(ctx context.Context, cli *config.CLI, model model.Model, statefulDB *statedb.StatefulDB, op *oatproxy.OATProxy, mdlw middleware.Middleware, atsync *atproto.ATProtoSynchronizer, bus *bus.Bus) (*Server, error) { +func NewServer(ctx context.Context, cli *config.CLI, model model.Model, statefulDB *statedb.StatefulDB, op *oatproxy.OATProxy, mdlw middleware.Middleware, atsync *atproto.ATProtoSynchronizer, bus *bus.Bus, ldb localdb.LocalDB) (*Server, error) { e := echo.New() s := &Server{ e: e, @@ -47,6 +49,7 @@ func NewServer(ctx context.Context, cli *config.CLI, model model.Model, stateful statefulDB: statefulDB, bus: bus, op: op, + localDB: ldb, } e.Use(s.ErrorHandlingMiddleware()) e.Use(s.ContextPreservingMiddleware()) diff --git a/pkg/storage/storage.go b/pkg/storage/storage.go index 3771f6a44..fe28ec8b9 100644 --- a/pkg/storage/storage.go +++ b/pkg/storage/storage.go @@ -10,13 +10,13 @@ import ( "golang.org/x/sync/errgroup" "stream.place/streamplace/pkg/aqtime" "stream.place/streamplace/pkg/config" + "stream.place/streamplace/pkg/localdb" "stream.place/streamplace/pkg/log" - "stream.place/streamplace/pkg/model" ) const moderationRetention = 120 * time.Second -func StartSegmentCleaner(ctx context.Context, mod model.Model, cli *config.CLI) error { +func StartSegmentCleaner(ctx context.Context, localDB localdb.LocalDB, cli *config.CLI) error { ctx = log.WithLogValues(ctx, "func", "StartSegmentCleaner") g, ctx := errgroup.WithContext(ctx) g.Go(func() error { @@ -25,14 +25,14 @@ func StartSegmentCleaner(ctx context.Context, mod model.Model, cli *config.CLI) case <-ctx.Done(): return nil case <-time.After(60 * time.Second): - expiredSegments, err := mod.GetExpiredSegments(ctx) + expiredSegments, err := localDB.GetExpiredSegments(ctx) if err != nil { return err } log.Log(ctx, "Cleaning expired segments", "count", len(expiredSegments)) for _, seg := range expiredSegments { g.Go(func() error { - err := deleteSegment(ctx, mod, cli, seg) + err := deleteSegment(ctx, localDB, cli, seg) if err != nil { log.Error(ctx, "Failed to delete segment", "error", err) } @@ -47,7 +47,7 @@ func StartSegmentCleaner(ctx context.Context, mod model.Model, cli *config.CLI) return g.Wait() } -func deleteSegment(ctx context.Context, mod model.Model, cli *config.CLI, seg model.Segment) error { +func deleteSegment(ctx context.Context, localDB localdb.LocalDB, cli *config.CLI, seg localdb.Segment) error { if time.Since(seg.StartTime) < moderationRetention { log.Debug(ctx, "Skipping deletion of segment", "id", seg.ID, "time since start", time.Since(seg.StartTime)) return nil @@ -61,7 +61,7 @@ func deleteSegment(ctx context.Context, mod model.Model, cli *config.CLI, seg mo if err != nil && !errors.Is(err, os.ErrNotExist) { return err } - err = mod.DeleteSegment(ctx, seg.ID) + err = localDB.DeleteSegment(ctx, seg.ID) if err != nil { return err }