Commit 538c3e9

Eric Bower  ·  2026-08-13 10:05:13 -0400 EDT
parent e5d5430
chore(pgs): uploader tests confirming pico+ only
1 files changed,  +178, -0
+178, -0
  1@@ -379,3 +379,181 @@ func TestRsyncDeleteNestedEmptyDirectories(t *testing.T) {
  2 	close(done)
  3 	time.Sleep(100 * time.Millisecond)
  4 }
  5+
  6+// setupPgsTestServer spins up an in-memory PGS SSH server for testing and
  7+// returns the server address and a teardown func.
  8+func setupPgsTestServer(t *testing.T, dbpool *pgsdb.MemoryDB) (string, func()) {
  9+	t.Helper()
 10+
 11+	opts := &slog.HandlerOptions{Level: slog.LevelDebug}
 12+	logger := slog.New(slog.NewTextHandler(io.Discard, opts))
 13+	slog.SetDefault(logger)
 14+
 15+	st, err := storage.NewStorageMemory(map[string]map[string]string{})
 16+	if err != nil {
 17+		t.Fatalf("failed to create storage: %v", err)
 18+	}
 19+
 20+	pubsub := NewPubsubChan()
 21+
 22+	_ = os.Setenv("PGS_SSH_PORT", "0")
 23+	_ = os.Setenv("PGS_PROM_PORT", "0")
 24+
 25+	cfg := NewPgsConfig(logger, dbpool, st, pubsub)
 26+	done := make(chan error)
 27+	readyCh := make(chan *pssh.SSHServer)
 28+	prometheus.DefaultRegisterer = prometheus.NewRegistry()
 29+
 30+	go StartSshServerForTesting(cfg, done, readyCh)
 31+
 32+	server := <-readyCh
 33+	if server == nil {
 34+		t.Fatal("failed to create ssh server")
 35+	}
 36+
 37+	var actualAddr string
 38+	for i := 0; i < 100; i++ {
 39+		server.Mu.Lock()
 40+		listener := server.Listener
 41+		server.Mu.Unlock()
 42+		if listener != nil {
 43+			actualAddr = listener.Addr().String()
 44+			break
 45+		}
 46+		time.Sleep(10 * time.Millisecond)
 47+	}
 48+	if actualAddr == "" {
 49+		t.Fatal("server listener not ready")
 50+	}
 51+
 52+	teardown := func() {
 53+		_ = pubsub.Close()
 54+		close(done)
 55+		time.Sleep(100 * time.Millisecond)
 56+	}
 57+
 58+	return actualAddr, teardown
 59+}
 60+
 61+// TestPlusFlagAllowsUpload verifies that a user with a valid pico+ feature flag
 62+// can successfully upload files via SFTP.
 63+func TestPlusFlagAllowsUpload(t *testing.T) {
 64+	dbpool := pgsdb.NewDBMemory(slog.Default())
 65+	// SetupTestData creates a user with a valid (non-expired) plus feature flag.
 66+	dbpool.SetupTestData()
 67+
 68+	addr, teardown := setupPgsTestServer(t, dbpool)
 69+	defer teardown()
 70+
 71+	user := GenerateUser()
 72+	dbpool.Pubkeys = append(dbpool.Pubkeys, &db.PublicKey{
 73+		ID:     "plus-user-pubkey",
 74+		UserID: dbpool.Users[0].ID,
 75+		Key:    shared.KeyForKeyText(user.signer.PublicKey()),
 76+	})
 77+
 78+	conn, err := user.NewClientAddr(addr)
 79+	if err != nil {
 80+		t.Fatalf("ssh dial failed: %v", err)
 81+	}
 82+	defer func() { _ = conn.Close() }()
 83+
 84+	client, err := sftp.NewClient(conn)
 85+	if err != nil {
 86+		t.Fatalf("sftp client failed: %v", err)
 87+	}
 88+	defer func() { _ = client.Close() }()
 89+
 90+	// Upload should succeed for a plus user.
 91+	f, err := client.Create("myproject/index.html")
 92+	if err != nil {
 93+		t.Fatalf("plus user upload was blocked unexpectedly: %v", err)
 94+	}
 95+	if _, err := f.Write([]byte("<html>hello</html>")); err != nil {
 96+		t.Fatalf("write failed: %v", err)
 97+	}
 98+	_ = f.Close()
 99+
100+	// Confirm the file is accessible via SFTP stat.
101+	fi, err := client.Lstat("myproject/index.html")
102+	if err != nil {
103+		t.Fatalf("file not found after upload: %v", err)
104+	}
105+	if fi.Size() == 0 {
106+		t.Errorf("uploaded file has zero size")
107+	}
108+
109+	// Confirm the project was created in the DB (no empty-project side-effect).
110+	projects, err := dbpool.FindProjectsByUser(dbpool.Users[0].ID)
111+	if err != nil {
112+		t.Fatalf("FindProjectsByUser: %v", err)
113+	}
114+	found := false
115+	for _, p := range projects {
116+		if p.Name == "myproject" {
117+			found = true
118+			break
119+		}
120+	}
121+	if !found {
122+		t.Errorf("project 'myproject' was not created in the DB after upload")
123+	}
124+}
125+
126+// TestNoPlusFlagBlocksUpload verifies that a user whose pico+ subscription has
127+// expired cannot upload files. The Validate() step (introduced in the last
128+// commit) must reject the connection before any project is created, so we also
129+// confirm that no empty project leaks into the DB.
130+func TestNoPlusFlagBlocksUpload(t *testing.T) {
131+	dbpool := pgsdb.NewDBMemory(slog.Default())
132+	// Create a user but give them an *expired* plus feature flag.
133+	dbpool.SetupTestData()
134+	expired := time.Now().Add(-time.Hour)
135+	dbpool.Feature.ExpiresAt = &expired // mark as expired
136+
137+	addr, teardown := setupPgsTestServer(t, dbpool)
138+	defer teardown()
139+
140+	user := GenerateUser()
141+	dbpool.Pubkeys = append(dbpool.Pubkeys, &db.PublicKey{
142+		ID:     "no-plus-user-pubkey",
143+		UserID: dbpool.Users[0].ID,
144+		Key:    shared.KeyForKeyText(user.signer.PublicKey()),
145+	})
146+
147+	conn, err := user.NewClientAddr(addr)
148+	if err != nil {
149+		t.Fatalf("ssh dial failed: %v", err)
150+	}
151+	defer func() { _ = conn.Close() }()
152+
153+	// Validate() fires at the SFTP-subsystem session level, so the server may
154+	// close the connection before the SFTP version handshake completes.
155+	// Accept rejection at either sftp.NewClient or client.Create as valid.
156+	client, sftpErr := sftp.NewClient(conn)
157+	if sftpErr != nil {
158+		// Server rejected the session during Validate() — this is the expected
159+		// behavior for a non-plus user. Verify no project was leaked.
160+		t.Logf("upload correctly blocked at SFTP handshake: %v", sftpErr)
161+	} else {
162+		defer func() { _ = client.Close() }()
163+
164+		// Upload must be rejected because the user's pico+ has expired.
165+		_, err = client.Create("someproject/index.html")
166+		if err == nil {
167+			t.Fatal("expected upload to be blocked for non-plus user, but it succeeded")
168+		}
169+		t.Logf("upload correctly blocked at Create: %v", err)
170+	}
171+
172+	// Most importantly: no empty project should have been created in the DB,
173+	// because the refactored Validate() checks the feature flag *before*
174+	// upserting the project bucket/quota.
175+	projects, err := dbpool.FindProjectsByUser(dbpool.Users[0].ID)
176+	if err != nil {
177+		t.Fatalf("FindProjectsByUser: %v", err)
178+	}
179+	if len(projects) != 0 {
180+		t.Errorf("expected 0 projects for non-plus user, got %d (empty project leak)", len(projects))
181+	}
182+}