Files
jiang13-bbs/backend/service/operations_test.go

627 lines
19 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package service
import (
"bufio"
"context"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"encoding/base64"
"encoding/json"
"encoding/pem"
"errors"
"fmt"
"github.com/freefire/jiang13-bbs/config"
"github.com/freefire/jiang13-bbs/model"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/logger"
"io"
"math/big"
"net"
"net/http"
"net/http/httptest"
"net/http/httputil"
"net/url"
"os"
"path/filepath"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
)
func testKey(t *testing.T) {
t.Helper()
t.Setenv("SETTINGS_MASTER_KEY", base64.StdEncoding.EncodeToString(make([]byte, 32)))
}
func testOps(t *testing.T) *Operations {
t.Helper()
dsn := os.Getenv("OPS_TEST_DATABASE_URL")
if dsn == "" {
t.Skip("set OPS_TEST_DATABASE_URL to an isolated PostgreSQL database")
}
u, e := url.Parse(dsn)
if e != nil || !strings.HasPrefix(strings.TrimPrefix(u.Path, "/"), "ops_test") {
t.Fatal("test database name must begin ops_test")
}
db, e := gorm.Open(postgres.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
if e != nil {
t.Fatal(e)
}
schema := "ops_test_" + fmt.Sprint(time.Now().UnixNano())
if e = db.Exec("CREATE SCHEMA " + schema).Error; e != nil {
t.Fatal(e)
}
q := u.Query()
q.Set("search_path", schema)
u.RawQuery = q.Encode()
scoped, e := gorm.Open(postgres.Open(u.String()), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
if e != nil {
t.Fatal(e)
}
t.Cleanup(func() {
sql, _ := scoped.DB()
_ = sql.Close()
db.Exec("DROP SCHEMA " + schema + " CASCADE")
sql, _ = db.DB()
_ = sql.Close()
})
if e = scoped.AutoMigrate(&model.TemporaryUpload{}, &model.ModuleConfig{}, &model.SiteSetting{}, &model.SettingsAudit{}, &model.ActionCounter{}, &model.MailTask{}, &model.EmailChallenge{}, &model.StoredObject{}, &model.User{}, &model.PostAttachment{}); e != nil {
t.Fatal(e)
}
dir := t.TempDir()
_ = os.MkdirAll(filepath.Join(dir, "uploads"), 0700)
_ = os.MkdirAll(filepath.Join(dir, "private"), 0700)
return NewOperations(scoped, &config.Config{DataDir: dir, DevMode: true, SiteURL: "https://forum.example.com"})
}
func raw(v any) json.RawMessage { b, _ := json.Marshal(v); return b }
func TestOperationalSecrets(t *testing.T) {
testKey(t)
o := &Operations{}
enc, e := o.seal("do-not-leak", "mail:password")
if e != nil {
t.Fatal(e)
}
enc2, _ := o.seal("do-not-leak", "mail:password")
if enc == enc2 || strings.Contains(enc, "do-not-leak") {
t.Fatal("nonce or encryption failure")
}
p, e := o.open(enc, "mail:password")
if e != nil || p != "do-not-leak" {
t.Fatal("roundtrip")
}
if _, e = o.open(enc, "storage:secret_key"); e == nil {
t.Fatal("domain substitution accepted")
}
t.Setenv("SETTINGS_MASTER_KEY", "")
if _, e = o.open(enc, "mail:password"); e == nil {
t.Fatal("missing key accepted")
}
}
func TestFilterNormalizationExceptionsAndMarkdown(t *testing.T) {
c := FilterConfig{Enabled: true, Rules: []FilterRule{{ID: "a", Word: "bad", Scopes: []string{"body"}, Action: "block", Enabled: true, Exceptions: []string{"badminton"}}, {ID: "b", Word: "hello", Scopes: []string{"body"}, Action: "log", Enabled: true}}}
r := MatchFilter(c, "body", "HELLO badminton and BAD **bad** [label](https://bad.example)\n\n\u0060bad\u0060\n\n\u0060\u0060\u0060go\nbad\n\u0060\u0060\u0060")
if r.Result != "block" {
t.Fatal(r)
}
blocks, excepted := 0, 0
for _, h := range r.Hits {
if h.Action == "block" {
if h.Excepted {
excepted++
} else {
blocks++
}
}
}
if blocks != 2 || excepted != 1 {
t.Fatalf("blocks=%d excepted=%d text=%q", blocks, excepted, r.Text)
}
if MatchFilter(c, "title", "bad").Result != "pass" {
t.Fatal("scope leak")
}
if MatchFilter(c, "body", "badminton").Result != "pass" {
t.Fatal("exception not applied")
}
}
func TestRateLimiterUsesRequestCategory(t *testing.T) {
r := NewRateLimiter()
r.SetLimit(RateLogin, 1)
if !r.Allow("login:127.0.0.1") || r.Allow("login:127.0.0.1") {
t.Fatal("per-IP key bypassed limit")
}
if !r.Allow("login:127.0.0.2") {
t.Fatal("different IP denied")
}
}
func TestPrivateServiceTargetDenied(t *testing.T) {
t.Setenv("SERVICE_PRIVATE_HOSTS", "")
c, e := safeDial(context.Background(), "tcp", "127.0.0.1:25")
if c != nil {
c.Close()
}
if e == nil || !strings.Contains(e.Error(), "受限") {
t.Fatal(e)
}
}
func TestModuleConflictAndAtomicDependency(t *testing.T) {
testKey(t)
o := testOps(t)
cfg := defaultModule("security").(*SecurityConfig)
cfg.CommentInterval = 15
var wins atomic.Int32
var wg sync.WaitGroup
for range 8 {
wg.Add(1)
go func() {
defer wg.Done()
e := o.Save("security", 0, raw(cfg), nil, 1)
if e == nil {
wins.Add(1)
} else if !errors.Is(e, ErrConfigConflict) {
t.Errorf("unexpected: %v", e)
}
}()
}
wg.Wait()
if wins.Load() != 1 {
t.Fatalf("concurrent winners: %d", wins.Load())
}
reopened := NewOperations(o.db, o.cfg)
s, e := reopened.Security()
if e != nil || s.CommentInterval != 15 {
t.Fatal("restart persistence", e)
}
cfg.VerifyEmail = true
cfg.AllowRegister = false
if e = o.Save("security", 1, raw(cfg), nil, 1); e == nil {
t.Fatal("dependency accepted")
}
s, _ = o.Security()
if !s.AllowRegister || s.VerifyEmail {
t.Fatal("failed save changed active config")
}
m := defaultModule("mail").(*MailConfig)
m.Enabled = true
m.Host = "smtp.example.com"
m.Username = "user"
m.Password = "secret-value"
m.From = "sender@example.com"
if e = o.Save("mail", 0, raw(m), nil, 1); e != nil {
t.Fatal(e)
}
got, e := o.Read("mail")
if e != nil {
t.Fatal(e)
}
b, _ := json.Marshal(got)
if strings.Contains(string(b), "secret-value") || strings.Contains(string(b), "mail:password") {
t.Fatal("secret exposed")
}
cfg.AllowRegister = true
if e = o.Save("security", 1, raw(cfg), nil, 1); e != nil {
t.Fatal(e)
}
if e = o.Save("mail", 1, raw(map[string]any{"enabled": false}), nil, 1); e == nil {
t.Fatal("mail dependency bypass")
}
if e = o.Save("mail", 1, raw(map[string]any{}), []string{"password"}, 1); e == nil {
t.Fatal("dependent credential cleared")
}
if e = o.Save("mail", 1, raw(map[string]any{"password": ""}), nil, 1); e != nil {
t.Fatal("blank failed to retain secret", e)
}
}
func TestSharedQuotaAndFiniteLoginRecovery(t *testing.T) {
o := testOps(t)
var wins atomic.Int32
var wg sync.WaitGroup
for range 20 {
wg.Add(1)
go func() {
defer wg.Done()
wait, e := o.Quota("same", 5, 60)
if e != nil {
t.Error(e)
} else if wait == 0 {
wins.Add(1)
}
}()
}
wg.Wait()
if wins.Load() != 5 {
t.Fatal(wins.Load())
}
cfg := *defaultModule("security").(*SecurityConfig)
for range cfg.LoginFailures {
o.RecordFailure("Victim", "ip", cfg)
}
wait, e := o.FailureWait("victim", "other", cfg)
if e != nil || wait <= 0 || wait > cfg.LoginWindow+1 {
t.Fatal(wait, e)
}
o.db.Model(&model.ActionCounter{}).Where("key = ?", counterKey("failure:account:victim")).Update("expires_at", time.Now().Add(-time.Second))
if wait, e = o.FailureWait("victim", "other", cfg); e != nil || wait != 0 {
t.Fatal("account did not recover")
}
}
func smtpFixture(t *testing.T, authFail, stall bool) (string, int, *x509.CertPool, *atomic.Int32) {
t.Helper()
key, _ := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
tpl := &x509.Certificate{SerialNumber: big.NewInt(1), Subject: pkix.Name{CommonName: "localhost"}, NotBefore: time.Now().Add(-time.Hour), NotAfter: time.Now().Add(time.Hour), IPAddresses: []net.IP{net.ParseIP("127.0.0.1")}, KeyUsage: x509.KeyUsageDigitalSignature, ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}}
der, e := x509.CreateCertificate(rand.Reader, tpl, tpl, &key.PublicKey, key)
if e != nil {
t.Fatal(e)
}
keyDER, _ := x509.MarshalECPrivateKey(key)
cert, _ := tls.X509KeyPair(pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}), pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER}))
pool := x509.NewCertPool()
pool.AppendCertsFromPEM(pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}))
ln, e := tls.Listen("tcp", "127.0.0.1:0", &tls.Config{Certificates: []tls.Certificate{cert}})
if e != nil {
t.Fatal(e)
}
t.Cleanup(func() { ln.Close() })
var accepted atomic.Int32
go func() {
for {
conn, e := ln.Accept()
if e != nil {
return
}
go func() {
defer conn.Close()
_ = conn.SetDeadline(time.Now().Add(4 * time.Second))
if stall {
time.Sleep(3 * time.Second)
return
}
io.WriteString(conn, "220 test SMTP\r\n")
r := bufio.NewReader(conn)
for {
line, e := r.ReadString('\n')
if e != nil {
return
}
switch {
case strings.HasPrefix(line, "EHLO"):
io.WriteString(conn, "250-test\r\n250 AUTH PLAIN\r\n")
case strings.HasPrefix(line, "AUTH"):
if authFail {
io.WriteString(conn, "535 rejected\r\n")
} else {
io.WriteString(conn, "235 accepted\r\n")
}
case strings.HasPrefix(line, "DATA"):
io.WriteString(conn, "354 go\r\n")
for {
l, e := r.ReadString('\n')
if e != nil {
return
}
if l == ".\r\n" {
break
}
}
accepted.Add(1)
io.WriteString(conn, "250 queued\r\n")
case strings.HasPrefix(line, "QUIT"):
io.WriteString(conn, "221 bye\r\n")
return
default:
io.WriteString(conn, "250 ok\r\n")
}
}
}()
}
}()
host, port, _ := net.SplitHostPort(ln.Addr().String())
n := 0
fmt.Sscan(port, &n)
return host, n, pool, &accepted
}
func TestSMTPConnectionAcceptAuthAndTimeout(t *testing.T) {
testKey(t)
t.Setenv("SERVICE_PRIVATE_HOSTS", "127.0.0.1")
for _, mode := range []string{"success", "auth", "timeout", "certificate"} {
t.Run(mode, func(t *testing.T) {
host, port, roots, accepted := smtpFixture(t, mode == "auth", mode == "timeout")
o := &Operations{smtpRoots: roots}
if mode == "certificate" {
o.smtpRoots = nil
}
secret, _ := o.seal("smtp-test-secret", "mail:password")
c := MailConfig{Host: host, Port: port, TLS: "tls", Username: "test", Password: secret, From: "test@example.com", FromName: "Test", Timeout: 2}
e := o.smtp(context.Background(), c, &mailPayload{To: "admin@example.com", Subject: "Test", Body: "<p>test</p>"}, "test")
if mode == "success" {
if e != nil || accepted.Load() != 1 {
t.Fatal(e, accepted.Load())
}
} else if e == nil {
t.Fatal("failure expected")
}
})
}
}
func TestStorageProbeAndHistoricPrivateObject(t *testing.T) {
testKey(t)
t.Setenv("SERVICE_PRIVATE_HOSTS", "127.0.0.1")
o := testOps(t)
var mu sync.Mutex
objects := map[string][]byte{}
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
mu.Lock()
defer mu.Unlock()
if r.Header.Get("Authorization") == "" {
w.WriteHeader(403)
return
}
switch r.Method {
case "PUT":
var body io.Reader = r.Body
if strings.Contains(r.Header.Get("Content-Encoding"), "aws-chunked") || strings.HasPrefix(r.Header.Get("X-Amz-Content-Sha256"), "STREAMING-") {
body = httputil.NewChunkedReader(r.Body)
}
b, _ := io.ReadAll(body)
objects[r.URL.Path] = b
w.Header().Set("ETag", `"test"`)
w.WriteHeader(200)
case "GET", "HEAD":
b, ok := objects[r.URL.Path]
if !ok {
w.WriteHeader(404)
return
}
w.Header().Set("Content-Length", fmt.Sprint(len(b)))
w.Header().Set("Content-Type", "application/octet-stream")
w.Header().Set("ETag", `"test"`)
w.Header().Set("Last-Modified", time.Now().UTC().Format(http.TimeFormat))
if r.Method == "GET" {
w.Write(b)
}
case "DELETE":
delete(objects, r.URL.Path)
w.WriteHeader(204)
default:
w.WriteHeader(400)
}
}))
defer srv.Close()
c := defaultModule("storage").(*StorageConfig)
c.Backend = "s3"
c.Endpoint = srv.URL
c.Bucket = "test-bucket"
c.AccessKey = "test-access"
c.SecretKey = "test-secret"
if e := o.TestStorage(context.Background(), raw(c), nil); e != nil {
t.Fatal(e)
}
if len(objects) != 0 {
t.Fatal("probe leaked object")
}
if e := o.Save("storage", 0, raw(c), nil, 1); e != nil {
t.Fatal(e)
}
path := filepath.Join(t.TempDir(), "file.txt")
os.WriteFile(path, []byte("historic"), 0600)
id, e := o.StoreFile(path, "text/plain", false)
if e != nil {
t.Fatal(e)
}
c.Backend = "local"
c.AccessKey = ""
c.SecretKey = ""
if e = o.Save("storage", 1, raw(c), nil, 1); e != nil {
t.Fatal(e)
}
if _, _, e = o.OpenObject(context.Background(), id, true); e == nil {
t.Fatal("private object publicly accessible")
}
r, _, e := o.OpenObject(context.Background(), id, false)
if e != nil {
t.Fatal(e)
}
defer r.Close()
b, _ := io.ReadAll(r)
if string(b) != "historic" {
t.Fatalf("%q", b)
}
_ = r.Close()
if e = o.RemoveObject(id); e != nil {
t.Fatal(e)
}
if _, _, e = o.OpenObject(context.Background(), id, false); e == nil {
t.Fatal("deleted object accessible")
}
}
func TestTemporaryCleanupProtectsActiveReferencedAndDraftFiles(t *testing.T) {
o := testOps(t)
if e := o.db.Exec("CREATE TABLE posts (content text); CREATE TABLE comments (content text)").Error; e != nil {
t.Fatal(e)
}
dir := filepath.Join(o.cfg.DataDir, "uploads", "images")
_ = os.MkdirAll(dir, 0700)
path := filepath.Join(dir, "managed.png.partial")
release, e := o.BeginTemporary(path)
if e != nil {
t.Fatal(e)
}
_ = os.WriteFile(path, []byte("temporary"), 0600)
old := time.Now().Add(-8 * 24 * time.Hour)
_ = os.Chtimes(path, old, old)
o.db.Model(&model.TemporaryUpload{}).Where("id = ?", counterKey("uploads/images/managed.png.partial")).Update("created_at", old)
candidates, e := o.temporaryCandidates()
if e != nil || len(candidates) != 0 {
t.Fatal("active upload selected", e)
}
release()
candidates, e = o.temporaryCandidates()
if e != nil || len(candidates) != 1 {
t.Fatal("expired inactive upload not selected", e)
}
_ = o.db.Exec("INSERT INTO posts(content) VALUES (?)", "referenced "+filepath.Base(path)).Error
r, e := o.CleanTemporary([]string{candidates[0].ID})
if e != nil || r["removed"] != 0 {
t.Fatal("reference deleted", r, e)
}
o.db.Exec("DELETE FROM posts")
draft := filepath.Join(dir, "draft.png")
_ = os.WriteFile(draft, []byte("draft"), 0600)
_ = os.Chtimes(draft, old, old)
unknown := filepath.Join(dir, "legacy.png.partial")
_ = os.WriteFile(unknown, []byte("unknown"), 0600)
_ = os.Chtimes(unknown, old, old)
r, e = o.CleanTemporary([]string{candidates[0].ID})
if e != nil || r["removed"] != 1 {
t.Fatal(r, e)
}
for _, p := range []string{draft, unknown} {
if _, e = os.Stat(p); e != nil {
t.Fatal("protected file removed")
}
}
r, e = o.CleanTemporary([]string{candidates[0].ID})
if e != nil || r["removed"] != 0 {
t.Fatal("duplicate cleanup", r, e)
}
}
func TestDurableMailRetriesAndCodeSingleUse(t *testing.T) {
testKey(t)
t.Setenv("SERVICE_PRIVATE_HOSTS", "127.0.0.1")
o := testOps(t)
host, port, roots, _ := smtpFixture(t, true, false)
o.smtpRoots = roots
c := defaultModule("mail").(*MailConfig)
c.Enabled = true
c.Host = host
c.Port = port
c.From = "test@example.com"
c.FromName = "Test"
c.Username = "test"
c.Password = "test"
if e := o.Save("mail", 0, raw(c), nil, 1); e != nil {
t.Fatal(e)
}
if e := o.EnqueueMail(o.db, "user@example.com", "register", "single-code", "same-task"); e != nil {
t.Fatal(e)
}
if e := o.EnqueueMail(o.db, "user@example.com", "register", "single-code", "same-task"); e != nil {
t.Fatal(e)
}
for range 3 {
if e := o.ProcessMail(context.Background()); e != nil {
t.Fatal(e)
}
o.db.Model(&model.MailTask{}).Where("dedupe = ?", "same-task").Update("next_at", time.Now().Add(-time.Second))
}
rows, e := o.MailRows()
if e != nil || len(rows) != 1 || rows[0].Status != "failed" || rows[0].Attempts != 3 || rows[0].Payload != "" {
t.Fatal(rows, e)
}
challenge := model.EmailChallenge{Hash: counterKey("user@example.com:register:single-code"), Email: "user@example.com", Purpose: "register", ExpiresAt: time.Now().Add(time.Minute)}
if e = o.db.Create(&challenge).Error; e != nil {
t.Fatal(e)
}
if e = o.ConsumeCode("user@example.com", "register", "single-code"); e != nil {
t.Fatal(e)
}
if e = o.ConsumeCode("user@example.com", "register", "single-code"); e == nil {
t.Fatal("code replay")
}
}
func TestMailQueueRecoveryAndDestinationQuota(t *testing.T) {
testKey(t)
t.Setenv("SERVICE_PRIVATE_HOSTS", "127.0.0.1")
o := testOps(t)
host, port, roots, accepted := smtpFixture(t, false, false)
c := defaultModule("mail").(*MailConfig)
c.Enabled = true
c.Host = host
c.Port = port
c.Username = "test"
c.Password = "secret"
c.From = "test@example.com"
c.FromName = "Test"
if e := o.Save("mail", 0, raw(c), nil, 1); e != nil {
t.Fatal(e)
}
security := defaultModule("security").(*SecurityConfig)
security.VerifyEmail = true
if e := o.Save("security", 0, raw(security), nil, 1); e != nil {
t.Fatal(e)
}
if wait, e := o.SendCode("USER@example.com", "register", "192.0.2.1"); e != nil || wait != 0 {
t.Fatal(wait, e)
}
if wait, e := o.SendCode("user@example.com", "register", "192.0.2.2"); e != nil || wait < 1 {
t.Fatal("destination quota bypassed", wait, e)
}
var task model.MailTask
if e := o.db.First(&task).Error; e != nil {
t.Fatal(e)
}
if strings.Contains(task.Payload, "user@example.com") {
t.Fatal("unencrypted queue")
}
o.db.Model(&task).Updates(map[string]any{"status": "sending", "attempts": 1, "next_at": time.Now().Add(-time.Minute)})
restarted := NewOperations(o.db, o.cfg)
restarted.smtpRoots = roots
if e := restarted.ProcessMail(context.Background()); e != nil {
t.Fatal(e)
}
var finished model.MailTask
o.db.First(&finished, task.ID)
if finished.Status != "accepted" || finished.Attempts != 2 || finished.Payload != "" || accepted.Load() != 1 {
t.Fatal("queue did not recover", finished.Status, finished.Attempts)
}
}
func TestMailTemplateGuards(t *testing.T) {
var cfg MailConfig
cfg.applyTemplateDefaults()
if cfg.BodyTemplate != defaultMailBody || cfg.SubjectTemplate != defaultMailSubject {
t.Fatal("defaults not applied")
}
customBody := "<p>{{code}}</p>"
cfg.BodyTemplate = customBody
cfg.applyTemplateDefaults()
if cfg.BodyTemplate != customBody {
t.Fatal("自定义正文不应被默认模板覆盖")
}
o := &Operations{cfg: &config.Config{SiteURL: "https://forum.example.com", DevMode: true}}
if e := o.validate(nil, "mail", &MailConfig{Port: 465, TLS: "tls", Timeout: 10, Retention: 30, BodyTemplate: "<script></script><p>no code</p>"}); e == nil || !strings.Contains(e.Error(), "{{code}}") {
t.Fatal(e)
}
custom := `<!DOCTYPE html><html><body><style>.x{color:red}</style><script>void 0</script><p class="x">{{code}}</p><a href="{{link}}">go</a></body></html>`
p := o.renderMail(MailConfig{BodyTemplate: custom}, "<b>x</b>", "reset")
if strings.Contains(p.Body, "<b>x</b>") || !strings.Contains(p.Body, "&lt;b&gt;x&lt;/b&gt;") || !strings.Contains(p.Body, "<script>") {
t.Fatal(p.Body)
}
if strings.Contains(strings.ToLower(p.Body), "<style") {
t.Fatal("style blocks should be inlined for mail clients", p.Body)
}
if !strings.Contains(p.Body, "color:red") && !strings.Contains(p.Body, "color: red") {
t.Fatal("class styles should be inlined", p.Body)
}
if !strings.Contains(p.Subject, "密码找回") || strings.Contains(p.Subject, "\n") {
t.Fatal(p.Subject)
}
if !strings.Contains(p.Body, "https://forum.example.com/reset-password") {
t.Fatal(p.Body)
}
def := o.renderMail(MailConfig{}, "123456", "")
if strings.Contains(strings.ToLower(def.Body), "<style") {
t.Fatal("default template style should be inlined")
}
if !strings.Contains(def.Body, "background:#f5f7fb") && !strings.Contains(def.Body, "background: #f5f7fb") {
t.Fatal("default wrap style missing after prepare", def.Body)
}
}