diff --git a/pkg/atproto/jwks.go b/pkg/atproto/jwks.go deleted file mode 100644 index ba4fbc1f2..000000000 --- a/pkg/atproto/jwks.go +++ /dev/null @@ -1,45 +0,0 @@ -package atproto - -import ( - "context" - "encoding/json" - "os" - - "github.com/lestrrat-go/jwx/v2/jwk" - oauth_helpers "github.com/streamplace/atproto-oauth-golang/helpers" - "stream.place/streamplace/pkg/log" -) - -func EnsureJWK(ctx context.Context, fPath string) (jwk.Key, error) { - var key jwk.Key - _, err := os.Stat(fPath) - if err == nil { - b, err := os.ReadFile(fPath) - if err != nil { - return nil, err - } - key, err = jwk.ParseKey(b) - if err != nil { - return nil, err - } - } else if os.IsNotExist(err) { - key, err = oauth_helpers.GenerateKey(nil) - if err != nil { - return nil, err - } - - b, err := json.Marshal(key) - if err != nil { - return nil, err - } - - if err := os.WriteFile(fPath, b, 0600); err != nil { - return nil, err - } - log.Log(ctx, "generated JWK", "path", fPath) - } else { - return nil, err - } - - return key, nil -} diff --git a/pkg/cmd/streamplace.go b/pkg/cmd/streamplace.go index 438ed5cb2..28b2dbad8 100644 --- a/pkg/cmd/streamplace.go +++ b/pkg/cmd/streamplace.go @@ -325,15 +325,13 @@ func start(build *config.BuildFlags, platformJobs []jobFunc) error { } } - jwkPath := cli.DataFilePath([]string{"jwk.json"}) - jwk, err := atproto.EnsureJWK(ctx, jwkPath) + jwk, err := statefulDB.EnsureJWK(ctx, "jwk") if err != nil { return err } cli.JWK = jwk - accessJWKPath := cli.DataFilePath([]string{"access-jwk.json"}) - accessJWK, err := atproto.EnsureJWK(ctx, accessJWKPath) + accessJWK, err := statefulDB.EnsureJWK(ctx, "access-jwk") if err != nil { return err } diff --git a/pkg/statedb/config.go b/pkg/statedb/config.go new file mode 100644 index 000000000..9722820db --- /dev/null +++ b/pkg/statedb/config.go @@ -0,0 +1,34 @@ +package statedb + +import ( + "errors" + "time" + + "gorm.io/gorm" +) + +type Config struct { + Key string `gorm:"column:key;primarykey"` + Value []byte `gorm:"column:value"` + CreatedAt time.Time `gorm:"column:created_at"` + UpdatedAt time.Time `gorm:"column:updated_at"` +} + +func (state *StatefulDB) GetConfig(key string) (*Config, error) { + var config Config + if err := state.DB.Where("key = ?", key).First(&config).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + return nil, err + } + return &config, nil +} + +func (state *StatefulDB) PutConfig(key string, value []byte) error { + config := Config{ + Key: key, + Value: value, + } + return state.DB.Save(&config).Error +} diff --git a/pkg/statedb/jwks.go b/pkg/statedb/jwks.go new file mode 100644 index 000000000..38fb490fa --- /dev/null +++ b/pkg/statedb/jwks.go @@ -0,0 +1,73 @@ +package statedb + +import ( + "context" + "encoding/json" + "fmt" + "os" + + "github.com/lestrrat-go/jwx/v2/jwk" + oauth_helpers "github.com/streamplace/atproto-oauth-golang/helpers" + "stream.place/streamplace/pkg/log" +) + +func (state *StatefulDB) EnsureJWK(ctx context.Context, name string) (jwk.Key, error) { + var key jwk.Key + + conf, err := state.GetConfig(name) + if err != nil { + return nil, err + } + + // happy path: we found the jwk in the database, use that + if conf != nil { + key, err = jwk.ParseKey(conf.Value) + if err != nil { + return nil, err + } + return key, nil + } + + // migration path: maybe we have an old one on disk. + key, _ = state.getOldJWK(ctx, name) + + // new path: found neither, generate a new one + if key == nil { + log.Warn(ctx, "no JWK found, generating new one", "name", name) + key, err = oauth_helpers.GenerateKey(nil) + if err != nil { + return nil, fmt.Errorf("failed to generate JWK: %w", err) + } + } + + b, err := json.Marshal(key) + if err != nil { + return nil, fmt.Errorf("failed to marshal JWK: %w", err) + } + err = state.PutConfig(name, b) + if err != nil { + return nil, fmt.Errorf("failed to save JWK: %w", err) + } + + return key, nil +} + +// migration for the old one we stored on disk +func (state *StatefulDB) getOldJWK(ctx context.Context, name string) (jwk.Key, error) { + var key jwk.Key + jwkPath := state.CLI.DataFilePath([]string{name + ".json"}) + _, err := os.Stat(jwkPath) + if err == nil { + b, err := os.ReadFile(jwkPath) + if err != nil { + return nil, err + } + key, err = jwk.ParseKey(b) + if err != nil { + return nil, err + } + log.Warn(ctx, "found old JWK on disk, migrating to stateful database", "path", jwkPath) + return key, nil + } + return nil, nil +} diff --git a/pkg/statedb/notification.go b/pkg/statedb/notification.go index 162675eb5..448b17102 100644 --- a/pkg/statedb/notification.go +++ b/pkg/statedb/notification.go @@ -3,44 +3,41 @@ package statedb import ( "fmt" "time" - - "gorm.io/gorm" ) type Notification struct { - Token string `gorm:"primarykey"` - RepoDID string `json:"repoDID,omitempty" gorm:"column:repo_did;index"` - CreatedAt time.Time - UpdatedAt time.Time - DeletedAt gorm.DeletedAt `gorm:"index"` + Token string `gorm:"column:token;primarykey"` + RepoDID string `json:"repoDID,omitempty" gorm:"column:repo_did;index"` + CreatedAt time.Time `gorm:"column:created_at"` + UpdatedAt time.Time `gorm:"column:updated_at"` } -func (db *StatefulDB) CreateNotification(token string, repoDID string) error { +func (state *StatefulDB) CreateNotification(token string, repoDID string) error { not := Notification{ Token: token, } if repoDID != "" { not.RepoDID = repoDID } - err := db.DB.Save(¬).Error + err := state.DB.Save(¬).Error if err != nil { return err } return nil } -func (db *StatefulDB) ListNotifications() ([]Notification, error) { +func (state *StatefulDB) ListNotifications() ([]Notification, error) { nots := []Notification{} - err := db.DB.Find(¬s).Error + err := state.DB.Find(¬s).Error if err != nil { return nil, fmt.Errorf("error retrieving notifications: %w", err) } return nots, nil } -func (db *StatefulDB) ListUserNotifications(userDID string) ([]Notification, error) { +func (state *StatefulDB) ListUserNotifications(userDID string) ([]Notification, error) { nots := []Notification{} - err := db.DB.Where("repo_did = ?", userDID).Find(¬s).Error + err := state.DB.Where("repo_did = ?", userDID).Find(¬s).Error if err != nil { return nil, fmt.Errorf("error retrieving notifications: %w", err) } @@ -48,10 +45,10 @@ func (db *StatefulDB) ListUserNotifications(userDID string) ([]Notification, err } // todo fixme we don't have followers in this database -func (db *StatefulDB) GetFollowersNotificationTokens(userDID string) ([]string, error) { +func (state *StatefulDB) GetFollowersNotificationTokens(userDID string) ([]string, error) { var tokens []string - err := db.DB.Model(&Notification{}). + err := state.DB.Model(&Notification{}). Distinct("notifications.token"). Joins("JOIN follows ON follows.user_did = notifications.repo_did"). Where("follows.subject_did = ?", userDID). @@ -63,7 +60,7 @@ func (db *StatefulDB) GetFollowersNotificationTokens(userDID string) ([]string, } // also you prolly wanna get one for yourself - nots, err := db.ListUserNotifications(userDID) + nots, err := state.ListUserNotifications(userDID) if err != nil { return nil, fmt.Errorf("error retrieving user notifications: %w", err) } diff --git a/pkg/statedb/oauth_session.go b/pkg/statedb/oauth_session.go index 9b79c5a3e..bbd254059 100644 --- a/pkg/statedb/oauth_session.go +++ b/pkg/statedb/oauth_session.go @@ -7,13 +7,13 @@ import ( "gorm.io/gorm" ) -func (db *StatefulDB) CreateOAuthSession(id string, session *oatproxy.OAuthSession) error { - return db.DB.Create(session).Error +func (state *StatefulDB) CreateOAuthSession(id string, session *oatproxy.OAuthSession) error { + return state.DB.Create(session).Error } -func (db *StatefulDB) LoadOAuthSession(id string) (*oatproxy.OAuthSession, error) { +func (state *StatefulDB) LoadOAuthSession(id string) (*oatproxy.OAuthSession, error) { var session oatproxy.OAuthSession - if err := db.DB.Where("downstream_dpop_jkt = ?", id).First(&session).Error; err != nil { + if err := state.DB.Where("downstream_dpop_jkt = ?", id).First(&session).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, nil } @@ -22,8 +22,8 @@ func (db *StatefulDB) LoadOAuthSession(id string) (*oatproxy.OAuthSession, error return &session, nil } -func (db *StatefulDB) UpdateOAuthSession(id string, session *oatproxy.OAuthSession) error { - res := db.DB.Model(&oatproxy.OAuthSession{}).Where("downstream_dpop_jkt = ?", id).Updates(session) +func (state *StatefulDB) UpdateOAuthSession(id string, session *oatproxy.OAuthSession) error { + res := state.DB.Model(&oatproxy.OAuthSession{}).Where("downstream_dpop_jkt = ?", id).Updates(session) if res.Error != nil { return res.Error } @@ -33,17 +33,17 @@ func (db *StatefulDB) UpdateOAuthSession(id string, session *oatproxy.OAuthSessi return nil } -func (db *StatefulDB) ListOAuthSessions() ([]oatproxy.OAuthSession, error) { +func (state *StatefulDB) ListOAuthSessions() ([]oatproxy.OAuthSession, error) { var sessions []oatproxy.OAuthSession - if err := db.DB.Find(&sessions).Error; err != nil { + if err := state.DB.Find(&sessions).Error; err != nil { return nil, err } return sessions, nil } -func (db *StatefulDB) GetSessionByDID(did string) (*oatproxy.OAuthSession, error) { +func (state *StatefulDB) GetSessionByDID(did string) (*oatproxy.OAuthSession, error) { var session oatproxy.OAuthSession - if err := db.DB.Where("repo_did = ? AND revoked_at IS NULL", did).Order("updated_at DESC").First(&session).Error; err != nil { + if err := state.DB.Where("repo_did = ? AND revoked_at IS NULL", did).Order("updated_at DESC").First(&session).Error; err != nil { return nil, err } return &session, nil diff --git a/pkg/statedb/statedb.go b/pkg/statedb/statedb.go index a01e2d78b..8c78469c5 100644 --- a/pkg/statedb/statedb.go +++ b/pkg/statedb/statedb.go @@ -23,6 +23,13 @@ type StatefulDB struct { CLI *config.CLI } +// list tables here so we can migrate them +var StatefulDBModels = []any{ + oatproxy.OAuthSession{}, + Notification{}, + Config{}, +} + var NoPostgresDatabaseCode = "3D000" // Stateful database for storing private streamplace state @@ -68,16 +75,13 @@ func MakeDB(cli *config.CLI) (*StatefulDB, error) { } sqlDB.SetMaxOpenConns(1) } - for _, model := range []any{ - oatproxy.OAuthSession{}, - Notification{}, - } { + for _, model := range StatefulDBModels { err = db.AutoMigrate(model) if err != nil { return nil, err } } - return &StatefulDB{DB: db}, nil + return &StatefulDB{DB: db, CLI: cli}, nil } func openDB(dial gorm.Dialector) (*gorm.DB, error) {