Commit 7b5a5aa
Eric Bower
·
2026-08-17 09:35:36 -0400 EDT
parent c58eb5b
chore(pgs): check for pgs feature flag
11 files changed,
+290,
-23
+37,
-0
1@@ -0,0 +1,37 @@
2+package main
3+
4+import (
5+ "log/slog"
6+ "os"
7+
8+ "github.com/picosh/pico/pkg/db/postgres"
9+)
10+
11+func main() {
12+ logger := slog.Default()
13+ DbURL := os.Getenv("DATABASE_URL")
14+ dbpool := postgres.NewDB(DbURL, logger)
15+
16+ args := os.Args
17+ if len(args) < 3 {
18+ logger.Error("usage: go run ./cmd/scripts/feature-flag <username> <feature>")
19+ os.Exit(1)
20+ }
21+
22+ username := args[1]
23+ feature := args[2]
24+
25+ logger.Info(
26+ "Adding feature flag to user",
27+ "username", username,
28+ "feature", feature,
29+ )
30+
31+ err := dbpool.AddFeatureUser(username, feature)
32+ if err != nil {
33+ logger.Error("Failed to add feature flag to user", "err", err)
34+ os.Exit(1)
35+ } else {
36+ logger.Info("Successfully added feature flag to user")
37+ }
38+}
+3,
-5
1@@ -230,13 +230,11 @@ This means only you can access the site through a web tunnel or by downloading t
2 }
3
4 func (c *Cmd) stats(cfgMaxSize uint64) error {
5- ff, err := c.Dbpool.FindFeature(c.User.ID, "plus")
6+ ff, err := findFeatureFlag(c.Dbpool, c.Cfg, c.User.ID)
7 if err != nil {
8- ff = db.NewFeatureFlag(c.User.ID, "plus", cfgMaxSize, 0, 0)
9+ ff = db.NewFeatureFlag(c.User.ID, "pgs", cfgMaxSize, 0, 0)
10 }
11- // this is jank
12- ff.Data.StorageMax = ff.FindStorageMax(cfgMaxSize)
13- storageMax := ff.Data.StorageMax
14+ storageMax := ff.FindStorageMax(cfgMaxSize)
15
16 bucketName := shared.GetAssetBucketName(c.User.ID)
17 bucket, err := c.Store.GetBucket(bucketName)
+1,
-1
1@@ -62,7 +62,7 @@ func (c *PgsConfig) StaticPath(fname string) string {
2 return filepath.Join("pkg", "apps", "pgs", fname)
3 }
4
5-var maxSize = uint64(25 * shared.MB)
6+var maxSize = uint64(50 * shared.MB)
7 var maxAssetSize = int64(10 * shared.MB)
8
9 // Needs to be small for caching files like _headers and _redirects.
+10,
-1
1@@ -16,6 +16,7 @@ type MemoryDB struct {
2 Projects []*db.Project
3 Pubkeys []*db.PublicKey
4 Feature *db.FeatureFlag
5+ Features []*db.FeatureFlag
6 FormEntries []*db.FormEntry
7 }
8
9@@ -82,7 +83,15 @@ func (me *MemoryDB) FindUserByName(name string) (*db.User, error) {
10 }
11
12 func (me *MemoryDB) FindFeature(userID, name string) (*db.FeatureFlag, error) {
13- return me.Feature, nil
14+ for _, ff := range me.Features {
15+ if ff.UserID == userID && ff.Name == name {
16+ return ff, nil
17+ }
18+ }
19+ if me.Feature != nil && (me.Feature.Name == name || me.Feature.Name == "") {
20+ return me.Feature, nil
21+ }
22+ return nil, fmt.Errorf("feature flag not found")
23 }
24
25 func (me *MemoryDB) Close() error {
+43,
-0
1@@ -0,0 +1,43 @@
2+package pgs
3+
4+import (
5+ "fmt"
6+ "strings"
7+
8+ pgsdb "github.com/picosh/pico/pkg/apps/pgs/db"
9+ "github.com/picosh/pico/pkg/db"
10+)
11+
12+func setFeatureLimits(ff *db.FeatureFlag, cfg *PgsConfig) {
13+ ff.Data.StorageMax = ff.FindStorageMax(cfg.MaxSize)
14+ ff.Data.FileMax = ff.FindFileMax(cfg.MaxAssetSize)
15+ ff.Data.SpecialFileMax = ff.FindSpecialFileMax(cfg.MaxSpecialFileSize)
16+}
17+
18+func findFeatureFlag(dbpool pgsdb.PgsDB, cfg *PgsConfig, userID string) (*db.FeatureFlag, error) {
19+ ff, err := dbpool.FindFeature(userID, "plus")
20+ if err == nil {
21+ if ff.IsValid() {
22+ setFeatureLimits(ff, cfg)
23+ return ff, nil
24+ }
25+ err = fmt.Errorf("your pico+ has expired")
26+ }
27+
28+ ffPgs, pgsErr := dbpool.FindFeature(userID, "pgs")
29+ if pgsErr == nil {
30+ if ffPgs.IsValid() {
31+ setFeatureLimits(ffPgs, cfg)
32+ return ffPgs, nil
33+ }
34+ pgsErr = fmt.Errorf("your pgs access has expired")
35+ }
36+
37+ if err != nil && strings.Contains(err.Error(), "expired") {
38+ return nil, err
39+ }
40+ if pgsErr != nil {
41+ return nil, pgsErr
42+ }
43+ return nil, err
44+}
+1,
-16
1@@ -15,7 +15,6 @@ import (
2 "sync"
3 "time"
4
5- pgsdb "github.com/picosh/pico/pkg/apps/pgs/db"
6 "github.com/picosh/pico/pkg/db"
7 "github.com/picosh/pico/pkg/pssh"
8 sendutils "github.com/picosh/pico/pkg/send/utils"
9@@ -219,7 +218,7 @@ func (h *UploadAssetHandler) Validate(s *pssh.SSHServerConnSession) error {
10 return err
11 }
12
13- ff, err := findPlusFF(h.Cfg.DB, h.Cfg, user.ID)
14+ ff, err := findFeatureFlag(h.Cfg.DB, h.Cfg, user.ID)
15 if err != nil {
16 return err
17 }
18@@ -273,20 +272,6 @@ func (h *UploadAssetHandler) findDenylist(bucket storage.Bucket, project *db.Pro
19 return str, nil
20 }
21
22-func findPlusFF(dbpool pgsdb.PgsDB, cfg *PgsConfig, userID string) (*db.FeatureFlag, error) {
23- ff, err := dbpool.FindFeature(userID, "plus")
24- if err != nil {
25- return nil, err
26- }
27- if !ff.IsValid() {
28- return nil, fmt.Errorf("your pico+ has expired")
29- }
30- ff.Data.StorageMax = ff.FindStorageMax(cfg.MaxSize)
31- ff.Data.FileMax = ff.FindFileMax(cfg.MaxAssetSize)
32- ff.Data.SpecialFileMax = ff.FindSpecialFileMax(cfg.MaxSpecialFileSize)
33- return ff, nil
34-}
35-
36 func mtimeToTime(entry *sendutils.FileEntry) time.Time {
37 var mtime time.Time
38 if entry.Mtime > 0 {
+157,
-0
1@@ -557,3 +557,160 @@ func TestNoPlusFlagBlocksUpload(t *testing.T) {
2 t.Errorf("expected 0 projects for non-plus user, got %d (empty project leak)", len(projects))
3 }
4 }
5+
6+func TestFindFeatureFlag(t *testing.T) {
7+ logger := slog.Default()
8+ cfg := &PgsConfig{
9+ MaxSize: uint64(100 * shared.MB),
10+ MaxAssetSize: int64(10 * shared.MB),
11+ MaxSpecialFileSize: int64(5 * shared.KB),
12+ }
13+
14+ validExpires := time.Now().Add(24 * time.Hour)
15+ expiredExpires := time.Now().Add(-24 * time.Hour)
16+ userID := "user-1"
17+
18+ t.Run("plus valid returns plus", func(t *testing.T) {
19+ dbpool := pgsdb.NewDBMemory(logger)
20+ plusFF := db.NewFeatureFlag(userID, "plus", uint64(50*shared.MB), int64(5*shared.MB), int64(2*shared.KB))
21+ plusFF.ExpiresAt = &validExpires
22+ dbpool.Features = []*db.FeatureFlag{plusFF}
23+
24+ ff, err := findFeatureFlag(dbpool, cfg, userID)
25+ if err != nil {
26+ t.Fatalf("unexpected error: %v", err)
27+ }
28+ if ff.Name != "plus" {
29+ t.Errorf("expected plus, got %s", ff.Name)
30+ }
31+ })
32+
33+ t.Run("pgs valid returns pgs when no plus", func(t *testing.T) {
34+ dbpool := pgsdb.NewDBMemory(logger)
35+ pgsFF := db.NewFeatureFlag(userID, "pgs", uint64(50*shared.MB), int64(5*shared.MB), int64(2*shared.KB))
36+ pgsFF.ExpiresAt = &validExpires
37+ dbpool.Features = []*db.FeatureFlag{pgsFF}
38+
39+ ff, err := findFeatureFlag(dbpool, cfg, userID)
40+ if err != nil {
41+ t.Fatalf("unexpected error: %v", err)
42+ }
43+ if ff.Name != "pgs" {
44+ t.Errorf("expected pgs, got %s", ff.Name)
45+ }
46+ })
47+
48+ t.Run("plus preferred over pgs when both valid", func(t *testing.T) {
49+ dbpool := pgsdb.NewDBMemory(logger)
50+ plusFF := db.NewFeatureFlag(userID, "plus", uint64(50*shared.MB), int64(5*shared.MB), int64(2*shared.KB))
51+ plusFF.ExpiresAt = &validExpires
52+ pgsFF := db.NewFeatureFlag(userID, "pgs", uint64(30*shared.MB), int64(3*shared.MB), int64(1*shared.KB))
53+ pgsFF.ExpiresAt = &validExpires
54+ dbpool.Features = []*db.FeatureFlag{plusFF, pgsFF}
55+
56+ ff, err := findFeatureFlag(dbpool, cfg, userID)
57+ if err != nil {
58+ t.Fatalf("unexpected error: %v", err)
59+ }
60+ if ff.Name != "plus" {
61+ t.Errorf("expected plus to be picked first, got %s", ff.Name)
62+ }
63+ })
64+
65+ t.Run("pgs picked when plus is expired", func(t *testing.T) {
66+ dbpool := pgsdb.NewDBMemory(logger)
67+ plusFF := db.NewFeatureFlag(userID, "plus", uint64(50*shared.MB), int64(5*shared.MB), int64(2*shared.KB))
68+ plusFF.ExpiresAt = &expiredExpires
69+ pgsFF := db.NewFeatureFlag(userID, "pgs", uint64(30*shared.MB), int64(3*shared.MB), int64(1*shared.KB))
70+ pgsFF.ExpiresAt = &validExpires
71+ dbpool.Features = []*db.FeatureFlag{plusFF, pgsFF}
72+
73+ ff, err := findFeatureFlag(dbpool, cfg, userID)
74+ if err != nil {
75+ t.Fatalf("unexpected error: %v", err)
76+ }
77+ if ff.Name != "pgs" {
78+ t.Errorf("expected pgs when plus is expired, got %s", ff.Name)
79+ }
80+ })
81+
82+ t.Run("error when both plus and pgs expired", func(t *testing.T) {
83+ dbpool := pgsdb.NewDBMemory(logger)
84+ plusFF := db.NewFeatureFlag(userID, "plus", uint64(50*shared.MB), int64(5*shared.MB), int64(2*shared.KB))
85+ plusFF.ExpiresAt = &expiredExpires
86+ pgsFF := db.NewFeatureFlag(userID, "pgs", uint64(30*shared.MB), int64(3*shared.MB), int64(1*shared.KB))
87+ pgsFF.ExpiresAt = &expiredExpires
88+ dbpool.Features = []*db.FeatureFlag{plusFF, pgsFF}
89+
90+ _, err := findFeatureFlag(dbpool, cfg, userID)
91+ if err == nil {
92+ t.Fatal("expected error, got nil")
93+ }
94+ })
95+
96+ t.Run("error when neither plus nor pgs exists", func(t *testing.T) {
97+ dbpool := pgsdb.NewDBMemory(logger)
98+ _, err := findFeatureFlag(dbpool, cfg, userID)
99+ if err == nil {
100+ t.Fatal("expected error, got nil")
101+ }
102+ })
103+}
104+
105+func TestPgsFlagAllowsUpload(t *testing.T) {
106+ dbpool := pgsdb.NewDBMemory(slog.Default())
107+ dbpool.SetupTestData()
108+ // Replace default plus flag with a valid pgs flag
109+ dbpool.Feature = nil
110+ valid := time.Now().Add(24 * time.Hour)
111+ pgsFF := db.NewFeatureFlag(
112+ dbpool.Users[0].ID,
113+ "pgs",
114+ uint64(25*shared.MB),
115+ int64(10*shared.MB),
116+ int64(5*shared.KB),
117+ )
118+ pgsFF.ExpiresAt = &valid
119+ dbpool.Features = []*db.FeatureFlag{pgsFF}
120+
121+ addr, teardown := setupPgsTestServer(t, dbpool)
122+ defer teardown()
123+
124+ user := GenerateUser()
125+ dbpool.Pubkeys = append(dbpool.Pubkeys, &db.PublicKey{
126+ ID: "pgs-user-pubkey",
127+ UserID: dbpool.Users[0].ID,
128+ Key: shared.KeyForKeyText(user.signer.PublicKey()),
129+ })
130+
131+ conn, err := user.NewClientAddr(addr)
132+ if err != nil {
133+ t.Fatalf("ssh dial failed: %v", err)
134+ }
135+ defer func() { _ = conn.Close() }()
136+
137+ client, err := sftp.NewClient(conn)
138+ if err != nil {
139+ t.Fatalf("sftp client failed: %v", err)
140+ }
141+ defer func() { _ = client.Close() }()
142+
143+ // Upload should succeed for a pgs user.
144+ f, err := client.Create("myproject/index.html")
145+ if err != nil {
146+ t.Fatalf("pgs user upload was blocked unexpectedly: %v", err)
147+ }
148+ if _, err := f.Write([]byte("<html>hello from pgs</html>")); err != nil {
149+ t.Fatalf("write failed: %v", err)
150+ }
151+ _ = f.Close()
152+
153+ // Confirm the file is accessible via SFTP stat.
154+ fi, err := client.Lstat("myproject/index.html")
155+ if err != nil {
156+ t.Fatalf("file not found after upload: %v", err)
157+ }
158+ if fi.Size() == 0 {
159+ t.Errorf("uploaded file has zero size")
160+ }
161+}
+1,
-0
1@@ -609,6 +609,7 @@ type DB interface {
2 VisitUrlNotFound(opts *SummaryOpts) ([]*VisitUrl, error)
3
4 AddPicoPlusUser(username, email, paymentType, txId string) error
5+ AddFeatureUser(username, name string) error
6 FindFeature(userID string, feature string) (*FeatureFlag, error)
7 FindFeaturesByUser(userID string) ([]*FeatureFlag, error)
8 HasFeatureByUser(userID string, feature string) bool
+11,
-0
1@@ -1902,6 +1902,17 @@ func (me *PsqlDB) AddPicoPlusUser(username, email, paymentType, txId string) err
2 return tx.Commit()
3 }
4
5+func (me *PsqlDB) AddFeatureUser(username, name string) error {
6+ user, err := me.FindUserByName(username)
7+ if err != nil {
8+ return err
9+ }
10+
11+ expiresAt := me.createFeatureExpiresAt(user.ID, name)
12+ _, err = me.InsertFeature(user.ID, name, expiresAt)
13+ return err
14+}
15+
16 func (me *PsqlDB) UpsertProject(userID, projectName, projectDir string) (*db.Project, error) {
17 project, err := me.FindProjectByName(userID, projectName)
18 if err == nil {
+22,
-0
1@@ -1081,6 +1081,28 @@ func TestAddPicoPlusUser(t *testing.T) {
2 }
3 }
4
5+func TestAddFeatureUser(t *testing.T) {
6+ cleanupTestData(t)
7+
8+ user, _ := testDB.RegisterUser("featureuserowner", "ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAI featureuserowner", "comment", "")
9+
10+ err := testDB.AddFeatureUser("featureuserowner", "pgs")
11+ if err != nil {
12+ t.Fatalf("AddFeatureUser failed: %v", err)
13+ }
14+
15+ ff, err := testDB.FindFeature(user.ID, "pgs")
16+ if err != nil {
17+ t.Fatalf("FindFeature failed: %v", err)
18+ }
19+ if ff.Name != "pgs" {
20+ t.Errorf("expected feature name pgs, got %s", ff.Name)
21+ }
22+ if !ff.IsValid() {
23+ t.Errorf("expected feature to be valid")
24+ }
25+}
26+
27 // ============ Feed Items Tests ============
28
29 func TestInsertFeedItems(t *testing.T) {
+4,
-0
1@@ -204,6 +204,10 @@ func (me *StubDB) AddPicoPlusUser(username, email, paymentType, txId string) err
2 return errNotImpl
3 }
4
5+func (me *StubDB) AddFeatureUser(username, name string) error {
6+ return errNotImpl
7+}
8+
9 func (me *StubDB) FindUserStats(userID string) (*db.UserStats, error) {
10 return nil, errNotImpl
11 }