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 }