Compare commits
9
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0e068c6ade | ||
|
|
308b0b9edd | ||
|
|
8bcb5c63be | ||
|
|
532ee44da3 | ||
|
|
407a023a68 | ||
|
|
bbdd4af66d | ||
|
|
47d627c2e7 | ||
|
|
f2fe4bdbbe | ||
|
|
d777a1e3d3 |
@@ -0,0 +1,170 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/rand"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/config"
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
adminPasswordLen = 20
|
||||||
|
minAdminPasswordLen = 12
|
||||||
|
adminPasswordChars = "ABCDEFGHJKLMNPQRSTUVWXYZabcdefghijkmnopqrstuvwxyz23456789"
|
||||||
|
)
|
||||||
|
|
||||||
|
func cmdAdmin(args []string) error {
|
||||||
|
if len(args) < 1 {
|
||||||
|
return errors.New("usage: nixmsg admin <init|set-password>")
|
||||||
|
}
|
||||||
|
switch args[0] {
|
||||||
|
case "init":
|
||||||
|
return cmdAdminInit(args[1:])
|
||||||
|
case "set-password":
|
||||||
|
return cmdAdminSetPassword(args[1:])
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("unknown admin command: %s", args[0])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdAdminInit(_ []string) error {
|
||||||
|
cfg, err := loadAndValidateConfig()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
ctx := context.Background()
|
||||||
|
db, err := store.Open(cfg.DataDir, cfg.SQLiteSynchronous)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer func() { _ = db.Close() }()
|
||||||
|
|
||||||
|
ok, err := store.HasAdminPassword(ctx, db.Write)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if ok {
|
||||||
|
return errors.New("admin already initialized; use admin set-password")
|
||||||
|
}
|
||||||
|
|
||||||
|
password, err := generateAdminPassword(adminPasswordLen)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
phc, err := auth.HashPassword(password)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := store.SetAdminPasswordHash(ctx, db.Write, phc); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
// 密码只打印到终端一次,不进日志。
|
||||||
|
fmt.Printf("admin password: %s\n", password)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func cmdAdminSetPassword(args []string) error {
|
||||||
|
password, err := parseSetPasswordArgs(args, os.Stdin)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if vErr := validateAdminPassword(password); vErr != nil {
|
||||||
|
return vErr
|
||||||
|
}
|
||||||
|
cfg, err := loadAndValidateConfig()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
ctx := context.Background()
|
||||||
|
db, err := store.Open(cfg.DataDir, cfg.SQLiteSynchronous)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer func() { _ = db.Close() }()
|
||||||
|
|
||||||
|
phc, err := auth.HashPassword(password)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := store.SetAdminPasswordHash(ctx, db.Write, phc); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
fmt.Fprintln(os.Stderr, "admin password updated")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseSetPasswordArgs(args []string, stdin io.Reader) (string, error) {
|
||||||
|
var password string
|
||||||
|
for i := 0; i < len(args); i++ {
|
||||||
|
switch args[i] {
|
||||||
|
case "--password":
|
||||||
|
if i+1 >= len(args) {
|
||||||
|
return "", errors.New("usage: nixmsg admin set-password --password <password>")
|
||||||
|
}
|
||||||
|
password = args[i+1]
|
||||||
|
i++
|
||||||
|
default:
|
||||||
|
if strings.HasPrefix(args[i], "-") {
|
||||||
|
return "", fmt.Errorf("unknown flag: %s", args[i])
|
||||||
|
}
|
||||||
|
if password != "" {
|
||||||
|
return "", errors.New("usage: nixmsg admin set-password --password <password>")
|
||||||
|
}
|
||||||
|
password = args[i]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if password == "" {
|
||||||
|
b, err := io.ReadAll(io.LimitReader(stdin, 4096))
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
password = strings.TrimSpace(string(b))
|
||||||
|
}
|
||||||
|
if password == "" {
|
||||||
|
return "", errors.New("password required (pass --password or stdin)")
|
||||||
|
}
|
||||||
|
return password, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func validateAdminPassword(password string) error {
|
||||||
|
if len(password) < minAdminPasswordLen {
|
||||||
|
return fmt.Errorf("admin password must be at least %d characters", minAdminPasswordLen)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func generateAdminPassword(n int) (string, error) {
|
||||||
|
buf := make([]byte, n)
|
||||||
|
charset := []byte(adminPasswordChars)
|
||||||
|
for i := 0; i < n; {
|
||||||
|
var b [1]byte
|
||||||
|
if _, err := rand.Read(b[:]); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
if int(b[0]) >= 256-(256%len(charset)) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
buf[i] = charset[int(b[0])%len(charset)]
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
return string(buf), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func loadAndValidateConfig() (config.Config, error) {
|
||||||
|
cfgPath := config.PathFromEnv()
|
||||||
|
cfg, err := config.Load(cfgPath)
|
||||||
|
if err != nil {
|
||||||
|
return config.Config{}, err
|
||||||
|
}
|
||||||
|
if err := cfg.Validate(); err != nil {
|
||||||
|
return config.Config{}, err
|
||||||
|
}
|
||||||
|
return cfg, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,63 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||||
|
)
|
||||||
|
|
||||||
|
func cmdBackup(args []string) error {
|
||||||
|
outPath, err := parseBackupArgs(args)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
cfg, err := loadAndValidateConfig()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if dir := filepath.Dir(outPath); dir != "" && dir != "." {
|
||||||
|
if mkErr := os.MkdirAll(dir, 0o755); mkErr != nil {
|
||||||
|
return fmt.Errorf("mkdir backup dir: %w", mkErr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
absOut, err := filepath.Abs(outPath)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
ctx := context.Background()
|
||||||
|
// 对运行中的库:单独打开写连接执行 VACUUM INTO(可与 serve 并存,WAL 下安全)。
|
||||||
|
write, err := store.OpenWriter(cfg.DataDir, cfg.SQLiteSynchronous)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer func() { _ = write.Close() }()
|
||||||
|
if err := store.VacuumInto(ctx, write, filepath.ToSlash(absOut)); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
fmt.Fprintf(os.Stderr, "backup written to %s\n", absOut)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseBackupArgs(args []string) (string, error) {
|
||||||
|
var out string
|
||||||
|
for i := 0; i < len(args); i++ {
|
||||||
|
switch args[i] {
|
||||||
|
case "--out":
|
||||||
|
if i+1 >= len(args) {
|
||||||
|
return "", errors.New("usage: nixmsg backup --out <file.db>")
|
||||||
|
}
|
||||||
|
out = args[i+1]
|
||||||
|
i++
|
||||||
|
default:
|
||||||
|
return "", fmt.Errorf("unknown argument: %s", args[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if out == "" {
|
||||||
|
return "", errors.New("usage: nixmsg backup --out <file.db>")
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,24 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func cmdCheckConfig(args []string) error {
|
||||||
|
path := config.PathFromEnv()
|
||||||
|
if len(args) > 0 {
|
||||||
|
path = args[0]
|
||||||
|
}
|
||||||
|
cfg, err := config.Load(path)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := cfg.Validate(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
fmt.Fprintln(os.Stderr, "config ok")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,153 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestCheckConfigRejectsLargeBody(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
path := filepath.Join(dir, "bad.yaml")
|
||||||
|
body := "listen: \":0\"\ndata_dir: \"" + filepath.ToSlash(dir) + "\"\nlimits:\n max_body_bytes: 262145\n"
|
||||||
|
if err := os.WriteFile(path, []byte(body), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
t.Setenv("NIXMSG_CONFIG", path)
|
||||||
|
err := cmdCheckConfig(nil)
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "max_body_bytes") {
|
||||||
|
t.Fatalf("want max_body_bytes error, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCheckConfigOK(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
path := writeTestConfig(t, dir)
|
||||||
|
t.Setenv("NIXMSG_CONFIG", path)
|
||||||
|
if err := cmdCheckConfig(nil); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAdminInitOnceAndSetPassword(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
path := writeTestConfig(t, dir)
|
||||||
|
t.Setenv("NIXMSG_CONFIG", path)
|
||||||
|
|
||||||
|
var out bytes.Buffer
|
||||||
|
old := os.Stdout
|
||||||
|
r, w, err := os.Pipe()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
os.Stdout = w
|
||||||
|
errInit := cmdAdminInit(nil)
|
||||||
|
_ = w.Close()
|
||||||
|
os.Stdout = old
|
||||||
|
_, _ = out.ReadFrom(r)
|
||||||
|
_ = r.Close()
|
||||||
|
if errInit != nil {
|
||||||
|
t.Fatal(errInit)
|
||||||
|
}
|
||||||
|
text := out.String()
|
||||||
|
if !strings.Contains(text, "admin password:") {
|
||||||
|
t.Fatalf("password not printed: %q", text)
|
||||||
|
}
|
||||||
|
pass := strings.TrimSpace(strings.TrimPrefix(strings.TrimSpace(text), "admin password:"))
|
||||||
|
if len(pass) != 20 {
|
||||||
|
t.Fatalf("password len=%d value=%q", len(pass), pass)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err2 := cmdAdminInit(nil); err2 == nil || !strings.Contains(err2.Error(), "already initialized") {
|
||||||
|
t.Fatalf("want already initialized, got %v", err2)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err2 := cmdAdminSetPassword([]string{"--password", "short"}); err2 == nil {
|
||||||
|
t.Fatal("expected short password error")
|
||||||
|
}
|
||||||
|
if err2 := cmdAdminSetPassword([]string{"--password", "long-enough-password"}); err2 != nil {
|
||||||
|
t.Fatal(err2)
|
||||||
|
}
|
||||||
|
|
||||||
|
db, openErr := store.Open(dir, "FULL")
|
||||||
|
if openErr != nil {
|
||||||
|
t.Fatal(openErr)
|
||||||
|
}
|
||||||
|
defer func() { _ = db.Close() }()
|
||||||
|
ok, hasErr := store.HasAdminPassword(context.Background(), db.Write)
|
||||||
|
if hasErr != nil || !ok {
|
||||||
|
t.Fatalf("has admin: ok=%v err=%v", ok, hasErr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBackupVacuumInto(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
path := writeTestConfig(t, dir)
|
||||||
|
t.Setenv("NIXMSG_CONFIG", path)
|
||||||
|
initAdminForTest(t, dir)
|
||||||
|
|
||||||
|
out := filepath.Join(dir, "backup", "copy.db")
|
||||||
|
if err := os.MkdirAll(filepath.Dir(out), 0o755); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := cmdBackup([]string{"--out", out}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
st, err := os.Stat(out)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if st.Size() == 0 {
|
||||||
|
t.Fatal("backup empty")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHealthcheck(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
cfgPath := writeTestConfig(t, dir)
|
||||||
|
initAdminForTest(t, dir)
|
||||||
|
t.Setenv("NIXMSG_CONFIG", cfgPath)
|
||||||
|
|
||||||
|
cfg, err := loadAndValidateConfig()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
errCh := make(chan error, 1)
|
||||||
|
go func() { errCh <- runServe(ctx, cfg) }()
|
||||||
|
|
||||||
|
deadline := time.Now().Add(10 * time.Second)
|
||||||
|
for time.Now().Before(deadline) {
|
||||||
|
if b, readErr := os.ReadFile(filepath.Join(dir, "listen.addr")); readErr == nil && strings.TrimSpace(string(b)) != "" {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
time.Sleep(20 * time.Millisecond)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := cmdHealthcheck(nil); err != nil {
|
||||||
|
cancel()
|
||||||
|
<-errCh
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
cancel()
|
||||||
|
<-errCh
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseSetPasswordArgs(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
pass, err := parseSetPasswordArgs([]string{"--password", "abcdefghijkl"}, nil)
|
||||||
|
if err != nil || pass != "abcdefghijkl" {
|
||||||
|
t.Fatalf("pass=%q err=%v", pass, err)
|
||||||
|
}
|
||||||
|
pass, err = parseSetPasswordArgs([]string{"twelvechars!!"}, nil)
|
||||||
|
if err != nil || pass != "twelvechars!!" {
|
||||||
|
t.Fatalf("pass=%q err=%v", pass, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,38 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func cmdHealthcheck(_ []string) error {
|
||||||
|
cfg, err := loadAndValidateConfig()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
addrBytes, err := os.ReadFile(filepath.Join(cfg.DataDir, "listen.addr"))
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("read listen.addr: %w", err)
|
||||||
|
}
|
||||||
|
addr := strings.TrimSpace(string(addrBytes))
|
||||||
|
if addr == "" {
|
||||||
|
return fmt.Errorf("listen.addr is empty")
|
||||||
|
}
|
||||||
|
url := "http://" + addr + "/healthz"
|
||||||
|
client := &http.Client{Timeout: 5 * time.Second}
|
||||||
|
resp, err := client.Get(url)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer func() { _ = resp.Body.Close() }()
|
||||||
|
body, _ := io.ReadAll(io.LimitReader(resp.Body, 1024))
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
return fmt.Errorf("healthz status %d: %s", resp.StatusCode, strings.TrimSpace(string(body)))
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
+16
-4
@@ -10,16 +10,28 @@ func main() {
|
|||||||
fmt.Fprintln(os.Stderr, "usage: nixmsg <command>")
|
fmt.Fprintln(os.Stderr, "usage: nixmsg <command>")
|
||||||
os.Exit(2)
|
os.Exit(2)
|
||||||
}
|
}
|
||||||
|
var err error
|
||||||
switch os.Args[1] {
|
switch os.Args[1] {
|
||||||
case "version":
|
case "version":
|
||||||
cmdVersion(os.Args[2:])
|
cmdVersion(os.Args[2:])
|
||||||
case "serve":
|
case "serve":
|
||||||
if err := cmdServe(os.Args[2:]); err != nil {
|
err = cmdServe(os.Args[2:])
|
||||||
fmt.Fprintln(os.Stderr, err)
|
// P-WIRE-BEGIN
|
||||||
os.Exit(1)
|
case "admin":
|
||||||
}
|
err = cmdAdmin(os.Args[2:])
|
||||||
|
case "backup":
|
||||||
|
err = cmdBackup(os.Args[2:])
|
||||||
|
case "check-config":
|
||||||
|
err = cmdCheckConfig(os.Args[2:])
|
||||||
|
case "healthcheck":
|
||||||
|
err = cmdHealthcheck(os.Args[2:])
|
||||||
|
// P-WIRE-END
|
||||||
default:
|
default:
|
||||||
fmt.Fprintf(os.Stderr, "unknown command: %s\n", os.Args[1])
|
fmt.Fprintf(os.Stderr, "unknown command: %s\n", os.Args[1])
|
||||||
os.Exit(2)
|
os.Exit(2)
|
||||||
}
|
}
|
||||||
|
if err != nil {
|
||||||
|
fmt.Fprintln(os.Stderr, err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+55
-1
@@ -4,11 +4,13 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"os"
|
"os"
|
||||||
"os/signal"
|
"os/signal"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
"syscall"
|
"syscall"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -23,6 +25,12 @@ func cmdServe(_ []string) error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
// P-WIRE-BEGIN
|
||||||
|
if err := cfg.Validate(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
setupJSONLogger(cfg.Log)
|
||||||
|
// P-WIRE-END
|
||||||
ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
|
ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
|
||||||
defer stop()
|
defer stop()
|
||||||
return runServe(ctx, cfg)
|
return runServe(ctx, cfg)
|
||||||
@@ -41,6 +49,16 @@ func runServe(ctx context.Context, cfg config.Config) error {
|
|||||||
}
|
}
|
||||||
defer func() { _ = db.Close() }()
|
defer func() { _ = db.Close() }()
|
||||||
|
|
||||||
|
// P-WIRE-BEGIN
|
||||||
|
ok, err := store.HasAdminPassword(ctx, db.Write)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if !ok {
|
||||||
|
return errors.New("admin password not initialized; run: nixmsg admin init")
|
||||||
|
}
|
||||||
|
// P-WIRE-END
|
||||||
|
|
||||||
// 启动恢复入口已挂上(假实现为空操作);M 线替换 message.Service 后生效。
|
// 启动恢复入口已挂上(假实现为空操作);M 线替换 message.Service 后生效。
|
||||||
if recoverErr := deps.Messages.RecoverOnStart(ctx); recoverErr != nil {
|
if recoverErr := deps.Messages.RecoverOnStart(ctx); recoverErr != nil {
|
||||||
return fmt.Errorf("message recover: %w", recoverErr)
|
return fmt.Errorf("message recover: %w", recoverErr)
|
||||||
@@ -62,6 +80,17 @@ func runServe(ctx context.Context, cfg config.Config) error {
|
|||||||
w.WriteHeader(http.StatusOK)
|
w.WriteHeader(http.StatusOK)
|
||||||
_, _ = w.Write([]byte("ok"))
|
_, _ = w.Write([]byte("ok"))
|
||||||
})
|
})
|
||||||
|
// P-WIRE-BEGIN
|
||||||
|
mux.HandleFunc("GET /readyz", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if readyErr := db.Ready(r.Context()); readyErr != nil {
|
||||||
|
slog.Error("readyz failed", "err", readyErr)
|
||||||
|
http.Error(w, "not ready", http.StatusServiceUnavailable)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
_, _ = w.Write([]byte("ok"))
|
||||||
|
})
|
||||||
|
// P-WIRE-END
|
||||||
// 确保前端资源被链接进二进制;完整静态托管由后续任务完善。
|
// 确保前端资源被链接进二进制;完整静态托管由后续任务完善。
|
||||||
_ = web.Dist()
|
_ = web.Dist()
|
||||||
|
|
||||||
@@ -87,9 +116,17 @@ func runServe(ctx context.Context, cfg config.Config) error {
|
|||||||
|
|
||||||
select {
|
select {
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
shutdownCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
// P-WIRE-BEGIN
|
||||||
|
// 先停止接受新连接,再等写队列最多 10 秒,然后断开并退出。
|
||||||
|
shutdownCtx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
_ = srv.Shutdown(shutdownCtx)
|
_ = srv.Shutdown(shutdownCtx)
|
||||||
|
drainCtx, drainCancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||||
|
defer drainCancel()
|
||||||
|
if drainErr := db.Queue.Drain(drainCtx); drainErr != nil && !errors.Is(drainErr, context.DeadlineExceeded) {
|
||||||
|
slog.Error("write queue drain", "err", drainErr)
|
||||||
|
}
|
||||||
|
// P-WIRE-END
|
||||||
serveErr := <-errCh
|
serveErr := <-errCh
|
||||||
if serveErr != nil && !errors.Is(serveErr, http.ErrServerClosed) {
|
if serveErr != nil && !errors.Is(serveErr, http.ErrServerClosed) {
|
||||||
return serveErr
|
return serveErr
|
||||||
@@ -107,3 +144,20 @@ func writeListenAddr(dataDir, addr string) error {
|
|||||||
path := filepath.Join(dataDir, "listen.addr")
|
path := filepath.Join(dataDir, "listen.addr")
|
||||||
return os.WriteFile(path, []byte(addr+"\n"), 0o644)
|
return os.WriteFile(path, []byte(addr+"\n"), 0o644)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// P-WIRE-BEGIN
|
||||||
|
func setupJSONLogger(cfg config.LogConfig) {
|
||||||
|
level := slog.LevelInfo
|
||||||
|
switch strings.ToLower(strings.TrimSpace(cfg.Level)) {
|
||||||
|
case "debug":
|
||||||
|
level = slog.LevelDebug
|
||||||
|
case "warn", "warning":
|
||||||
|
level = slog.LevelWarn
|
||||||
|
case "error":
|
||||||
|
level = slog.LevelError
|
||||||
|
}
|
||||||
|
h := slog.NewJSONHandler(os.Stdout, &slog.HandlerOptions{Level: level})
|
||||||
|
slog.SetDefault(slog.New(h))
|
||||||
|
}
|
||||||
|
|
||||||
|
// P-WIRE-END
|
||||||
|
|||||||
@@ -10,23 +10,51 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||||
"git.asio.asia/nixevol/NixMsg/internal/config"
|
"git.asio.asia/nixevol/NixMsg/internal/config"
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
func writeTestConfig(t *testing.T, dataDir string) string {
|
||||||
|
t.Helper()
|
||||||
|
cfgPath := filepath.Join(dataDir, "config.yaml")
|
||||||
|
cfgYAML := []byte("listen: \"127.0.0.1:0\"\ndata_dir: \"" + filepath.ToSlash(dataDir) + "\"\nlog:\n level: error\n")
|
||||||
|
if err := os.WriteFile(cfgPath, cfgYAML, 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return cfgPath
|
||||||
|
}
|
||||||
|
|
||||||
|
func initAdminForTest(t *testing.T, dataDir string) {
|
||||||
|
t.Helper()
|
||||||
|
db, err := store.Open(dataDir, "FULL")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer func() { _ = db.Close() }()
|
||||||
|
phc, err := auth.HashPassword("test-admin-password-xx")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := store.SetAdminPasswordHash(context.Background(), db.Write, phc); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestServeHealthzAndListenAddr(t *testing.T) {
|
func TestServeHealthzAndListenAddr(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
dataDir := t.TempDir()
|
dataDir := t.TempDir()
|
||||||
cfgPath := filepath.Join(dataDir, "config.yaml")
|
cfgPath := writeTestConfig(t, dataDir)
|
||||||
cfgYAML := []byte("listen: \"127.0.0.1:0\"\ndata_dir: \"" + filepath.ToSlash(dataDir) + "\"\n")
|
initAdminForTest(t, dataDir)
|
||||||
if err := os.WriteFile(cfgPath, cfgYAML, 0o644); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
cfg, err := config.Load(cfgPath)
|
cfg, err := config.Load(cfgPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
|
if vErr := cfg.Validate(); vErr != nil {
|
||||||
|
t.Fatal(vErr)
|
||||||
|
}
|
||||||
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
defer cancel()
|
defer cancel()
|
||||||
@@ -37,7 +65,7 @@ func TestServeHealthzAndListenAddr(t *testing.T) {
|
|||||||
}()
|
}()
|
||||||
|
|
||||||
var addr string
|
var addr string
|
||||||
deadline := time.Now().Add(5 * time.Second)
|
deadline := time.Now().Add(10 * time.Second)
|
||||||
for time.Now().Before(deadline) {
|
for time.Now().Before(deadline) {
|
||||||
b, readErr := os.ReadFile(filepath.Join(dataDir, "listen.addr"))
|
b, readErr := os.ReadFile(filepath.Join(dataDir, "listen.addr"))
|
||||||
if readErr == nil {
|
if readErr == nil {
|
||||||
@@ -65,6 +93,15 @@ func TestServeHealthzAndListenAddr(t *testing.T) {
|
|||||||
t.Fatalf("unexpected body: %q", body)
|
t.Fatalf("unexpected body: %q", body)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
ready, err := http.Get("http://" + addr + "/readyz")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("readyz: %v", err)
|
||||||
|
}
|
||||||
|
defer func() { _ = ready.Body.Close() }()
|
||||||
|
if ready.StatusCode != http.StatusOK {
|
||||||
|
t.Fatalf("readyz status=%d", ready.StatusCode)
|
||||||
|
}
|
||||||
|
|
||||||
if _, err := os.Stat(filepath.Join(dataDir, "nixmsg.db")); err != nil {
|
if _, err := os.Stat(filepath.Join(dataDir, "nixmsg.db")); err != nil {
|
||||||
t.Fatalf("db missing: %v", err)
|
t.Fatalf("db missing: %v", err)
|
||||||
}
|
}
|
||||||
@@ -75,7 +112,23 @@ func TestServeHealthzAndListenAddr(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("serve exit: %v", err)
|
t.Fatalf("serve exit: %v", err)
|
||||||
}
|
}
|
||||||
case <-time.After(5 * time.Second):
|
case <-time.After(10 * time.Second):
|
||||||
t.Fatal("serve did not stop")
|
t.Fatal("serve did not stop")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestServeRejectsWithoutAdmin(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
dataDir := t.TempDir()
|
||||||
|
cfgPath := writeTestConfig(t, dataDir)
|
||||||
|
cfg, err := config.Load(cfgPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
err = runServe(ctx, cfg)
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "admin password") {
|
||||||
|
t.Fatalf("want admin init error, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+185
-4
@@ -152,19 +152,200 @@
|
|||||||
|
|
||||||
## 平台 P
|
## 平台 P
|
||||||
|
|
||||||
暂无。
|
### P1 2026-09-30
|
||||||
|
|
||||||
|
1. **admin set-password 传参方式**
|
||||||
|
- 原条款:DEVELOPMENT 11.2 仅列命令名,未规定密码如何传入。
|
||||||
|
- 实际做法:支持 `--password <pwd>`、位置参数,或从 stdin 读一行;最短 12 位。
|
||||||
|
- 原因:自动化测试与 Docker 非交互环境需要非交互传参。
|
||||||
|
- 备选方案:仅交互式 prompt。
|
||||||
|
- 影响:运维文档需写明推荐用 `--password` 或管道,勿把密码写进 shell 历史时可改用 stdin。
|
||||||
|
|
||||||
|
2. **admin init 密码字符集**
|
||||||
|
- 原条款:生成 20 位密码,未规定字符集。
|
||||||
|
- 实际做法:从去掉易混字符(0/O/1/I/l)的字母数字中均匀抽样 20 位。
|
||||||
|
- 原因:终端抄写友好。
|
||||||
|
- 备选方案:全 ASCII 可打印字符。
|
||||||
|
- 影响:熵略低于全字符集,对 20 位仍足够。
|
||||||
|
|
||||||
|
3. **PHC 哈希辅助先放在 auth,池在 P3**
|
||||||
|
- 原条款:argon2 并发池属 P3;admin init 需落库哈希属 P1。
|
||||||
|
- 实际做法:P1 在 `internal/auth` 提供 `HashPassword`/`VerifyPassword`(PHC),admin 命令直接调用;P3 再用同参数实现带并发上限的 `HashPool`。
|
||||||
|
- 原因:避免 P1 在 cmd 内复制算法,又不等到 P3 才做 init。
|
||||||
|
- 备选方案:P1 cmd 内联 argon2;或 P1/P3 合并提交。
|
||||||
|
- 影响:P3 合入后 admin 可改为走 HashPool(非必须,单次 init 无并发压力)。
|
||||||
|
|
||||||
|
4. **优雅停机顺序**
|
||||||
|
- 原条款:停接受 → 写队列最多 10 秒 → 断开连接退出。
|
||||||
|
- 实际做法:`http.Server.Shutdown`(先停接受并等待进行中的 HTTP)后,再 `Queue.Drain` 最多 10 秒;本期尚无长连接表,断开连接由 Shutdown 覆盖。
|
||||||
|
- 原因:当前 main 仅有 HTTP 健康检查;N 线接入后应在 Drain 前后显式踢连接。
|
||||||
|
- 备选方案:先关 listener、Drain、再 Shutdown。
|
||||||
|
- 影响:有 MQTT 长连接后需 N/P 联调停机路径。
|
||||||
|
|
||||||
|
### P2 2026-09-30
|
||||||
|
|
||||||
|
1. **合并写入队列替换 T0.3 简单实现**
|
||||||
|
- 原条款:DEVELOPMENT 7.2;TASKS P2。
|
||||||
|
- 实际做法:单写 goroutine;批次上限 256 或等待 2ms;每操作用 `SAVEPOINT`/`ROLLBACK TO`/`RELEASE`;Begin/Commit/SAVEPOINT 基础设施失败返回 `errors.Join(ErrBusy, err)` 并令 `IsReady()=false`;业务操作错误只回滚该 SAVEPOINT,不标 busy。
|
||||||
|
- 原因:满足每秒约 200 条写入的落盘合并需求。
|
||||||
|
- 备选方案:按固定时间窗无条件合并。
|
||||||
|
- 影响:调用方需用 `errors.Is(err, store.ErrBusy)` 映射协议 `busy`;`/readyz` 读 `DB.Ready`。
|
||||||
|
|
||||||
|
2. **调用方 context 取消与已入队任务**
|
||||||
|
- 原条款:未规定入队后取消。
|
||||||
|
- 实际做法:入队前检查 ctx;批次执行前再检查;若调用方在等待结果时取消,最多再等 30 秒取结果以免泄漏。
|
||||||
|
- 原因:写 goroutine 仍可能已执行该操作,不能静默丢结果。
|
||||||
|
- 备选方案:取消即从队列摘除(需可取消数据结构)。
|
||||||
|
- 影响:极端取消场景下调用方可能多等一会儿。
|
||||||
|
|
||||||
|
### P3 2026-09-30
|
||||||
|
|
||||||
|
1. **管理员/注册锁定阈值沿用登录 IP 档**
|
||||||
|
- 原条款:登录/对话密码阈值写清;管理员登录与注册安全码仅写「临时锁定」,未给数字。
|
||||||
|
- 实际做法:`LockAdminIP`、`LockRegisterIP`、`LockTalkPair`/`LockLoginEndpointIP` 均为 5 分钟窗口 10 次、锁 5 分钟;`LockTalkTarget`/`LockLoginEndpoint` 为 1 小时 50 次、锁 1 小时。
|
||||||
|
- 原因:与 PRD D18/F23「默认 5 分钟 10 次」叙述一致。
|
||||||
|
- 备选方案:管理员单独更严阈值。
|
||||||
|
- 影响:A/I/N 线直接用 `LoginLocks` 即可。
|
||||||
|
|
||||||
|
2. **真实 auth 实现未改 wire.go**
|
||||||
|
- 原条款:平台不改 `wire.go`;P3 实现接口。
|
||||||
|
- 实际做法:提供 `NewPool`/`NewSessionTokens`/`NewAPITokens`/`NewLoginLocks`;`wire()` 仍用 Stub,由总控或各线接线时替换。
|
||||||
|
- 原因:分工禁止改 wire.go。
|
||||||
|
- 备选方案:在 serve 旁路替换(会绕过 wire)。
|
||||||
|
- 影响:合入后需有一次接线才能在进程内用上真实哈希池。
|
||||||
|
|
||||||
|
3. **令牌随机部分用 RawURLEncoding**
|
||||||
|
- 原条款:32 字节随机数的 base64url。
|
||||||
|
- 实际做法:`encoding/base64.RawURLEncoding`(无 padding)。
|
||||||
|
- 原因:URL/Header 友好,与常见 token 惯例一致。
|
||||||
|
- 备选方案:StdEncoding 带 padding。
|
||||||
|
- 影响:SDK/文档示例需无 `=` 结尾。
|
||||||
|
|
||||||
|
### P4 2026-09-30
|
||||||
|
|
||||||
|
1. **指标包放在 `internal/metrics`**
|
||||||
|
- 原条款:TASKS 分工表未列 metrics 目录;P4 要求用 prometheus/client_golang 建注册表。
|
||||||
|
- 实际做法:新建 `internal/metrics`,提供 `New`/`Handler`/`Registry` 字段供各线打点;不在 `serve` 挂路由(访问规则属 A 线)。
|
||||||
|
- 原因:不宜塞进 auth/store/config。
|
||||||
|
- 备选方案:放 `internal/httpx`(A 线目录)。
|
||||||
|
- 影响:A/N 接线时 import 本包并挂 `/metrics`。
|
||||||
|
|
||||||
|
2. **指标命名**
|
||||||
|
- 原条款:列了指标含义,未规定 Prometheus 名字。
|
||||||
|
- 实际做法:`nixmsg_connections{transport}`、`nixmsg_endpoints`、`nixmsg_deliveries_pending`、`nixmsg_messages_scheduled`、`nixmsg_dispatch_to_push_duration_seconds`、`nixmsg_ack_duration_seconds`、`nixmsg_write_queue_length`、`nixmsg_write_batch_commit_duration_seconds`、`nixmsg_password_hash_queue_length`、`nixmsg_errors_total{code}`。
|
||||||
|
- 原因:固定可抓取文本便于联调。
|
||||||
|
- 备选方案:更短前缀或 HistogramVec。
|
||||||
|
- 影响:仪表盘按上述名字配置。
|
||||||
|
|
||||||
## 连接 N
|
## 连接 N
|
||||||
|
|
||||||
暂无。
|
### N1 / N2 2026-09-30
|
||||||
|
|
||||||
|
1. **未接线 `cmd/nixmsg`**
|
||||||
|
- 原条款:serve 最终应挂上端口识别、broker、`/mqtt`。
|
||||||
|
- 实际做法:本任务只交付 `internal/listener`、`internal/broker`;按总控要求不改 `cmd/nixmsg`。
|
||||||
|
- 原因:避免与平台/总控并行改 wire 冲突;合并时再接线。
|
||||||
|
- 备选方案:本分支顺带改 `wire.go`(与指令冲突)。
|
||||||
|
- 影响:当前 `serve` 仍是 T0.4 的简单 `/healthz` 监听,不含 MQTT。
|
||||||
|
|
||||||
|
2. **`listen.addr` / `admin.addr` 仅端口为 0 时写入**
|
||||||
|
- 原条款:DEVELOPMENT 4.1「端口写 0 时」写地址文件;T0.1 偏差曾改为 always write。
|
||||||
|
- 实际做法:`listener.Server` 仅当配置地址端口为 `0` 时写 `listen.addr` / `admin.addr`。
|
||||||
|
- 原因:本任务说明与 DEVELOPMENT 4.1 字面一致;T0.1 的 always write 在 `cmd/nixmsg`,本线未改。
|
||||||
|
- 备选方案:接线时统一为 always write 以兼容 harness。
|
||||||
|
- 影响:固定端口场景下 harness 若只读地址文件会读不到;接线时建议沿用 T0.1 超集或改 harness。
|
||||||
|
|
||||||
|
3. **登录校验为可替换接口,默认拒绝**
|
||||||
|
- 原条款:第 5 节完整会话令牌/密码/锁定属 N3。
|
||||||
|
- 实际做法:`broker.Authenticator` 接口 + 默认 `RejectAuthenticator`;内部错误在 `OnConnect` 返回 error;测试提供 `AllowAuthenticator`。
|
||||||
|
- 原因:N3 范围;N2 需可跑通装配与钩子。
|
||||||
|
- 备选方案:N2 内做假登录表(超出范围)。
|
||||||
|
- 影响:真实端连不上直到 N3;总控接线时注入 Authenticator。
|
||||||
|
|
||||||
|
4. **大帧并发名额释放策略**
|
||||||
|
- 原条款:DEVELOPMENT 7.5 大于 64KiB 全局同时不超过 64;PUBACK / 超时 / 断线释放。
|
||||||
|
- 实际做法:发布前申请名额;QoS 0 发布成功立即释放;QoS 1 在 `OnQosComplete` 且 payload>64KiB 时释放,断线 `releaseAllLarge`;未单独做「确认超时」计时释放(确认超时属 M 线推送循环)。
|
||||||
|
- 原因:N2 无投递确认计时器;与 M 线推送超时释放衔接。
|
||||||
|
- 备选方案:broker 内对大帧自建超时(与 M 重复)。
|
||||||
|
- 影响:若客户端永不 PUBACK 且不断线,名额可能占满直到断开;M 线超时踢线或回调 Disconnect 可释放。
|
||||||
|
|
||||||
|
5. **`OnPublishDropped` 仅打日志**
|
||||||
|
- 原条款:清「已推送」标记并 1 秒后重推。
|
||||||
|
- 实际做法:钩子记录 debug 日志;清标记/重推留给消息 M。
|
||||||
|
- 原因:投递状态在 M/store,N2 无投递表。
|
||||||
|
- 备选方案:N2 暴露回调给 M 注册。
|
||||||
|
- 影响:接线后 M 需订阅或包装该钩子;当前接口可后续加 `OnPublishDropped` 回调字段。
|
||||||
|
|
||||||
## 消息 M
|
## 消息 M
|
||||||
|
|
||||||
暂无。
|
### M1 2026-09-30
|
||||||
|
|
||||||
|
1. **提交时分发做成最小正确版**
|
||||||
|
- 原条款:DEVELOPMENT 7.3 步骤 8 / 7.4:`send_at` 已到则同一写操作内完整分发(停用拒绝、`queue_full`、`expire_at`/宽限、无接收者 `completed`、回执等)。
|
||||||
|
- 实际做法:单聊只插一条 `pending`;群按当时 `group_members` 去掉发送者各插 `pending`;消息改为 `dispatched`。不设 `expire_at`,不检查接收端配额/在线/停用,不因无接收者改为 `completed`,不写回执,不唤醒推送循环。
|
||||||
|
- 原因:M1 范围是提交;完整分发与推送属 M2。
|
||||||
|
- 备选方案:M1 直接实现完整 7.4(抢 M2)。
|
||||||
|
- 影响:到点消息已有投递行,但停用成员仍会有 `pending`;无成员群仍为 `dispatched` 且无投递;推送需等 M2。
|
||||||
|
|
||||||
|
2. **请求频率突发容量写死为 100**
|
||||||
|
- 原条款:DEVELOPMENT 6.10 每端每秒 50、突发 100;配置示例仅有 `requests_per_second`。
|
||||||
|
- 实际做法:`Limits.RequestBurst` 默认 100;`requests_per_second<=0` 时不限速(便于测试)。速率桶挂在 `message.App` 的 `Submit` 入口;`ack`/`receipt_ack` 尚未实现故未接桶。
|
||||||
|
- 原因:配置无独立 burst 字段。
|
||||||
|
- 备选方案:配置增加 `request_burst`;由连接线在上行统一限流。
|
||||||
|
- 影响:改 `requests_per_second` 不改突发;正式接线后若 N 线也限流可能双重计数。
|
||||||
|
|
||||||
|
3. **未接线 `cmd/nixmsg`**
|
||||||
|
- 原条款:可替换 T0.4 假实现。
|
||||||
|
- 实际做法:新增 `message.App` 实现 `Submit`;保留 `Stub`;按任务隔离要求未改 `cmd/nixmsg`/`wire.go`。
|
||||||
|
- 原因:本任务禁止改 `cmd/nixmsg`;总控接线或后续任务再换。
|
||||||
|
- 备选方案:本任务直接改 `wire.go`。
|
||||||
|
- 影响:进程内仍用 Stub,需显式构造 `message.New` 才能用真实提交。
|
||||||
|
|
||||||
|
4. **防重键在、消息行已删时返回 `not_found`**
|
||||||
|
- 原条款:防重命中返回原消息当前状态;未写明消息行已被清理时的提交重试行为(状态查询为 `not_found`)。
|
||||||
|
- 实际做法:`send_keys` 指纹相同但 `messages` 无行时返回 `not_found`。
|
||||||
|
- 原因:无法构造 `send_at`/`state`。
|
||||||
|
- 备选方案:在 `send_keys` 冗余存结果快照。
|
||||||
|
- 影响:保留期过后的重试不再幂等成功。
|
||||||
|
|
||||||
## 身份 I
|
## 身份 I
|
||||||
|
|
||||||
暂无。
|
### I1 2026-09-30
|
||||||
|
|
||||||
|
1. **注册做成可挂载 Handler,不改 cmd/listener**
|
||||||
|
- 原条款:TASKS I1 / DEVELOPMENT 6.9 在 `listen` 上提供 `POST /api/client/register`;依赖 N1 端口识别。
|
||||||
|
- 实际做法:`identity.NewRegisterHandler` / `identity.NewServer().Handler()` 返回 `http.Handler`,由接线方 `mux.Handle("/api/client/register", h)`;本线不改 `cmd/nixmsg`、`internal/listener`(N 线未合入)。
|
||||||
|
- 原因:隔离交付,避免抢 N/P 接线。
|
||||||
|
- 备选方案:本线直接改 `wire.go` 挂路由。
|
||||||
|
- 影响:合入后需总控或 N/A 接线才对外可访问。
|
||||||
|
|
||||||
|
2. **密码哈希与锁定走 auth 接口,本分支用可替换假实现测**
|
||||||
|
- 原条款:依赖 P3 argon2 池与锁定计数器。
|
||||||
|
- 实际做法:`RegisterConfig.Hash`/`Locks` 注入 `auth.HashPool`、`auth.LoginLocks`;测试用 `auth.NewStubHashPool` + 仅实现 `LockRegisterIP`(5 分钟 10 次)的测试锁定器,不在本线重写 argon2。
|
||||||
|
- 原因:P3 尚未在本分支。
|
||||||
|
- 备选方案:等 P3 合入后再写 I1。
|
||||||
|
- 影响:生产须注入 P3 实现;StubLoginLocks 永不锁定,不能直接用于开放注册。
|
||||||
|
|
||||||
|
3. **settings 开关取值**
|
||||||
|
- 原条款:`settings.registration_enabled`,未规定字符串字面量。
|
||||||
|
- 实际做法:`1`/`true`/`yes`/`on`(大小写不敏感)视为开启,其余(含缺省)关闭;安全码键 `registration_code`。
|
||||||
|
- 原因:与 store 测试写入的 `"0"`/`"1"` 对齐并兼容常见布尔字面量。
|
||||||
|
- 备选方案:仅认 `"1"`。
|
||||||
|
- 影响:A 线写注册设置时宜写 `"1"`/`"0"`。
|
||||||
|
|
||||||
|
4. **客户端 IP**
|
||||||
|
- 原条款:DEVELOPMENT 4.5 受信任代理下用 `X-Forwarded-For`。
|
||||||
|
- 实际做法:Handler 默认取 `RemoteAddr` 的 host;可通过 `RegisterConfig.ClientIP` 注入。本线不做 `trusted_proxies` 解析(属 listener/接线)。
|
||||||
|
- 原因:不改 listener;代理 IP 应由外层在挂载前算好或注入。
|
||||||
|
- 备选方案:在 identity 内复制 4.5 逻辑。
|
||||||
|
- 影响:经代理部署时接线方必须注入真实 IP,否则锁定按直连 IP 计。
|
||||||
|
|
||||||
|
5. **生成登录密码长度**
|
||||||
|
- 原条款:F01 留空则生成,8–128 字符,不以 `nst_` 开头;未规定生成长度。
|
||||||
|
- 实际做法:生成 20 位字母数字;若偶然以 `nst_` 开头则重抽。
|
||||||
|
- 原因:与管理员 init 量级接近,满足规则。
|
||||||
|
- 备选方案:16/32 位。
|
||||||
|
- 影响:无产品行为差异。
|
||||||
|
|
||||||
## 后台接口 A
|
## 后台接口 A
|
||||||
|
|
||||||
|
|||||||
@@ -3,17 +3,30 @@ module git.asio.asia/nixevol/NixMsg
|
|||||||
go 1.27
|
go 1.27
|
||||||
|
|
||||||
require (
|
require (
|
||||||
|
github.com/coder/websocket v1.8.14
|
||||||
|
github.com/mochi-mqtt/server/v2 v2.7.9
|
||||||
|
github.com/prometheus/client_golang v1.24.1
|
||||||
go.yaml.in/yaml/v3 v3.0.5
|
go.yaml.in/yaml/v3 v3.0.5
|
||||||
|
golang.org/x/crypto v0.57.0
|
||||||
modernc.org/sqlite v1.60.1
|
modernc.org/sqlite v1.60.1
|
||||||
)
|
)
|
||||||
|
|
||||||
require (
|
require (
|
||||||
|
github.com/beorn7/perks v1.0.1 // indirect
|
||||||
|
github.com/cespare/xxhash/v2 v2.3.0 // indirect
|
||||||
github.com/dustin/go-humanize v1.0.1 // indirect
|
github.com/dustin/go-humanize v1.0.1 // indirect
|
||||||
github.com/google/uuid v1.6.0 // indirect
|
github.com/google/uuid v1.6.0 // indirect
|
||||||
|
github.com/gorilla/websocket v1.5.0 // indirect
|
||||||
github.com/mattn/go-isatty v0.0.24 // indirect
|
github.com/mattn/go-isatty v0.0.24 // indirect
|
||||||
|
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect
|
||||||
github.com/ncruces/go-strftime v1.0.0 // indirect
|
github.com/ncruces/go-strftime v1.0.0 // indirect
|
||||||
|
github.com/prometheus/client_model v0.6.2 // indirect
|
||||||
|
github.com/prometheus/common v0.70.1 // indirect
|
||||||
|
github.com/prometheus/procfs v0.21.1 // indirect
|
||||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
|
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
|
||||||
|
github.com/rs/xid v1.4.0 // indirect
|
||||||
golang.org/x/sys v0.48.0 // indirect
|
golang.org/x/sys v0.48.0 // indirect
|
||||||
|
google.golang.org/protobuf v1.36.11 // indirect
|
||||||
modernc.org/libc v1.77.1 // indirect
|
modernc.org/libc v1.77.1 // indirect
|
||||||
modernc.org/mathutil v1.7.1 // indirect
|
modernc.org/mathutil v1.7.1 // indirect
|
||||||
modernc.org/memory v1.12.1 // indirect
|
modernc.org/memory v1.12.1 // indirect
|
||||||
|
|||||||
@@ -1,19 +1,61 @@
|
|||||||
|
github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM=
|
||||||
|
github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw=
|
||||||
|
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
|
||||||
|
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
||||||
|
github.com/coder/websocket v1.8.14 h1:9L0p0iKiNOibykf283eHkKUHHrpG7f65OE3BhhO7v9g=
|
||||||
|
github.com/coder/websocket v1.8.14/go.mod h1:NX3SzP+inril6yawo5CQXx8+fk145lPDC6pumgx0mVg=
|
||||||
|
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
||||||
|
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||||
github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY=
|
github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY=
|
||||||
github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto=
|
github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto=
|
||||||
|
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
|
||||||
|
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
|
||||||
github.com/google/pprof v0.0.0-20260802141513-ef3492d7dac3 h1:LMLX+LgTNWpfvCBdFebv6EsYotImrt/Ppc5cXIriCSo=
|
github.com/google/pprof v0.0.0-20260802141513-ef3492d7dac3 h1:LMLX+LgTNWpfvCBdFebv6EsYotImrt/Ppc5cXIriCSo=
|
||||||
github.com/google/pprof v0.0.0-20260802141513-ef3492d7dac3/go.mod h1:jl5iWTm0/hd5PjEYEOuwAJ57L/CibdZfrqZ5XA5GrCk=
|
github.com/google/pprof v0.0.0-20260802141513-ef3492d7dac3/go.mod h1:jl5iWTm0/hd5PjEYEOuwAJ57L/CibdZfrqZ5XA5GrCk=
|
||||||
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
|
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
|
||||||
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
|
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
|
||||||
|
github.com/gorilla/websocket v1.5.0 h1:PPwGk2jz7EePpoHN/+ClbZu8SPxiqlu12wZP/3sWmnc=
|
||||||
|
github.com/gorilla/websocket v1.5.0/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
|
||||||
github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k=
|
github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k=
|
||||||
github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM=
|
github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM=
|
||||||
|
github.com/jinzhu/copier v0.3.5 h1:GlvfUwHk62RokgqVNvYsku0TATCF7bAHVwEXoBh3iJg=
|
||||||
|
github.com/jinzhu/copier v0.3.5/go.mod h1:DfbEm0FYsaqBcKcFuvmOZb218JkPGtvSHsKg8S8hyyg=
|
||||||
|
github.com/klauspost/compress v1.19.1 h1:VsB4HPswih7mmZ8WleSFQ75c/Ui1M4trX5oAsJnhSlk=
|
||||||
|
github.com/klauspost/compress v1.19.1/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ=
|
||||||
|
github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc=
|
||||||
|
github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw=
|
||||||
github.com/mattn/go-isatty v0.0.24 h1:tGZZoVgT/KiqK1c8ocVLeDS8BSWMRd47J3Lbz7vsReI=
|
github.com/mattn/go-isatty v0.0.24 h1:tGZZoVgT/KiqK1c8ocVLeDS8BSWMRd47J3Lbz7vsReI=
|
||||||
github.com/mattn/go-isatty v0.0.24/go.mod h1:nMCL3Zebbrt45jsMDgnfIwz6ydEQApk5oEI3HqDio6A=
|
github.com/mattn/go-isatty v0.0.24/go.mod h1:nMCL3Zebbrt45jsMDgnfIwz6ydEQApk5oEI3HqDio6A=
|
||||||
|
github.com/mochi-mqtt/server/v2 v2.7.9 h1:y0g4vrSLAag7T07l2oCzOa/+nKVLoazKEWAArwqBNYI=
|
||||||
|
github.com/mochi-mqtt/server/v2 v2.7.9/go.mod h1:lZD3j35AVNqJL5cezlnSkuG05c0FCHSsfAKSPBOSbqc=
|
||||||
|
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
|
||||||
|
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
|
||||||
github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w=
|
github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w=
|
||||||
github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls=
|
github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls=
|
||||||
|
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
||||||
|
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
|
||||||
|
github.com/prometheus/client_golang v1.24.1 h1:JnJkREXzWxUdCuPFpIWZiPispT9xVV59uiuyR2bPlnU=
|
||||||
|
github.com/prometheus/client_golang v1.24.1/go.mod h1:F+oSRECHg4sse5ucfYpYDeIv/hu68Zo0uoHKetWnzcE=
|
||||||
|
github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk=
|
||||||
|
github.com/prometheus/client_model v0.6.2/go.mod h1:y3m2F6Gdpfy6Ut/GBsUqTWZqCUvMVzSfMLjcu6wAwpE=
|
||||||
|
github.com/prometheus/common v0.70.1 h1:1HvjP4D5oL3t8RsPlwxA9onvvStjtIHYE5XuuwOi/PY=
|
||||||
|
github.com/prometheus/common v0.70.1/go.mod h1:VdFUQDMZK3VLkurFUVhia6uys/0suUp86TJz5qbJRhc=
|
||||||
|
github.com/prometheus/procfs v0.21.1 h1:GljZCt+zSTS+NZq88cyQ1LjZ+RCHp3uVuabBWA5+OJI=
|
||||||
|
github.com/prometheus/procfs v0.21.1/go.mod h1:aB55Cww9pdSJVHk0hUf0inxWyyjPogFIjmHKYgMKmtY=
|
||||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=
|
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=
|
||||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
|
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
|
||||||
|
github.com/rs/xid v1.4.0 h1:qd7wPTDkN6KQx2VmMBLrpHkiyQwgFXRnkOLacUiaSNY=
|
||||||
|
github.com/rs/xid v1.4.0/go.mod h1:trrq9SKmegXys3aeAKXMUTdJsYXVwGY3RLcfgqegfbg=
|
||||||
|
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
|
||||||
|
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
|
||||||
|
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
|
||||||
|
go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
|
||||||
|
go.yaml.in/yaml/v2 v2.4.4 h1:tuyd0P+2Ont/d6e2rl3be67goVK4R6deVxCUX5vyPaQ=
|
||||||
|
go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ=
|
||||||
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
|
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
|
||||||
go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg=
|
go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg=
|
||||||
|
golang.org/x/crypto v0.57.0 h1:3ZVCjf8Ggz7zneR/EHRVx68Ctf+2pmIMP2UFhh9cC6M=
|
||||||
|
golang.org/x/crypto v0.57.0/go.mod h1:Fdz0i5U6CoizGwLda9DttjSk6qlZo25zYNtR+ycvuZA=
|
||||||
golang.org/x/mod v0.41.0 h1:qJmnOUb4YB+FsEuM3HcWucdZASCPGhsX6uljO6pog0c=
|
golang.org/x/mod v0.41.0 h1:qJmnOUb4YB+FsEuM3HcWucdZASCPGhsX6uljO6pog0c=
|
||||||
golang.org/x/mod v0.41.0/go.mod h1:Ek9pY8RKWXwsWvd3rQiHYtMqkjSUV+s1Rj7j4H5Ur6o=
|
golang.org/x/mod v0.41.0/go.mod h1:Ek9pY8RKWXwsWvd3rQiHYtMqkjSUV+s1Rj7j4H5Ur6o=
|
||||||
golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk=
|
golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk=
|
||||||
@@ -22,6 +64,10 @@ golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo=
|
|||||||
golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og=
|
golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og=
|
||||||
golang.org/x/tools v0.50.0 h1:c2ifzfcuY7L90lZ2aKd8S4K2NpASF08SZx9ZuJkHmSU=
|
golang.org/x/tools v0.50.0 h1:c2ifzfcuY7L90lZ2aKd8S4K2NpASF08SZx9ZuJkHmSU=
|
||||||
golang.org/x/tools v0.50.0/go.mod h1:7ulVMw3831Mwi5EZD6RomGyffr4VFjuNYXf2BbCEAV0=
|
golang.org/x/tools v0.50.0/go.mod h1:7ulVMw3831Mwi5EZD6RomGyffr4VFjuNYXf2BbCEAV0=
|
||||||
|
google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
|
||||||
|
google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
|
||||||
|
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
||||||
|
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||||
modernc.org/cc/v4 v4.29.7 h1:q+NXGJ0bK3b4TXFYQQVr9pYETGnmwFWkrUzJnMya/Tg=
|
modernc.org/cc/v4 v4.29.7 h1:q+NXGJ0bK3b4TXFYQQVr9pYETGnmwFWkrUzJnMya/Tg=
|
||||||
modernc.org/cc/v4 v4.29.7/go.mod h1:OnovgIhbbMXMu1aISnJ0wvVD1KnW+cAUJkIrAWh+kVI=
|
modernc.org/cc/v4 v4.29.7/go.mod h1:OnovgIhbbMXMu1aISnJ0wvVD1KnW+cAUJkIrAWh+kVI=
|
||||||
modernc.org/ccgo/v4 v4.36.1 h1:ZNIUZAryN0UgnJwtyxrdEzcFc3yD4Cu4AzjfPXsLsIE=
|
modernc.org/ccgo/v4 v4.36.1 h1:ZNIUZAryN0UgnJwtyxrdEzcFc3yD4Cu4AzjfPXsLsIE=
|
||||||
|
|||||||
@@ -0,0 +1,417 @@
|
|||||||
|
package identity
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/subtle"
|
||||||
|
"database/sql"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
"unicode/utf8"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
maxRegisterBodyBytes = 4 * 1024
|
||||||
|
|
||||||
|
settingRegistrationEnabled = "registration_enabled"
|
||||||
|
settingRegistrationCode = "registration_code"
|
||||||
|
|
||||||
|
sourceSelf = "self"
|
||||||
|
|
||||||
|
idAlphabet = "abcdefghijklmnopqrstuvwxyz0123456789"
|
||||||
|
)
|
||||||
|
|
||||||
|
// APIError 是注册 HTTP/业务错误,带 HTTP 状态与协议错误码。
|
||||||
|
type APIError struct {
|
||||||
|
Status int
|
||||||
|
Code string
|
||||||
|
Message string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e *APIError) Error() string {
|
||||||
|
if e == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
if e.Message == "" {
|
||||||
|
return e.Code
|
||||||
|
}
|
||||||
|
return e.Code + ": " + e.Message
|
||||||
|
}
|
||||||
|
|
||||||
|
func apiErr(status int, code, msg string) *APIError {
|
||||||
|
return &APIError{Status: status, Code: code, Message: msg}
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterConfig 是可挂载注册处理器的依赖。
|
||||||
|
// Hash / Locks 用 auth 接口;P3 未合入时测试可注入 StubHashPool 与可锁定的 LoginLocks。
|
||||||
|
type RegisterConfig struct {
|
||||||
|
DB *store.DB
|
||||||
|
Hash auth.HashPool
|
||||||
|
Locks auth.LoginLocks
|
||||||
|
Logger *slog.Logger
|
||||||
|
// Now 可测;nil 则用 time.Now。
|
||||||
|
Now func() time.Time
|
||||||
|
// ClientIP 可测;nil 则从 RemoteAddr 取 host。
|
||||||
|
ClientIP func(*http.Request) string
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterHandler 处理 POST/OPTIONS /api/client/register(可挂到任意 ServeMux)。
|
||||||
|
type RegisterHandler struct {
|
||||||
|
cfg RegisterConfig
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewRegisterHandler 构造可挂载的注册 Handler。DB/Hash/Locks 必填。
|
||||||
|
func NewRegisterHandler(cfg RegisterConfig) *RegisterHandler {
|
||||||
|
if cfg.Logger == nil {
|
||||||
|
cfg.Logger = slog.Default()
|
||||||
|
}
|
||||||
|
if cfg.Now == nil {
|
||||||
|
cfg.Now = time.Now
|
||||||
|
}
|
||||||
|
if cfg.ClientIP == nil {
|
||||||
|
cfg.ClientIP = clientIPFromRemoteAddr
|
||||||
|
}
|
||||||
|
return &RegisterHandler{cfg: cfg}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ServeHTTP 实现 http.Handler。
|
||||||
|
func (h *RegisterHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||||
|
setCORS(w)
|
||||||
|
switch r.Method {
|
||||||
|
case http.MethodOptions:
|
||||||
|
w.Header().Set("Access-Control-Allow-Methods", "POST, OPTIONS")
|
||||||
|
w.Header().Set("Access-Control-Allow-Headers", "Content-Type")
|
||||||
|
w.WriteHeader(http.StatusNoContent)
|
||||||
|
return
|
||||||
|
case http.MethodPost:
|
||||||
|
h.handlePost(w, r)
|
||||||
|
default:
|
||||||
|
writeRegisterError(w, apiErr(http.StatusMethodNotAllowed, protocol.CodeBadRequest, "method not allowed"))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *RegisterHandler) handlePost(w http.ResponseWriter, r *http.Request) {
|
||||||
|
ip := h.cfg.ClientIP(r)
|
||||||
|
r.Body = http.MaxBytesReader(w, r.Body, maxRegisterBodyBytes)
|
||||||
|
body, err := io.ReadAll(r.Body)
|
||||||
|
if err != nil {
|
||||||
|
var maxErr *http.MaxBytesError
|
||||||
|
if errors.As(err, &maxErr) || errors.Is(err, io.ErrUnexpectedEOF) || isBodyTooLarge(err) {
|
||||||
|
h.logResult("bad_request", "", ip)
|
||||||
|
writeRegisterError(w, apiErr(http.StatusBadRequest, protocol.CodeBadRequest, "body too large"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
h.logResult("bad_request", "", ip)
|
||||||
|
writeRegisterError(w, apiErr(http.StatusBadRequest, protocol.CodeBadRequest, "read body failed"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
req, err := protocol.DecodeRegister(body)
|
||||||
|
if err != nil {
|
||||||
|
h.logResult("bad_request", "", ip)
|
||||||
|
writeRegisterError(w, apiErr(http.StatusBadRequest, protocol.CodeBadRequest, "invalid json"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
result, apiErr := h.register(r.Context(), req, ip)
|
||||||
|
if apiErr != nil {
|
||||||
|
id := ""
|
||||||
|
if req != nil {
|
||||||
|
id = req.ID
|
||||||
|
}
|
||||||
|
h.logResult(apiErr.Code, id, ip)
|
||||||
|
writeRegisterError(w, apiErr)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
h.logResult("ok", result.ID, ip)
|
||||||
|
writeRegisterOK(w, result)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *RegisterHandler) register(ctx context.Context, req *protocol.RegisterRequest, ip string) (RegisterResult, *APIError) {
|
||||||
|
enabled, storedCode, err := h.loadRegistrationSettings(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "settings unavailable")
|
||||||
|
}
|
||||||
|
if !enabled {
|
||||||
|
return RegisterResult{}, apiErr(http.StatusForbidden, protocol.CodeRegistrationClosed, "registration closed")
|
||||||
|
}
|
||||||
|
|
||||||
|
lockKey := auth.LockKey{Kind: auth.LockRegisterIP, IP: ip}
|
||||||
|
if locked, _ := h.cfg.Locks.Check(lockKey); locked {
|
||||||
|
return RegisterResult{}, apiErr(http.StatusTooManyRequests, protocol.CodeRateLimited, "rate limited")
|
||||||
|
}
|
||||||
|
|
||||||
|
if !constantTimeEqual(req.RegistrationCode, storedCode) {
|
||||||
|
h.cfg.Locks.Fail(lockKey)
|
||||||
|
return RegisterResult{}, apiErr(http.StatusForbidden, protocol.CodeRegistrationCodeInvalid, "registration code invalid")
|
||||||
|
}
|
||||||
|
|
||||||
|
if valErr := req.Validate(); valErr != nil {
|
||||||
|
code, msg := protocol.CodeBadRequest, valErr.Error()
|
||||||
|
var pe *protocol.Error
|
||||||
|
if errors.As(valErr, &pe) && pe != nil {
|
||||||
|
code, msg = pe.Code, pe.Message
|
||||||
|
}
|
||||||
|
return RegisterResult{}, apiErr(http.StatusBadRequest, code, msg)
|
||||||
|
}
|
||||||
|
|
||||||
|
id := req.ID
|
||||||
|
loginPassword := req.LoginPassword
|
||||||
|
passwordGenerated := false
|
||||||
|
if loginPassword == "" {
|
||||||
|
pw, genErr := generateLoginPassword()
|
||||||
|
if genErr != nil {
|
||||||
|
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "generate password failed")
|
||||||
|
}
|
||||||
|
loginPassword = pw
|
||||||
|
passwordGenerated = true
|
||||||
|
}
|
||||||
|
|
||||||
|
loginHash, err := h.cfg.Hash.Hash(ctx, auth.PasswordLogin, loginPassword)
|
||||||
|
if err != nil {
|
||||||
|
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "hash failed")
|
||||||
|
}
|
||||||
|
|
||||||
|
var talkHash sql.NullString
|
||||||
|
if req.TalkPassword != "" {
|
||||||
|
th, hashErr := h.cfg.Hash.Hash(ctx, auth.PasswordTalk, req.TalkPassword)
|
||||||
|
if hashErr != nil {
|
||||||
|
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "hash failed")
|
||||||
|
}
|
||||||
|
talkHash = sql.NullString{String: th, Valid: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
nowMs := h.cfg.Now().UnixMilli()
|
||||||
|
const maxIDAttempts = 8
|
||||||
|
for attempt := 0; attempt < maxIDAttempts; attempt++ {
|
||||||
|
useID := id
|
||||||
|
if useID == "" {
|
||||||
|
genID, genErr := generateEndpointID()
|
||||||
|
if genErr != nil {
|
||||||
|
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "generate id failed")
|
||||||
|
}
|
||||||
|
useID = genID
|
||||||
|
}
|
||||||
|
|
||||||
|
insertErr := h.cfg.DB.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||||
|
_, execErr := tx.ExecContext(ctx, `
|
||||||
|
INSERT INTO endpoints(
|
||||||
|
id, name, remark, source, login_hash, talk_hash, talk_version,
|
||||||
|
default_delay_ms, enabled, created_at
|
||||||
|
) VALUES (?, ?, '', ?, ?, ?, 0, 0, 1, ?)`,
|
||||||
|
useID, req.Name, sourceSelf, loginHash, talkHash, nowMs,
|
||||||
|
)
|
||||||
|
return execErr
|
||||||
|
})
|
||||||
|
if insertErr == nil {
|
||||||
|
out := RegisterResult{ID: useID}
|
||||||
|
if passwordGenerated {
|
||||||
|
out.LoginPassword = loginPassword
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
if isUniqueConstraint(insertErr) {
|
||||||
|
if id != "" {
|
||||||
|
return RegisterResult{}, apiErr(http.StatusConflict, protocol.CodeIDTaken, "id taken")
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "insert failed")
|
||||||
|
}
|
||||||
|
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "generate id exhausted")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *RegisterHandler) loadRegistrationSettings(ctx context.Context) (enabled bool, code string, err error) {
|
||||||
|
var enabledVal, codeVal sql.NullString
|
||||||
|
row := h.cfg.DB.Read.QueryRowContext(ctx, `SELECT value FROM settings WHERE key = ?`, settingRegistrationEnabled)
|
||||||
|
if scanErr := row.Scan(&enabledVal); scanErr != nil && !errors.Is(scanErr, sql.ErrNoRows) {
|
||||||
|
return false, "", scanErr
|
||||||
|
}
|
||||||
|
row = h.cfg.DB.Read.QueryRowContext(ctx, `SELECT value FROM settings WHERE key = ?`, settingRegistrationCode)
|
||||||
|
if scanErr := row.Scan(&codeVal); scanErr != nil && !errors.Is(scanErr, sql.ErrNoRows) {
|
||||||
|
return false, "", scanErr
|
||||||
|
}
|
||||||
|
return settingTruthy(enabledVal.String), codeVal.String, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *RegisterHandler) logResult(result, id, ip string) {
|
||||||
|
h.cfg.Logger.Info("register", "result", result, "id", id, "ip", ip)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Server 实现 identity.Service:I1 只实现 Register,其余仍为未实现。
|
||||||
|
type Server struct {
|
||||||
|
handler *RegisterHandler
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewServer 用同一套依赖构造 Service(Register)与可挂载 Handler。
|
||||||
|
func NewServer(cfg RegisterConfig) *Server {
|
||||||
|
return &Server{handler: NewRegisterHandler(cfg)}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Handler 返回可挂载的注册 HTTP 处理器。
|
||||||
|
func (s *Server) Handler() http.Handler { return s.handler }
|
||||||
|
|
||||||
|
// Register 实现自助注册(source 固定为 self;RemoteIP 用于锁定)。
|
||||||
|
func (s *Server) Register(ctx context.Context, req RegisterRequest) (RegisterResult, error) {
|
||||||
|
preq := &protocol.RegisterRequest{
|
||||||
|
RegistrationCode: req.RegistrationCode,
|
||||||
|
ID: req.ID,
|
||||||
|
LoginPassword: req.LoginPassword,
|
||||||
|
Name: req.Name,
|
||||||
|
TalkPassword: req.TalkPassword,
|
||||||
|
}
|
||||||
|
ip := req.RemoteIP
|
||||||
|
if ip == "" {
|
||||||
|
ip = "0.0.0.0"
|
||||||
|
}
|
||||||
|
result, err := s.handler.register(ctx, preq, ip)
|
||||||
|
if err != nil {
|
||||||
|
return RegisterResult{}, err
|
||||||
|
}
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) SelfGet(context.Context, string) (SelfInfo, error) {
|
||||||
|
return SelfInfo{}, ErrNotImplemented
|
||||||
|
}
|
||||||
|
func (s *Server) SelfUpdate(context.Context, string, *protocol.SelfUpdate) error {
|
||||||
|
return ErrNotImplemented
|
||||||
|
}
|
||||||
|
func (s *Server) SelfSetTalkPassword(context.Context, string, string) error {
|
||||||
|
return ErrNotImplemented
|
||||||
|
}
|
||||||
|
func (s *Server) SelfChangeLoginPassword(context.Context, string, string, string) (string, error) {
|
||||||
|
return "", ErrNotImplemented
|
||||||
|
}
|
||||||
|
func (s *Server) SelfLogout(context.Context, string) error { return ErrNotImplemented }
|
||||||
|
func (s *Server) UnlockTalk(context.Context, string, string, string) error {
|
||||||
|
return ErrNotImplemented
|
||||||
|
}
|
||||||
|
func (s *Server) HasTalkGrant(context.Context, string, string) (bool, error) {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
func (s *Server) Disable(context.Context, string) error { return ErrNotImplemented }
|
||||||
|
func (s *Server) Enable(context.Context, string) error { return ErrNotImplemented }
|
||||||
|
func (s *Server) Delete(context.Context, string) error { return ErrNotImplemented }
|
||||||
|
|
||||||
|
var _ Service = (*Server)(nil)
|
||||||
|
var _ http.Handler = (*RegisterHandler)(nil)
|
||||||
|
|
||||||
|
func setCORS(w http.ResponseWriter) {
|
||||||
|
w.Header().Set("Access-Control-Allow-Origin", "*")
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeRegisterOK(w http.ResponseWriter, result RegisterResult) {
|
||||||
|
setCORS(w)
|
||||||
|
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
_ = protocol.Encode(w, protocol.RegisterResponse{
|
||||||
|
OK: true,
|
||||||
|
Data: protocol.RegisterData{
|
||||||
|
ID: result.ID,
|
||||||
|
LoginPassword: result.LoginPassword,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeRegisterError(w http.ResponseWriter, err *APIError) {
|
||||||
|
setCORS(w)
|
||||||
|
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||||
|
status := http.StatusInternalServerError
|
||||||
|
code, msg := protocol.CodeBusy, "internal error"
|
||||||
|
if err != nil {
|
||||||
|
status = err.Status
|
||||||
|
code, msg = err.Code, err.Message
|
||||||
|
}
|
||||||
|
w.WriteHeader(status)
|
||||||
|
_ = protocol.Encode(w, protocol.RegisterResponse{
|
||||||
|
OK: false,
|
||||||
|
Error: &protocol.ErrorBody{Code: code, Message: msg},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func clientIPFromRemoteAddr(r *http.Request) string {
|
||||||
|
host, _, err := net.SplitHostPort(r.RemoteAddr)
|
||||||
|
if err != nil {
|
||||||
|
return r.RemoteAddr
|
||||||
|
}
|
||||||
|
return host
|
||||||
|
}
|
||||||
|
|
||||||
|
func settingTruthy(v string) bool {
|
||||||
|
switch strings.TrimSpace(strings.ToLower(v)) {
|
||||||
|
case "1", "true", "yes", "on":
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func constantTimeEqual(a, b string) bool {
|
||||||
|
// 长度不同时 ConstantTimeCompare 直接失败;先按较短侧对齐比较,再核对长度,避免过早返回。
|
||||||
|
ab := []byte(a)
|
||||||
|
bb := []byte(b)
|
||||||
|
if len(ab) != len(bb) {
|
||||||
|
dummy := make([]byte, len(ab))
|
||||||
|
subtle.ConstantTimeCompare(ab, dummy)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return subtle.ConstantTimeCompare(ab, bb) == 1
|
||||||
|
}
|
||||||
|
|
||||||
|
func generateEndpointID() (string, error) {
|
||||||
|
b := make([]byte, 8)
|
||||||
|
if _, err := rand.Read(b); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
out := make([]byte, 8)
|
||||||
|
for i := range b {
|
||||||
|
out[i] = idAlphabet[int(b[i])%len(idAlphabet)]
|
||||||
|
}
|
||||||
|
return "e_" + string(out), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func generateLoginPassword() (string, error) {
|
||||||
|
// 20 字节可读字符,满足 8–128,且不以 nst_ 开头(字母数字混合,冲突概率极低;若撞前缀则重抽)。
|
||||||
|
const alphabet = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789"
|
||||||
|
for range 8 {
|
||||||
|
b := make([]byte, 20)
|
||||||
|
if _, err := rand.Read(b); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
out := make([]byte, 20)
|
||||||
|
for i := range b {
|
||||||
|
out[i] = alphabet[int(b[i])%len(alphabet)]
|
||||||
|
}
|
||||||
|
pw := string(out)
|
||||||
|
if !strings.HasPrefix(pw, protocol.SessionTokenPrefix) && utf8.RuneCountInString(pw) >= protocol.MinLoginPasswordLen {
|
||||||
|
return pw, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return "", errors.New("identity: generate login password failed")
|
||||||
|
}
|
||||||
|
|
||||||
|
func isUniqueConstraint(err error) bool {
|
||||||
|
if err == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
msg := strings.ToLower(err.Error())
|
||||||
|
return strings.Contains(msg, "unique constraint") || strings.Contains(msg, "constraint failed")
|
||||||
|
}
|
||||||
|
|
||||||
|
func isBodyTooLarge(err error) bool {
|
||||||
|
if err == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
msg := strings.ToLower(err.Error())
|
||||||
|
return strings.Contains(msg, "request body too large") || strings.Contains(msg, "http: request body too large")
|
||||||
|
}
|
||||||
@@ -0,0 +1,446 @@
|
|||||||
|
package identity
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"encoding/json"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||||
|
)
|
||||||
|
|
||||||
|
// registerIPLocker 仅实现 LockRegisterIP:5 分钟窗口内 10 次失败则锁定 5 分钟。
|
||||||
|
// P3 的完整 LoginLocks 未合入本分支时,测试用此可替换实现覆盖 F23 锁定验收。
|
||||||
|
type registerIPLocker struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
fails map[string][]time.Time
|
||||||
|
lockedUntil map[string]time.Time
|
||||||
|
now func() time.Time
|
||||||
|
window time.Duration
|
||||||
|
limit int
|
||||||
|
lockFor time.Duration
|
||||||
|
}
|
||||||
|
|
||||||
|
func newRegisterIPLocker(now func() time.Time) *registerIPLocker {
|
||||||
|
if now == nil {
|
||||||
|
now = time.Now
|
||||||
|
}
|
||||||
|
return ®isterIPLocker{
|
||||||
|
fails: make(map[string][]time.Time),
|
||||||
|
lockedUntil: make(map[string]time.Time),
|
||||||
|
now: now,
|
||||||
|
window: 5 * time.Minute,
|
||||||
|
limit: 10,
|
||||||
|
lockFor: 5 * time.Minute,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (l *registerIPLocker) Check(key auth.LockKey) (bool, time.Duration) {
|
||||||
|
if key.Kind != auth.LockRegisterIP {
|
||||||
|
return false, 0
|
||||||
|
}
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
until, ok := l.lockedUntil[key.IP]
|
||||||
|
if !ok {
|
||||||
|
return false, 0
|
||||||
|
}
|
||||||
|
now := l.now()
|
||||||
|
if now.Before(until) {
|
||||||
|
return true, until.Sub(now)
|
||||||
|
}
|
||||||
|
delete(l.lockedUntil, key.IP)
|
||||||
|
return false, 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func (l *registerIPLocker) Fail(key auth.LockKey) (bool, time.Duration) {
|
||||||
|
if key.Kind != auth.LockRegisterIP {
|
||||||
|
return false, 0
|
||||||
|
}
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
now := l.now()
|
||||||
|
if until, ok := l.lockedUntil[key.IP]; ok && now.Before(until) {
|
||||||
|
return true, until.Sub(now)
|
||||||
|
}
|
||||||
|
cutoff := now.Add(-l.window)
|
||||||
|
list := l.fails[key.IP]
|
||||||
|
kept := list[:0]
|
||||||
|
for _, t := range list {
|
||||||
|
if t.After(cutoff) {
|
||||||
|
kept = append(kept, t)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
kept = append(kept, now)
|
||||||
|
l.fails[key.IP] = kept
|
||||||
|
if len(kept) >= l.limit {
|
||||||
|
until := now.Add(l.lockFor)
|
||||||
|
l.lockedUntil[key.IP] = until
|
||||||
|
return true, l.lockFor
|
||||||
|
}
|
||||||
|
return false, 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func (l *registerIPLocker) ClearEndpoint(string) {}
|
||||||
|
func (l *registerIPLocker) Clear(key auth.LockKey) {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
delete(l.fails, key.IP)
|
||||||
|
delete(l.lockedUntil, key.IP)
|
||||||
|
}
|
||||||
|
|
||||||
|
var _ auth.LoginLocks = (*registerIPLocker)(nil)
|
||||||
|
|
||||||
|
type testEnv struct {
|
||||||
|
db *store.DB
|
||||||
|
hash auth.HashPool
|
||||||
|
locks *registerIPLocker
|
||||||
|
logBuf *bytes.Buffer
|
||||||
|
handler http.Handler
|
||||||
|
fixedIP string
|
||||||
|
}
|
||||||
|
|
||||||
|
func openTestEnv(t *testing.T) *testEnv {
|
||||||
|
t.Helper()
|
||||||
|
db, err := store.Open(t.TempDir(), "FULL")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { _ = db.Close() })
|
||||||
|
|
||||||
|
buf := &bytes.Buffer{}
|
||||||
|
logger := slog.New(slog.NewTextHandler(buf, &slog.HandlerOptions{Level: slog.LevelInfo}))
|
||||||
|
locks := newRegisterIPLocker(time.Now)
|
||||||
|
env := &testEnv{
|
||||||
|
db: db,
|
||||||
|
hash: auth.NewStubHashPool(),
|
||||||
|
locks: locks,
|
||||||
|
logBuf: buf,
|
||||||
|
fixedIP: "203.0.113.10",
|
||||||
|
}
|
||||||
|
env.handler = NewRegisterHandler(RegisterConfig{
|
||||||
|
DB: db,
|
||||||
|
Hash: env.hash,
|
||||||
|
Locks: locks,
|
||||||
|
Logger: logger,
|
||||||
|
ClientIP: func(*http.Request) string {
|
||||||
|
return env.fixedIP
|
||||||
|
},
|
||||||
|
})
|
||||||
|
return env
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e *testEnv) setRegistration(t *testing.T, enabled bool, code string) {
|
||||||
|
t.Helper()
|
||||||
|
en := "0"
|
||||||
|
if enabled {
|
||||||
|
en = "1"
|
||||||
|
}
|
||||||
|
now := time.Now().UnixMilli()
|
||||||
|
err := e.db.Queue.Do(context.Background(), func(tx *sql.Tx) error {
|
||||||
|
if _, err := tx.Exec(`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)
|
||||||
|
ON CONFLICT(key) DO UPDATE SET value=excluded.value, updated_at=excluded.updated_at`,
|
||||||
|
settingRegistrationEnabled, en, now); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
_, err := tx.Exec(`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)
|
||||||
|
ON CONFLICT(key) DO UPDATE SET value=excluded.value, updated_at=excluded.updated_at`,
|
||||||
|
settingRegistrationCode, code, now)
|
||||||
|
return err
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e *testEnv) insertEndpoint(t *testing.T, id, loginHash string) {
|
||||||
|
t.Helper()
|
||||||
|
err := e.db.Queue.Do(context.Background(), func(tx *sql.Tx) error {
|
||||||
|
_, err := tx.Exec(`INSERT INTO endpoints(
|
||||||
|
id, name, remark, source, login_hash, talk_hash, talk_version,
|
||||||
|
default_delay_ms, enabled, created_at
|
||||||
|
) VALUES (?, '', '', 'admin', ?, NULL, 0, 0, 1, ?)`, id, loginHash, time.Now().UnixMilli())
|
||||||
|
return err
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e *testEnv) getEndpoint(t *testing.T, id string) (source, loginHash string, ok bool) {
|
||||||
|
t.Helper()
|
||||||
|
err := e.db.Read.QueryRow(`SELECT source, login_hash FROM endpoints WHERE id = ?`, id).Scan(&source, &loginHash)
|
||||||
|
if errorsIsNoRows(err) {
|
||||||
|
return "", "", false
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return source, loginHash, true
|
||||||
|
}
|
||||||
|
|
||||||
|
func errorsIsNoRows(err error) bool {
|
||||||
|
return err == sql.ErrNoRows
|
||||||
|
}
|
||||||
|
|
||||||
|
type registerResp struct {
|
||||||
|
OK bool `json:"ok"`
|
||||||
|
Data struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
LoginPassword string `json:"login_password"`
|
||||||
|
} `json:"data"`
|
||||||
|
Error *protocol.ErrorBody `json:"error"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e *testEnv) doRegister(t *testing.T, body string) (int, registerResp, http.Header) {
|
||||||
|
t.Helper()
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/client/register", strings.NewReader(body))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.RemoteAddr = e.fixedIP + ":54321"
|
||||||
|
rr := httptest.NewRecorder()
|
||||||
|
e.handler.ServeHTTP(rr, req)
|
||||||
|
var resp registerResp
|
||||||
|
if err := json.Unmarshal(rr.Body.Bytes(), &resp); err != nil {
|
||||||
|
t.Fatalf("decode resp: %v body=%s", err, rr.Body.String())
|
||||||
|
}
|
||||||
|
return rr.Code, resp, rr.Header()
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegisterF23_ClosedFails(t *testing.T) {
|
||||||
|
env := openTestEnv(t)
|
||||||
|
env.setRegistration(t, false, "secretcode")
|
||||||
|
|
||||||
|
code, resp, hdr := env.doRegister(t, `{"registration_code":"secretcode","id":"ep_closed","login_password":"password1"}`)
|
||||||
|
if code != http.StatusForbidden || resp.Error == nil || resp.Error.Code != protocol.CodeRegistrationClosed {
|
||||||
|
t.Fatalf("status=%d resp=%+v", code, resp)
|
||||||
|
}
|
||||||
|
if hdr.Get("Access-Control-Allow-Origin") != "*" {
|
||||||
|
t.Fatalf("missing CORS: %v", hdr)
|
||||||
|
}
|
||||||
|
if _, _, ok := env.getEndpoint(t, "ep_closed"); ok {
|
||||||
|
t.Fatal("endpoint should not be created when closed")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegisterF23_WrongCodeFails_RightCodeOK(t *testing.T) {
|
||||||
|
env := openTestEnv(t)
|
||||||
|
env.setRegistration(t, true, "good-code-01")
|
||||||
|
|
||||||
|
code, resp, _ := env.doRegister(t, `{"registration_code":"bad-code-xx","id":"ep_wrong","login_password":"password1"}`)
|
||||||
|
if code != http.StatusForbidden || resp.Error == nil || resp.Error.Code != protocol.CodeRegistrationCodeInvalid {
|
||||||
|
t.Fatalf("wrong code: status=%d resp=%+v", code, resp)
|
||||||
|
}
|
||||||
|
|
||||||
|
code, resp, hdr := env.doRegister(t, `{"registration_code":"good-code-01","id":"ep_ok1","login_password":"password1","name":"门口"}`)
|
||||||
|
if code != http.StatusOK || !resp.OK || resp.Data.ID != "ep_ok1" {
|
||||||
|
t.Fatalf("ok register: status=%d resp=%+v", code, resp)
|
||||||
|
}
|
||||||
|
if resp.Data.LoginPassword != "" {
|
||||||
|
t.Fatalf("provided password must not echo: %q", resp.Data.LoginPassword)
|
||||||
|
}
|
||||||
|
if hdr.Get("Access-Control-Allow-Origin") != "*" {
|
||||||
|
t.Fatal("missing CORS on success")
|
||||||
|
}
|
||||||
|
source, loginHash, ok := env.getEndpoint(t, "ep_ok1")
|
||||||
|
if !ok || source != "self" {
|
||||||
|
t.Fatalf("endpoint source=%q ok=%v", source, ok)
|
||||||
|
}
|
||||||
|
match, err := env.hash.Verify(context.Background(), auth.PasswordLogin, "password1", loginHash)
|
||||||
|
if err != nil || !match {
|
||||||
|
t.Fatalf("login hash verify: match=%v err=%v", match, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegisterF23_ChangeCode_OldFails_ExistingRemains(t *testing.T) {
|
||||||
|
env := openTestEnv(t)
|
||||||
|
env.setRegistration(t, true, "code-old-01")
|
||||||
|
|
||||||
|
code, resp, _ := env.doRegister(t, `{"registration_code":"code-old-01","id":"ep_keep","login_password":"password1"}`)
|
||||||
|
if code != http.StatusOK || resp.Data.ID != "ep_keep" {
|
||||||
|
t.Fatalf("first register: status=%d resp=%+v", code, resp)
|
||||||
|
}
|
||||||
|
_, oldHash, ok := env.getEndpoint(t, "ep_keep")
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("missing endpoint after register")
|
||||||
|
}
|
||||||
|
|
||||||
|
env.setRegistration(t, true, "code-new-02")
|
||||||
|
code, resp, _ = env.doRegister(t, `{"registration_code":"code-old-01","id":"ep_new","login_password":"password1"}`)
|
||||||
|
if code != http.StatusForbidden || resp.Error == nil || resp.Error.Code != protocol.CodeRegistrationCodeInvalid {
|
||||||
|
t.Fatalf("old code after rotate: status=%d resp=%+v", code, resp)
|
||||||
|
}
|
||||||
|
code, resp, _ = env.doRegister(t, `{"registration_code":"code-new-02","id":"ep_new","login_password":"password1"}`)
|
||||||
|
if code != http.StatusOK || resp.Data.ID != "ep_new" {
|
||||||
|
t.Fatalf("new code: status=%d resp=%+v", code, resp)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, hashAfter, ok := env.getEndpoint(t, "ep_keep")
|
||||||
|
if !ok || hashAfter != oldHash {
|
||||||
|
t.Fatalf("existing endpoint mutated: ok=%v hashEqual=%v", ok, hashAfter == oldHash)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegisterF23_WrongCodeLock(t *testing.T) {
|
||||||
|
env := openTestEnv(t)
|
||||||
|
env.setRegistration(t, true, "lock-code-1")
|
||||||
|
|
||||||
|
for i := 0; i < 10; i++ {
|
||||||
|
code, resp, _ := env.doRegister(t, `{"registration_code":"wrong-code","id":"ep_lock","login_password":"password1"}`)
|
||||||
|
if code != http.StatusForbidden || resp.Error == nil || resp.Error.Code != protocol.CodeRegistrationCodeInvalid {
|
||||||
|
t.Fatalf("fail #%d: status=%d resp=%+v", i+1, code, resp)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
code, resp, _ := env.doRegister(t, `{"registration_code":"lock-code-1","id":"ep_lock","login_password":"password1"}`)
|
||||||
|
if code != http.StatusTooManyRequests || resp.Error == nil || resp.Error.Code != protocol.CodeRateLimited {
|
||||||
|
t.Fatalf("locked with good code: status=%d resp=%+v", code, resp)
|
||||||
|
}
|
||||||
|
if _, _, ok := env.getEndpoint(t, "ep_lock"); ok {
|
||||||
|
t.Fatal("must not insert while rate limited")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegisterF23_IDTakenKeepsOriginal(t *testing.T) {
|
||||||
|
env := openTestEnv(t)
|
||||||
|
env.setRegistration(t, true, "taken-code")
|
||||||
|
env.insertEndpoint(t, "ep_taken", "stub$original-password-xx")
|
||||||
|
|
||||||
|
code, resp, _ := env.doRegister(t, `{"registration_code":"taken-code","id":"ep_taken","login_password":"password1"}`)
|
||||||
|
if code != http.StatusConflict || resp.Error == nil || resp.Error.Code != protocol.CodeIDTaken {
|
||||||
|
t.Fatalf("id taken: status=%d resp=%+v", code, resp)
|
||||||
|
}
|
||||||
|
source, loginHash, ok := env.getEndpoint(t, "ep_taken")
|
||||||
|
if !ok || source != "admin" || loginHash != "stub$original-password-xx" {
|
||||||
|
t.Fatalf("original endpoint changed: source=%q hash=%q", source, loginHash)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegister_GenerateIDAndPassword(t *testing.T) {
|
||||||
|
env := openTestEnv(t)
|
||||||
|
env.setRegistration(t, true, "gen-code-01")
|
||||||
|
|
||||||
|
code, resp, _ := env.doRegister(t, `{"registration_code":"gen-code-01","id":"","login_password":""}`)
|
||||||
|
if code != http.StatusOK || !resp.OK {
|
||||||
|
t.Fatalf("status=%d resp=%+v", code, resp)
|
||||||
|
}
|
||||||
|
if !strings.HasPrefix(resp.Data.ID, "e_") || len(resp.Data.ID) != 10 {
|
||||||
|
t.Fatalf("generated id=%q", resp.Data.ID)
|
||||||
|
}
|
||||||
|
if len(resp.Data.LoginPassword) < protocol.MinLoginPasswordLen {
|
||||||
|
t.Fatalf("generated password too short: %q", resp.Data.LoginPassword)
|
||||||
|
}
|
||||||
|
if strings.HasPrefix(resp.Data.LoginPassword, protocol.SessionTokenPrefix) {
|
||||||
|
t.Fatal("generated password starts with nst_")
|
||||||
|
}
|
||||||
|
source, _, ok := env.getEndpoint(t, resp.Data.ID)
|
||||||
|
if !ok || source != "self" {
|
||||||
|
t.Fatalf("source=%q ok=%v", source, ok)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegister_OPTIONS_CORS(t *testing.T) {
|
||||||
|
env := openTestEnv(t)
|
||||||
|
req := httptest.NewRequest(http.MethodOptions, "/api/client/register", nil)
|
||||||
|
rr := httptest.NewRecorder()
|
||||||
|
env.handler.ServeHTTP(rr, req)
|
||||||
|
if rr.Code != http.StatusNoContent {
|
||||||
|
t.Fatalf("status=%d", rr.Code)
|
||||||
|
}
|
||||||
|
if rr.Header().Get("Access-Control-Allow-Origin") != "*" {
|
||||||
|
t.Fatal("missing Allow-Origin")
|
||||||
|
}
|
||||||
|
if !strings.Contains(rr.Header().Get("Access-Control-Allow-Methods"), "POST") {
|
||||||
|
t.Fatalf("methods=%q", rr.Header().Get("Access-Control-Allow-Methods"))
|
||||||
|
}
|
||||||
|
if len(rr.Result().Cookies()) != 0 {
|
||||||
|
t.Fatal("must not set cookies")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegister_BodyTooLarge(t *testing.T) {
|
||||||
|
env := openTestEnv(t)
|
||||||
|
env.setRegistration(t, true, "big-code-01")
|
||||||
|
body := `{"registration_code":"big-code-01","id":"ep_big","login_password":"password1","name":"` + strings.Repeat("x", 5000) + `"}`
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/client/register", strings.NewReader(body))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
rr := httptest.NewRecorder()
|
||||||
|
env.handler.ServeHTTP(rr, req)
|
||||||
|
if rr.Code != http.StatusBadRequest {
|
||||||
|
t.Fatalf("status=%d body=%s", rr.Code, rr.Body.String())
|
||||||
|
}
|
||||||
|
var resp registerResp
|
||||||
|
_ = json.Unmarshal(rr.Body.Bytes(), &resp)
|
||||||
|
if resp.Error == nil || resp.Error.Code != protocol.CodeBadRequest {
|
||||||
|
t.Fatalf("resp=%+v", resp)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegister_LogOmitsSecrets(t *testing.T) {
|
||||||
|
env := openTestEnv(t)
|
||||||
|
env.setRegistration(t, true, "log-secret-code")
|
||||||
|
_, _, _ = env.doRegister(t, `{"registration_code":"log-secret-code","id":"ep_log","login_password":"supersecretpw"}`)
|
||||||
|
logged := env.logBuf.String()
|
||||||
|
if strings.Contains(logged, "log-secret-code") || strings.Contains(logged, "supersecretpw") {
|
||||||
|
t.Fatalf("log leaked secrets: %s", logged)
|
||||||
|
}
|
||||||
|
if !strings.Contains(logged, "ep_log") || !strings.Contains(logged, env.fixedIP) {
|
||||||
|
t.Fatalf("log missing id/ip: %s", logged)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegister_NSTPasswordRejected(t *testing.T) {
|
||||||
|
env := openTestEnv(t)
|
||||||
|
env.setRegistration(t, true, "nst-code-01")
|
||||||
|
code, resp, _ := env.doRegister(t, `{"registration_code":"nst-code-01","id":"ep_nst","login_password":"nst_notallowed"}`)
|
||||||
|
if code != http.StatusBadRequest || resp.Error == nil || resp.Error.Code != protocol.CodeBadRequest {
|
||||||
|
t.Fatalf("status=%d resp=%+v", code, resp)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMountOnServeMux(t *testing.T) {
|
||||||
|
env := openTestEnv(t)
|
||||||
|
env.setRegistration(t, true, "mux-code-01")
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
mux.Handle("/api/client/register", env.handler)
|
||||||
|
|
||||||
|
body := `{"registration_code":"mux-code-01","id":"ep_mux","login_password":"password1"}`
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/client/register", strings.NewReader(body))
|
||||||
|
rr := httptest.NewRecorder()
|
||||||
|
mux.ServeHTTP(rr, req)
|
||||||
|
if rr.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status=%d body=%s", rr.Code, rr.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServerRegisterInterface(t *testing.T) {
|
||||||
|
env := openTestEnv(t)
|
||||||
|
env.setRegistration(t, true, "svc-code-01")
|
||||||
|
svc := NewServer(RegisterConfig{
|
||||||
|
DB: env.db,
|
||||||
|
Hash: env.hash,
|
||||||
|
Locks: env.locks,
|
||||||
|
Logger: slog.New(slog.NewTextHandler(io.Discard, nil)),
|
||||||
|
})
|
||||||
|
res, err := svc.Register(context.Background(), RegisterRequest{
|
||||||
|
RegistrationCode: "svc-code-01",
|
||||||
|
ID: "ep_svc",
|
||||||
|
LoginPassword: "password1",
|
||||||
|
RemoteIP: "198.51.100.1",
|
||||||
|
Source: "self",
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if res.ID != "ep_svc" {
|
||||||
|
t.Fatalf("id=%q", res.ID)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,153 @@
|
|||||||
|
package message
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/config"
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 消息状态(DEVELOPMENT 7.1)。
|
||||||
|
const (
|
||||||
|
StateScheduled = "scheduled"
|
||||||
|
StateDispatched = "dispatched"
|
||||||
|
StateCompleted = "completed"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 投递状态。
|
||||||
|
const (
|
||||||
|
DeliveryPending = "pending"
|
||||||
|
)
|
||||||
|
|
||||||
|
// talk_grants.kind。
|
||||||
|
const (
|
||||||
|
GrantKindPassword = "password"
|
||||||
|
GrantKindReply = "reply"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 请求频率桶默认突发容量(DEVELOPMENT 6.10;配置无单独字段)。
|
||||||
|
const defaultRequestBurst = 100
|
||||||
|
|
||||||
|
// Limits 是提交所需的配置上限(来自 config.LimitsConfig)。
|
||||||
|
type Limits struct {
|
||||||
|
MaxBodyBytes int
|
||||||
|
MaxMetaBytes int
|
||||||
|
MaxFrameBytes int
|
||||||
|
MaxTTLSeconds int64
|
||||||
|
MaxScheduleSeconds int64
|
||||||
|
RequestsPerSecond float64
|
||||||
|
RequestBurst int
|
||||||
|
MaxPendingPerSender int
|
||||||
|
MaxPendingPerReceiver int
|
||||||
|
GraceSeconds int64
|
||||||
|
}
|
||||||
|
|
||||||
|
// LimitsFromConfig 从平台配置构造 Limits。
|
||||||
|
func LimitsFromConfig(c config.LimitsConfig) Limits {
|
||||||
|
burst := defaultRequestBurst
|
||||||
|
return Limits{
|
||||||
|
MaxBodyBytes: c.MaxBodyBytes,
|
||||||
|
MaxMetaBytes: c.MaxMetaBytes,
|
||||||
|
MaxFrameBytes: c.MaxFrameBytes,
|
||||||
|
MaxTTLSeconds: int64(c.MaxTTLSeconds),
|
||||||
|
MaxScheduleSeconds: int64(c.MaxScheduleSeconds),
|
||||||
|
RequestsPerSecond: float64(c.RequestsPerSecond),
|
||||||
|
RequestBurst: burst,
|
||||||
|
MaxPendingPerSender: c.MaxPendingPerSender,
|
||||||
|
MaxPendingPerReceiver: c.MaxPendingPerReceiver,
|
||||||
|
GraceSeconds: int64(c.GraceSeconds),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// App 实现 Service 的提交路径(M1);其余方法暂返回未实现或空操作。
|
||||||
|
type App struct {
|
||||||
|
db *store.DB
|
||||||
|
lim Limits
|
||||||
|
hash auth.HashPool
|
||||||
|
locks auth.LoginLocks
|
||||||
|
nowFn func() time.Time
|
||||||
|
rates *rateLimiter
|
||||||
|
}
|
||||||
|
|
||||||
|
// Option 配置 App。
|
||||||
|
type Option func(*App)
|
||||||
|
|
||||||
|
// WithNow 注入时钟(测试用)。
|
||||||
|
func WithNow(now func() time.Time) Option {
|
||||||
|
return func(a *App) { a.nowFn = now }
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithLocks 注入对话密码锁定计数器;nil 表示不锁定。
|
||||||
|
func WithLocks(locks auth.LoginLocks) Option {
|
||||||
|
return func(a *App) { a.locks = locks }
|
||||||
|
}
|
||||||
|
|
||||||
|
// New 创建消息服务实现。hash 用于校验对话密码;locks 可为 nil。
|
||||||
|
func New(db *store.DB, lim Limits, hash auth.HashPool, opts ...Option) *App {
|
||||||
|
if lim.RequestBurst <= 0 {
|
||||||
|
lim.RequestBurst = defaultRequestBurst
|
||||||
|
}
|
||||||
|
a := &App{
|
||||||
|
db: db,
|
||||||
|
lim: lim,
|
||||||
|
hash: hash,
|
||||||
|
nowFn: time.Now,
|
||||||
|
rates: newRateLimiter(lim.RequestsPerSecond, lim.RequestBurst),
|
||||||
|
}
|
||||||
|
for _, opt := range opts {
|
||||||
|
opt(a)
|
||||||
|
}
|
||||||
|
return a
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) now() time.Time {
|
||||||
|
return a.nowFn()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) protocolLimits() protocol.Limits {
|
||||||
|
return protocol.Limits{
|
||||||
|
MaxBodyBytes: a.lim.MaxBodyBytes,
|
||||||
|
MaxMetaBytes: a.lim.MaxMetaBytes,
|
||||||
|
MaxFrameBytes: a.lim.MaxFrameBytes,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) Ack(context.Context, string, *protocol.Ack) (AckResult, error) {
|
||||||
|
return AckResult{}, ErrNotImplemented
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) Recall(context.Context, string, *protocol.Recall) (protocol.RecallData, error) {
|
||||||
|
return protocol.RecallData{}, ErrNotImplemented
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) Status(context.Context, string, *protocol.Status) (any, error) {
|
||||||
|
return nil, ErrNotImplemented
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) ReceiptAck(context.Context, string, *protocol.ReceiptAck) error {
|
||||||
|
return ErrNotImplemented
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) DispatchDue(context.Context, int64, int) (int, error) {
|
||||||
|
return 0, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) PushPending(context.Context, string, port.ConnID) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) OnPublishDropped(context.Context, string, port.ConnID, []byte) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) CleanupOnce(context.Context, int64) error { return nil }
|
||||||
|
|
||||||
|
func (a *App) RecoverOnStart(context.Context) error { return nil }
|
||||||
|
|
||||||
|
func (a *App) WakePush(string) {}
|
||||||
|
|
||||||
|
var _ Service = (*App)(nil)
|
||||||
@@ -0,0 +1,52 @@
|
|||||||
|
package message
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||||
|
)
|
||||||
|
|
||||||
|
// dispatchMinimalTx 是 M1 最小分发:单聊插一条 pending;群按当前成员去掉发送者各插 pending;
|
||||||
|
// 消息改为 dispatched。完整 7.4 规则见 DEVIATIONS「消息 M」。
|
||||||
|
func dispatchMinimalTx(tx *sql.Tx, seq int64, senderID, destKind, destID string, sendAt int64, keep int, nowMs int64) (string, error) {
|
||||||
|
recipients := make([]string, 0, 8)
|
||||||
|
switch destKind {
|
||||||
|
case protocol.TargetEndpoint:
|
||||||
|
recipients = append(recipients, destID)
|
||||||
|
case protocol.TargetGroup:
|
||||||
|
rows, err := tx.Query(`
|
||||||
|
SELECT endpoint_id FROM group_members WHERE group_id = ? AND endpoint_id != ?`, destID, senderID)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
defer func() { _ = rows.Close() }()
|
||||||
|
for rows.Next() {
|
||||||
|
var id string
|
||||||
|
if err := rows.Scan(&id); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
recipients = append(recipients, id)
|
||||||
|
}
|
||||||
|
if err := rows.Err(); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return "", errCode(protocol.CodeBadRequest, "invalid dest_kind")
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, ep := range recipients {
|
||||||
|
if _, err := tx.Exec(`
|
||||||
|
INSERT INTO deliveries(seq, endpoint_id, send_at, keep, state, reason, expire_at, pushed_conn, pushed_at, attempts, updated_at)
|
||||||
|
VALUES(?,?,?,?,?,?,NULL,NULL,NULL,0,?)`,
|
||||||
|
seq, ep, sendAt, keep, DeliveryPending, "", nowMs,
|
||||||
|
); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
state := StateDispatched
|
||||||
|
if _, err := tx.Exec(`UPDATE messages SET state = ? WHERE seq = ?`, state, seq); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
return state, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,7 @@
|
|||||||
|
package message
|
||||||
|
|
||||||
|
import "git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||||
|
|
||||||
|
func errCode(code, msg string) *protocol.Error {
|
||||||
|
return &protocol.Error{Code: code, Message: msg}
|
||||||
|
}
|
||||||
@@ -0,0 +1,59 @@
|
|||||||
|
package message
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// rateLimiter 是每端一个令牌桶:速率 rps、容量 burst。
|
||||||
|
// rps<=0 表示不限速。
|
||||||
|
type rateLimiter struct {
|
||||||
|
rps float64
|
||||||
|
burst float64
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
|
m map[string]*tokenBucket
|
||||||
|
}
|
||||||
|
|
||||||
|
type tokenBucket struct {
|
||||||
|
tokens float64
|
||||||
|
last time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
func newRateLimiter(rps float64, burst int) *rateLimiter {
|
||||||
|
if burst <= 0 {
|
||||||
|
burst = defaultRequestBurst
|
||||||
|
}
|
||||||
|
return &rateLimiter{
|
||||||
|
rps: rps,
|
||||||
|
burst: float64(burst),
|
||||||
|
m: make(map[string]*tokenBucket),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// allow 消耗 1 个令牌;允许则 true。
|
||||||
|
func (r *rateLimiter) allow(endpointID string, now time.Time) bool {
|
||||||
|
if r == nil || r.rps <= 0 {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
b := r.m[endpointID]
|
||||||
|
if b == nil {
|
||||||
|
b = &tokenBucket{tokens: r.burst, last: now}
|
||||||
|
r.m[endpointID] = b
|
||||||
|
}
|
||||||
|
elapsed := now.Sub(b.last).Seconds()
|
||||||
|
if elapsed > 0 {
|
||||||
|
b.tokens += elapsed * r.rps
|
||||||
|
if b.tokens > r.burst {
|
||||||
|
b.tokens = r.burst
|
||||||
|
}
|
||||||
|
b.last = now
|
||||||
|
}
|
||||||
|
if b.tokens < 1 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
b.tokens--
|
||||||
|
return true
|
||||||
|
}
|
||||||
@@ -1,25 +0,0 @@
|
|||||||
package message
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"errors"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
|
||||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestStubSubmitNotImplemented(t *testing.T) {
|
|
||||||
s := NewStub()
|
|
||||||
_, err := s.Submit(context.Background(), "a", port.ConnInfo{}, &protocol.Send{})
|
|
||||||
if !errors.Is(err, ErrNotImplemented) {
|
|
||||||
t.Fatalf("got %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestStubRecoverNoop(t *testing.T) {
|
|
||||||
s := NewStub()
|
|
||||||
if err := s.RecoverOnStart(context.Background()); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -0,0 +1,560 @@
|
|||||||
|
package message
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"encoding/hex"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Submit 处理发送提交(DEVELOPMENT 7.3):防重 → 校验/配额/授权 → 写入 → 到点则最小分发。
|
||||||
|
func (a *App) Submit(ctx context.Context, senderID string, conn port.ConnInfo, req *protocol.Send) (SubmitResult, error) {
|
||||||
|
if req == nil {
|
||||||
|
return SubmitResult{}, errCode(protocol.CodeBadRequest, "nil send")
|
||||||
|
}
|
||||||
|
if senderID == "" || !protocol.ValidEndpointID(senderID) {
|
||||||
|
return SubmitResult{}, errCode(protocol.CodeBadRequest, "invalid sender")
|
||||||
|
}
|
||||||
|
now := a.now()
|
||||||
|
if !a.rates.allow(senderID, now) {
|
||||||
|
return SubmitResult{}, errCode(protocol.CodeRateLimited, "request rate exceeded")
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := req.Validate(a.protocolLimits()); err != nil {
|
||||||
|
return SubmitResult{}, err
|
||||||
|
}
|
||||||
|
fpHex, err := protocol.RequestFingerprint(req)
|
||||||
|
if err != nil {
|
||||||
|
return SubmitResult{}, err
|
||||||
|
}
|
||||||
|
fp, err := hex.DecodeString(fpHex)
|
||||||
|
if err != nil || len(fp) != 32 {
|
||||||
|
return SubmitResult{}, fmt.Errorf("message: fingerprint decode: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
body, err := protocol.DecodeBody(req.Body)
|
||||||
|
if err != nil {
|
||||||
|
return SubmitResult{}, err
|
||||||
|
}
|
||||||
|
metaJSON, err := protocol.MetaCanonicalJSON(req.Meta)
|
||||||
|
if err != nil {
|
||||||
|
return SubmitResult{}, errCode(protocol.CodeBadRequest, "invalid meta")
|
||||||
|
}
|
||||||
|
contentType := protocol.EffectiveContentType(req.Body)
|
||||||
|
keep := protocol.EffectiveOfflineKeep(req)
|
||||||
|
ttl := protocol.EffectiveOfflineTTL(req)
|
||||||
|
receipt := protocol.EffectiveReceipt(req)
|
||||||
|
if keep && a.lim.MaxTTLSeconds > 0 && ttl > a.lim.MaxTTLSeconds {
|
||||||
|
return SubmitResult{}, errCode(protocol.CodeBadRequest, "ttl_seconds exceeds max_ttl_seconds")
|
||||||
|
}
|
||||||
|
|
||||||
|
nowMs := now.UnixMilli()
|
||||||
|
|
||||||
|
// 防重命中可在读连接快速返回;写路径仍会再查一次以防竞态。
|
||||||
|
if res, hit, lookupErr := a.lookupIdempotent(ctx, senderID, req.ID, fp); lookupErr != nil {
|
||||||
|
return SubmitResult{}, lookupErr
|
||||||
|
} else if hit {
|
||||||
|
return res, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
sender, err := a.loadEndpoint(ctx, senderID)
|
||||||
|
if err != nil {
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return SubmitResult{}, errCode(protocol.CodeInvalidTarget, "sender not found")
|
||||||
|
}
|
||||||
|
return SubmitResult{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
sendAt, err := a.computeSendAt(req, sender.DefaultDelayMs, nowMs)
|
||||||
|
if err != nil {
|
||||||
|
return SubmitResult{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
needPassword bool
|
||||||
|
talkPHC string
|
||||||
|
targetEp *endpointRow
|
||||||
|
)
|
||||||
|
|
||||||
|
switch req.To.Kind {
|
||||||
|
case protocol.TargetEndpoint:
|
||||||
|
targetEp, err = a.loadEndpoint(ctx, req.To.ID)
|
||||||
|
if err != nil {
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return SubmitResult{}, errCode(protocol.CodeInvalidTarget, "target not found")
|
||||||
|
}
|
||||||
|
return SubmitResult{}, err
|
||||||
|
}
|
||||||
|
if targetEp.Enabled == 0 {
|
||||||
|
return SubmitResult{}, errCode(protocol.CodeEndpointDisabled, "target disabled")
|
||||||
|
}
|
||||||
|
if senderID != req.To.ID {
|
||||||
|
needPassword, talkPHC, _, err = a.dmAuthNeeded(ctx, senderID, targetEp)
|
||||||
|
if err != nil {
|
||||||
|
return SubmitResult{}, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case protocol.TargetGroup:
|
||||||
|
exists, member, gErr := a.groupMembership(ctx, req.To.ID, senderID)
|
||||||
|
if gErr != nil {
|
||||||
|
return SubmitResult{}, gErr
|
||||||
|
}
|
||||||
|
if !exists {
|
||||||
|
return SubmitResult{}, errCode(protocol.CodeInvalidTarget, "group not found")
|
||||||
|
}
|
||||||
|
if !member {
|
||||||
|
return SubmitResult{}, errCode(protocol.CodeNotMember, "not a group member")
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return SubmitResult{}, errCode(protocol.CodeBadRequest, "invalid to.kind")
|
||||||
|
}
|
||||||
|
passwordVerified := false
|
||||||
|
if needPassword {
|
||||||
|
if locked, _ := a.talkLocked(senderID, req.To.ID, conn.RemoteIP); locked {
|
||||||
|
return SubmitResult{}, errCode(protocol.CodeRateLimited, "talk password locked")
|
||||||
|
}
|
||||||
|
if req.TalkPassword == "" {
|
||||||
|
return SubmitResult{}, errCode(protocol.CodeTalkPasswordRequired, "talk password required")
|
||||||
|
}
|
||||||
|
if a.hash == nil {
|
||||||
|
return SubmitResult{}, fmt.Errorf("message: hash pool required")
|
||||||
|
}
|
||||||
|
ok, vErr := a.hash.Verify(ctx, auth.PasswordTalk, req.TalkPassword, talkPHC)
|
||||||
|
if vErr != nil {
|
||||||
|
return SubmitResult{}, vErr
|
||||||
|
}
|
||||||
|
if !ok {
|
||||||
|
a.talkFail(senderID, req.To.ID, conn.RemoteIP)
|
||||||
|
return SubmitResult{}, errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid")
|
||||||
|
}
|
||||||
|
passwordVerified = true
|
||||||
|
a.talkClear(senderID, req.To.ID)
|
||||||
|
}
|
||||||
|
|
||||||
|
keepInt := 0
|
||||||
|
if keep {
|
||||||
|
keepInt = 1
|
||||||
|
}
|
||||||
|
receiptInt := 0
|
||||||
|
if receipt {
|
||||||
|
receiptInt = 1
|
||||||
|
}
|
||||||
|
|
||||||
|
var result SubmitResult
|
||||||
|
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||||
|
if res, hit, e := lookupIdempotentTx(tx, senderID, req.ID, fp); e != nil {
|
||||||
|
return e
|
||||||
|
} else if hit {
|
||||||
|
result = res
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if e := checkQuotaTx(tx, senderID, a.lim.MaxPendingPerSender); e != nil {
|
||||||
|
return e
|
||||||
|
}
|
||||||
|
|
||||||
|
// 写事务内再确认目标与授权(防并发停用/退群)。
|
||||||
|
switch req.To.Kind {
|
||||||
|
case protocol.TargetEndpoint:
|
||||||
|
ep, e := loadEndpointTx(tx, req.To.ID)
|
||||||
|
if e != nil {
|
||||||
|
if errors.Is(e, sql.ErrNoRows) {
|
||||||
|
return errCode(protocol.CodeInvalidTarget, "target not found")
|
||||||
|
}
|
||||||
|
return e
|
||||||
|
}
|
||||||
|
if ep.Enabled == 0 {
|
||||||
|
return errCode(protocol.CodeEndpointDisabled, "target disabled")
|
||||||
|
}
|
||||||
|
if senderID != req.To.ID {
|
||||||
|
needed, phc, ver, ae := dmAuthNeededTx(tx, senderID, ep)
|
||||||
|
if ae != nil {
|
||||||
|
return ae
|
||||||
|
}
|
||||||
|
if needed {
|
||||||
|
if !passwordVerified {
|
||||||
|
if req.TalkPassword == "" {
|
||||||
|
return errCode(protocol.CodeTalkPasswordRequired, "talk password required")
|
||||||
|
}
|
||||||
|
return errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid")
|
||||||
|
}
|
||||||
|
// 密码版本在校验后变化则拒绝,避免写过期授权。
|
||||||
|
if ep.TalkHash == nil || *ep.TalkHash != phc || ep.TalkVersion != ver {
|
||||||
|
return errCode(protocol.CodeTalkPasswordInvalid, "talk password changed")
|
||||||
|
}
|
||||||
|
if ge := upsertGrantTx(tx, senderID, req.To.ID, ver, GrantKindPassword, nowMs); ge != nil {
|
||||||
|
return ge
|
||||||
|
}
|
||||||
|
} else if passwordVerified {
|
||||||
|
// 已有授权或未设防:带对密码时仍可刷新授权(文档:带对了则写入或更新)。
|
||||||
|
if ep.TalkHash != nil && *ep.TalkHash != "" {
|
||||||
|
if ge := upsertGrantTx(tx, senderID, req.To.ID, ep.TalkVersion, GrantKindPassword, nowMs); ge != nil {
|
||||||
|
return ge
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// 发送方设了对话密码且发给别人的单聊:给对方写回复授权。
|
||||||
|
snd, se := loadEndpointTx(tx, senderID)
|
||||||
|
if se != nil {
|
||||||
|
return se
|
||||||
|
}
|
||||||
|
if senderID != req.To.ID && snd.TalkHash != nil && *snd.TalkHash != "" {
|
||||||
|
if ge := upsertGrantTx(tx, req.To.ID, senderID, snd.TalkVersion, GrantKindReply, nowMs); ge != nil {
|
||||||
|
return ge
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case protocol.TargetGroup:
|
||||||
|
exists, member, e := groupMembershipTx(tx, req.To.ID, senderID)
|
||||||
|
if e != nil {
|
||||||
|
return e
|
||||||
|
}
|
||||||
|
if !exists {
|
||||||
|
return errCode(protocol.CodeInvalidTarget, "group not found")
|
||||||
|
}
|
||||||
|
if !member {
|
||||||
|
return errCode(protocol.CodeNotMember, "not a group member")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
state := StateScheduled
|
||||||
|
res, e := tx.Exec(`
|
||||||
|
INSERT INTO messages(
|
||||||
|
id, sender_id, dest_kind, dest_id, meta, content_type, body_enc,
|
||||||
|
send_at, keep, ttl_seconds, receipt, state, reason, created_at
|
||||||
|
) VALUES(?,?,?,?,?,?,?,?,?,?,?,?, '', ?)`,
|
||||||
|
req.ID, senderID, req.To.Kind, req.To.ID, string(metaJSON), contentType, req.Body.Enc,
|
||||||
|
sendAt, keepInt, ttl, receiptInt, state, nowMs,
|
||||||
|
)
|
||||||
|
if e != nil {
|
||||||
|
return e
|
||||||
|
}
|
||||||
|
seq, e := res.LastInsertId()
|
||||||
|
if e != nil {
|
||||||
|
return e
|
||||||
|
}
|
||||||
|
if _, e = tx.Exec(`INSERT INTO message_bodies(seq, body) VALUES(?, ?)`, seq, body); e != nil {
|
||||||
|
return e
|
||||||
|
}
|
||||||
|
if _, e = tx.Exec(
|
||||||
|
`INSERT INTO send_keys(sender_id, msg_id, request_sha256, created_at) VALUES(?,?,?,?)`,
|
||||||
|
senderID, req.ID, fp, nowMs,
|
||||||
|
); e != nil {
|
||||||
|
return e
|
||||||
|
}
|
||||||
|
|
||||||
|
finalState := state
|
||||||
|
if sendAt <= nowMs {
|
||||||
|
finalState, e = dispatchMinimalTx(tx, seq, senderID, req.To.Kind, req.To.ID, sendAt, keepInt, nowMs)
|
||||||
|
if e != nil {
|
||||||
|
return e
|
||||||
|
}
|
||||||
|
}
|
||||||
|
result = SubmitResult{ID: req.ID, SendAtMs: sendAt, State: finalState}
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return SubmitResult{}, err
|
||||||
|
}
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) computeSendAt(req *protocol.Send, defaultDelayMs, nowMs int64) (int64, error) {
|
||||||
|
if req.SendAtMs != nil && req.DelayMs != nil {
|
||||||
|
return 0, errCode(protocol.CodeBadRequest, "send_at_ms and delay_ms are mutually exclusive")
|
||||||
|
}
|
||||||
|
var sendAt int64
|
||||||
|
switch {
|
||||||
|
case req.SendAtMs != nil:
|
||||||
|
sendAt = *req.SendAtMs
|
||||||
|
case req.DelayMs != nil:
|
||||||
|
if *req.DelayMs < 0 {
|
||||||
|
return 0, errCode(protocol.CodeBadRequest, "delay_ms negative")
|
||||||
|
}
|
||||||
|
sendAt = nowMs + *req.DelayMs
|
||||||
|
default:
|
||||||
|
if defaultDelayMs < 0 {
|
||||||
|
defaultDelayMs = 0
|
||||||
|
}
|
||||||
|
sendAt = nowMs + defaultDelayMs
|
||||||
|
}
|
||||||
|
if a.lim.MaxScheduleSeconds > 0 {
|
||||||
|
maxAt := nowMs + a.lim.MaxScheduleSeconds*1000
|
||||||
|
if sendAt > maxAt {
|
||||||
|
return 0, errCode(protocol.CodeBadRequest, "send time exceeds max_schedule_seconds")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return sendAt, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type endpointRow struct {
|
||||||
|
ID string
|
||||||
|
DefaultDelayMs int64
|
||||||
|
TalkHash *string
|
||||||
|
TalkVersion int64
|
||||||
|
Enabled int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) loadEndpoint(ctx context.Context, id string) (*endpointRow, error) {
|
||||||
|
row := a.db.Read.QueryRowContext(ctx, `
|
||||||
|
SELECT id, default_delay_ms, talk_hash, talk_version, enabled
|
||||||
|
FROM endpoints WHERE id = ?`, id)
|
||||||
|
return scanEndpoint(row)
|
||||||
|
}
|
||||||
|
|
||||||
|
func loadEndpointTx(tx *sql.Tx, id string) (*endpointRow, error) {
|
||||||
|
row := tx.QueryRow(`
|
||||||
|
SELECT id, default_delay_ms, talk_hash, talk_version, enabled
|
||||||
|
FROM endpoints WHERE id = ?`, id)
|
||||||
|
return scanEndpoint(row)
|
||||||
|
}
|
||||||
|
|
||||||
|
func scanEndpoint(row *sql.Row) (*endpointRow, error) {
|
||||||
|
var ep endpointRow
|
||||||
|
var talk sql.NullString
|
||||||
|
if err := row.Scan(&ep.ID, &ep.DefaultDelayMs, &talk, &ep.TalkVersion, &ep.Enabled); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if talk.Valid {
|
||||||
|
s := talk.String
|
||||||
|
ep.TalkHash = &s
|
||||||
|
}
|
||||||
|
return &ep, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) groupMembership(ctx context.Context, groupID, endpointID string) (exists, member bool, err error) {
|
||||||
|
var one int
|
||||||
|
err = a.db.Read.QueryRowContext(ctx, `SELECT 1 FROM groups WHERE id = ?`, groupID).Scan(&one)
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return false, false, nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return false, false, err
|
||||||
|
}
|
||||||
|
err = a.db.Read.QueryRowContext(ctx,
|
||||||
|
`SELECT 1 FROM group_members WHERE group_id = ? AND endpoint_id = ?`, groupID, endpointID,
|
||||||
|
).Scan(&one)
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return true, false, nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return true, false, err
|
||||||
|
}
|
||||||
|
return true, true, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func groupMembershipTx(tx *sql.Tx, groupID, endpointID string) (exists, member bool, err error) {
|
||||||
|
var one int
|
||||||
|
err = tx.QueryRow(`SELECT 1 FROM groups WHERE id = ?`, groupID).Scan(&one)
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return false, false, nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return false, false, err
|
||||||
|
}
|
||||||
|
err = tx.QueryRow(
|
||||||
|
`SELECT 1 FROM group_members WHERE group_id = ? AND endpoint_id = ?`, groupID, endpointID,
|
||||||
|
).Scan(&one)
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return true, false, nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return true, false, err
|
||||||
|
}
|
||||||
|
return true, true, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// dmAuthNeeded 返回是否需要对话密码,以及对方当前 talk_hash / version。
|
||||||
|
func (a *App) dmAuthNeeded(ctx context.Context, senderID string, target *endpointRow) (needed bool, phc string, version int64, err error) {
|
||||||
|
if target.TalkHash == nil || *target.TalkHash == "" {
|
||||||
|
return false, "", target.TalkVersion, nil
|
||||||
|
}
|
||||||
|
ok, err := hasValidGrant(ctx, a.db.Read, senderID, target.ID, target.TalkVersion)
|
||||||
|
if err != nil {
|
||||||
|
return false, "", 0, err
|
||||||
|
}
|
||||||
|
if ok {
|
||||||
|
return false, *target.TalkHash, target.TalkVersion, nil
|
||||||
|
}
|
||||||
|
return true, *target.TalkHash, target.TalkVersion, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func dmAuthNeededTx(tx *sql.Tx, senderID string, target *endpointRow) (needed bool, phc string, version int64, err error) {
|
||||||
|
if target.TalkHash == nil || *target.TalkHash == "" {
|
||||||
|
return false, "", target.TalkVersion, nil
|
||||||
|
}
|
||||||
|
ok, err := hasValidGrantTx(tx, senderID, target.ID, target.TalkVersion)
|
||||||
|
if err != nil {
|
||||||
|
return false, "", 0, err
|
||||||
|
}
|
||||||
|
if ok {
|
||||||
|
return false, *target.TalkHash, target.TalkVersion, nil
|
||||||
|
}
|
||||||
|
return true, *target.TalkHash, target.TalkVersion, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func hasValidGrant(ctx context.Context, db *sql.DB, senderID, targetID string, talkVersion int64) (bool, error) {
|
||||||
|
var n int
|
||||||
|
err := db.QueryRowContext(ctx, `
|
||||||
|
SELECT 1 FROM talk_grants
|
||||||
|
WHERE sender_id = ? AND target_id = ? AND target_talk_version = ?
|
||||||
|
LIMIT 1`, senderID, targetID, talkVersion).Scan(&n)
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
return err == nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func hasValidGrantTx(tx *sql.Tx, senderID, targetID string, talkVersion int64) (bool, error) {
|
||||||
|
var n int
|
||||||
|
err := tx.QueryRow(`
|
||||||
|
SELECT 1 FROM talk_grants
|
||||||
|
WHERE sender_id = ? AND target_id = ? AND target_talk_version = ?
|
||||||
|
LIMIT 1`, senderID, targetID, talkVersion).Scan(&n)
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
return err == nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func upsertGrantTx(tx *sql.Tx, senderID, targetID string, talkVersion int64, kind string, nowMs int64) error {
|
||||||
|
_, err := tx.Exec(`
|
||||||
|
INSERT INTO talk_grants(sender_id, target_id, target_talk_version, kind, created_at)
|
||||||
|
VALUES(?,?,?,?,?)
|
||||||
|
ON CONFLICT(sender_id, target_id) DO UPDATE SET
|
||||||
|
target_talk_version = excluded.target_talk_version,
|
||||||
|
kind = excluded.kind,
|
||||||
|
created_at = excluded.created_at`,
|
||||||
|
senderID, targetID, talkVersion, kind, nowMs,
|
||||||
|
)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func checkQuotaTx(tx *sql.Tx, senderID string, maxPending int) error {
|
||||||
|
if maxPending <= 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
var n int
|
||||||
|
err := tx.QueryRow(`
|
||||||
|
SELECT COUNT(*) FROM messages
|
||||||
|
WHERE sender_id = ? AND state IN ('scheduled', 'dispatched')`, senderID).Scan(&n)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if n >= maxPending {
|
||||||
|
return errCode(protocol.CodeQuotaExceeded, "max_pending_per_sender exceeded")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) lookupIdempotent(ctx context.Context, senderID, msgID string, fp []byte) (SubmitResult, bool, error) {
|
||||||
|
var stored []byte
|
||||||
|
err := a.db.Read.QueryRowContext(ctx, `
|
||||||
|
SELECT request_sha256 FROM send_keys WHERE sender_id = ? AND msg_id = ?`,
|
||||||
|
senderID, msgID,
|
||||||
|
).Scan(&stored)
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return SubmitResult{}, false, nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return SubmitResult{}, false, err
|
||||||
|
}
|
||||||
|
if !bytesEqual(stored, fp) {
|
||||||
|
return SubmitResult{}, false, errCode(protocol.CodeConflict, "message id conflict")
|
||||||
|
}
|
||||||
|
res, err := loadSubmitResult(ctx, a.db.Read, senderID, msgID)
|
||||||
|
if err != nil {
|
||||||
|
return SubmitResult{}, false, err
|
||||||
|
}
|
||||||
|
return res, true, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func lookupIdempotentTx(tx *sql.Tx, senderID, msgID string, fp []byte) (SubmitResult, bool, error) {
|
||||||
|
var stored []byte
|
||||||
|
err := tx.QueryRow(`
|
||||||
|
SELECT request_sha256 FROM send_keys WHERE sender_id = ? AND msg_id = ?`,
|
||||||
|
senderID, msgID,
|
||||||
|
).Scan(&stored)
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return SubmitResult{}, false, nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return SubmitResult{}, false, err
|
||||||
|
}
|
||||||
|
if !bytesEqual(stored, fp) {
|
||||||
|
return SubmitResult{}, false, errCode(protocol.CodeConflict, "message id conflict")
|
||||||
|
}
|
||||||
|
res, err := loadSubmitResultTx(tx, senderID, msgID)
|
||||||
|
if err != nil {
|
||||||
|
return SubmitResult{}, false, err
|
||||||
|
}
|
||||||
|
return res, true, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func loadSubmitResult(ctx context.Context, db *sql.DB, senderID, msgID string) (SubmitResult, error) {
|
||||||
|
var res SubmitResult
|
||||||
|
err := db.QueryRowContext(ctx, `
|
||||||
|
SELECT id, send_at, state FROM messages WHERE sender_id = ? AND id = ?`,
|
||||||
|
senderID, msgID,
|
||||||
|
).Scan(&res.ID, &res.SendAtMs, &res.State)
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return SubmitResult{}, errCode(protocol.CodeNotFound, "idempotent key without message")
|
||||||
|
}
|
||||||
|
return res, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func loadSubmitResultTx(tx *sql.Tx, senderID, msgID string) (SubmitResult, error) {
|
||||||
|
var res SubmitResult
|
||||||
|
err := tx.QueryRow(`
|
||||||
|
SELECT id, send_at, state FROM messages WHERE sender_id = ? AND id = ?`,
|
||||||
|
senderID, msgID,
|
||||||
|
).Scan(&res.ID, &res.SendAtMs, &res.State)
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return SubmitResult{}, errCode(protocol.CodeNotFound, "idempotent key without message")
|
||||||
|
}
|
||||||
|
return res, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func bytesEqual(a, b []byte) bool {
|
||||||
|
if len(a) != len(b) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
var v byte
|
||||||
|
for i := range a {
|
||||||
|
v |= a[i] ^ b[i]
|
||||||
|
}
|
||||||
|
return v == 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) talkLocked(senderID, targetID, ip string) (bool, error) {
|
||||||
|
if a.locks == nil {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: ip}); locked {
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
|
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}); locked {
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) talkFail(senderID, targetID, ip string) {
|
||||||
|
if a.locks == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: ip})
|
||||||
|
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) talkClear(senderID, targetID string) {
|
||||||
|
if a.locks == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
a.locks.Clear(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID})
|
||||||
|
}
|
||||||
@@ -0,0 +1,382 @@
|
|||||||
|
package message
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"errors"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/config"
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestStubSubmitNotImplemented(t *testing.T) {
|
||||||
|
s := NewStub()
|
||||||
|
_, err := s.Submit(context.Background(), "a", port.ConnInfo{}, &protocol.Send{})
|
||||||
|
if !errors.Is(err, ErrNotImplemented) {
|
||||||
|
t.Fatalf("got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStubRecoverNoop(t *testing.T) {
|
||||||
|
s := NewStub()
|
||||||
|
if err := s.RecoverOnStart(context.Background()); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func openTestApp(t *testing.T, lim Limits) (*App, *store.DB) {
|
||||||
|
t.Helper()
|
||||||
|
dir := t.TempDir()
|
||||||
|
db, err := store.Open(filepath.Join(dir, "data"), "FULL")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { _ = db.Close() })
|
||||||
|
fixed := time.UnixMilli(1_700_000_000_000)
|
||||||
|
app := New(db, lim, auth.NewStubHashPool(),
|
||||||
|
WithNow(func() time.Time { return fixed }),
|
||||||
|
WithLocks(auth.NewStubLoginLocks()),
|
||||||
|
)
|
||||||
|
return app, db
|
||||||
|
}
|
||||||
|
|
||||||
|
func defaultTestLimits() Limits {
|
||||||
|
cfg := config.Default().Limits
|
||||||
|
lim := LimitsFromConfig(cfg)
|
||||||
|
lim.RequestsPerSecond = 0 // 测试默认不限速
|
||||||
|
return lim
|
||||||
|
}
|
||||||
|
|
||||||
|
func insertEndpoint(t *testing.T, db *store.DB, id string, talkPassword string, enabled int, defaultDelayMs int64) {
|
||||||
|
t.Helper()
|
||||||
|
ctx := context.Background()
|
||||||
|
var talk any
|
||||||
|
var talkVer int64
|
||||||
|
if talkPassword != "" {
|
||||||
|
talk = "stub$" + talkPassword
|
||||||
|
talkVer = 1
|
||||||
|
}
|
||||||
|
err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||||
|
_, e := tx.Exec(`
|
||||||
|
INSERT INTO endpoints(id, name, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at)
|
||||||
|
VALUES(?,?,?,?,?,?,?,?)`,
|
||||||
|
id, id, "stub$login", talk, talkVer, defaultDelayMs, enabled, 1_700_000_000_000)
|
||||||
|
return e
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func baseSend(id, to string) *protocol.Send {
|
||||||
|
return &protocol.Send{
|
||||||
|
V: protocol.Version,
|
||||||
|
Type: protocol.TypeSend,
|
||||||
|
RID: "r1",
|
||||||
|
ID: id,
|
||||||
|
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: to},
|
||||||
|
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hello"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func protoCode(err error) string {
|
||||||
|
var pe *protocol.Error
|
||||||
|
if errors.As(err, &pe) {
|
||||||
|
return pe.Code
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSubmitTable(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
t.Run("idempotent_hit", func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
lim := defaultTestLimits()
|
||||||
|
app, db := openTestApp(t, lim)
|
||||||
|
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||||
|
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||||
|
ctx := context.Background()
|
||||||
|
req := baseSend("msg-1", "bob")
|
||||||
|
first, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if first.State != StateDispatched {
|
||||||
|
t.Fatalf("state=%s", first.State)
|
||||||
|
}
|
||||||
|
second, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if second != first {
|
||||||
|
t.Fatalf("want %+v got %+v", first, second)
|
||||||
|
}
|
||||||
|
var n int
|
||||||
|
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM messages WHERE sender_id=? AND id=?`, "alice", "msg-1").Scan(&n); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if n != 1 {
|
||||||
|
t.Fatalf("messages=%d", n)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("conflict", func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
lim := defaultTestLimits()
|
||||||
|
app, db := openTestApp(t, lim)
|
||||||
|
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||||
|
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||||
|
ctx := context.Background()
|
||||||
|
req := baseSend("msg-2", "bob")
|
||||||
|
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
other := baseSend("msg-2", "bob")
|
||||||
|
other.Body.Data = "other"
|
||||||
|
_, err := app.Submit(ctx, "alice", port.ConnInfo{}, other)
|
||||||
|
if protoCode(err) != protocol.CodeConflict {
|
||||||
|
t.Fatalf("want conflict got %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("quota_exceeded", func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
lim := defaultTestLimits()
|
||||||
|
lim.MaxPendingPerSender = 1
|
||||||
|
app, db := openTestApp(t, lim)
|
||||||
|
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||||
|
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||||
|
ctx := context.Background()
|
||||||
|
delay := int64(60_000)
|
||||||
|
req1 := baseSend("q1", "bob")
|
||||||
|
req1.DelayMs = &delay
|
||||||
|
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, req1); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
req2 := baseSend("q2", "bob")
|
||||||
|
req2.DelayMs = &delay
|
||||||
|
_, err := app.Submit(ctx, "alice", port.ConnInfo{}, req2)
|
||||||
|
if protoCode(err) != protocol.CodeQuotaExceeded {
|
||||||
|
t.Fatalf("want quota_exceeded got %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("auth_required_and_grant", func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
lim := defaultTestLimits()
|
||||||
|
app, db := openTestApp(t, lim)
|
||||||
|
insertEndpoint(t, db, "alice", "alice-secret", 1, 0)
|
||||||
|
insertEndpoint(t, db, "bob", "secret", 1, 0)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
_, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("a1", "bob"))
|
||||||
|
if protoCode(err) != protocol.CodeTalkPasswordRequired {
|
||||||
|
t.Fatalf("want talk_password_required got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
bad := baseSend("a2", "bob")
|
||||||
|
bad.TalkPassword = "wrong"
|
||||||
|
_, err = app.Submit(ctx, "alice", port.ConnInfo{}, bad)
|
||||||
|
if protoCode(err) != protocol.CodeTalkPasswordInvalid {
|
||||||
|
t.Fatalf("want talk_password_invalid got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
okReq := baseSend("a3", "bob")
|
||||||
|
okReq.TalkPassword = "secret"
|
||||||
|
res, err := app.Submit(ctx, "alice", port.ConnInfo{}, okReq)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if res.State != StateDispatched {
|
||||||
|
t.Fatalf("state=%s", res.State)
|
||||||
|
}
|
||||||
|
// 已有授权后不带密码也可发
|
||||||
|
if _, submitErr := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("a4", "bob")); submitErr != nil {
|
||||||
|
t.Fatal(submitErr)
|
||||||
|
}
|
||||||
|
// 回复授权:bob→alice(因 alice 设了对话密码)
|
||||||
|
var kind string
|
||||||
|
err = db.Read.QueryRow(`
|
||||||
|
SELECT kind FROM talk_grants WHERE sender_id=? AND target_id=?`, "bob", "alice").Scan(&kind)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if kind != GrantKindReply {
|
||||||
|
t.Fatalf("reply grant kind=%s", kind)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("self_skip_talk_password", func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
lim := defaultTestLimits()
|
||||||
|
app, db := openTestApp(t, lim)
|
||||||
|
insertEndpoint(t, db, "alice", "secret", 1, 0)
|
||||||
|
ctx := context.Background()
|
||||||
|
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("self1", "alice")); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("delay_and_send_at_mutex", func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
lim := defaultTestLimits()
|
||||||
|
app, db := openTestApp(t, lim)
|
||||||
|
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||||
|
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||||
|
ctx := context.Background()
|
||||||
|
delay := int64(1000)
|
||||||
|
sendAt := int64(1_700_000_001_000)
|
||||||
|
req := baseSend("m-mutex", "bob")
|
||||||
|
req.DelayMs = &delay
|
||||||
|
req.SendAtMs = &sendAt
|
||||||
|
_, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
|
||||||
|
if protoCode(err) != protocol.CodeBadRequest {
|
||||||
|
t.Fatalf("want bad_request got %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("idempotent_before_disabled_check", func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
lim := defaultTestLimits()
|
||||||
|
app, db := openTestApp(t, lim)
|
||||||
|
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||||
|
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||||
|
ctx := context.Background()
|
||||||
|
req := baseSend("pre-disable", "bob")
|
||||||
|
first, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
err = db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||||
|
_, e := tx.Exec(`UPDATE endpoints SET enabled = 0 WHERE id = ?`, "bob")
|
||||||
|
return e
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
// 新消息应失败
|
||||||
|
_, err = app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("after-disable", "bob"))
|
||||||
|
if protoCode(err) != protocol.CodeEndpointDisabled {
|
||||||
|
t.Fatalf("want endpoint_disabled got %v", err)
|
||||||
|
}
|
||||||
|
// 原请求重试仍返回原结果
|
||||||
|
second, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if second != first {
|
||||||
|
t.Fatalf("want %+v got %+v", first, second)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("scheduled_not_dispatched", func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
lim := defaultTestLimits()
|
||||||
|
app, db := openTestApp(t, lim)
|
||||||
|
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||||
|
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||||
|
ctx := context.Background()
|
||||||
|
delay := int64(10_000)
|
||||||
|
req := baseSend("sched-1", "bob")
|
||||||
|
req.DelayMs = &delay
|
||||||
|
res, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if res.State != StateScheduled {
|
||||||
|
t.Fatalf("state=%s", res.State)
|
||||||
|
}
|
||||||
|
var n int
|
||||||
|
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries`).Scan(&n); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if n != 0 {
|
||||||
|
t.Fatalf("deliveries=%d", n)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("group_dispatch_excludes_sender", func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
lim := defaultTestLimits()
|
||||||
|
app, db := openTestApp(t, lim)
|
||||||
|
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||||
|
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||||
|
insertEndpoint(t, db, "carol", "", 1, 0)
|
||||||
|
ctx := context.Background()
|
||||||
|
err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||||
|
if _, e := tx.Exec(`INSERT INTO groups(id, name, owner_id, created_at) VALUES(?,?,?,?)`,
|
||||||
|
"g1", "g", "alice", 1_700_000_000_000); e != nil {
|
||||||
|
return e
|
||||||
|
}
|
||||||
|
for _, m := range []string{"alice", "bob", "carol"} {
|
||||||
|
if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||||
|
"g1", m, 1_700_000_000_000); e != nil {
|
||||||
|
return e
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
req := &protocol.Send{
|
||||||
|
V: protocol.Version,
|
||||||
|
Type: protocol.TypeSend,
|
||||||
|
RID: "r1",
|
||||||
|
ID: "gmsg-1",
|
||||||
|
To: protocol.Target{Kind: protocol.TargetGroup, ID: "g1"},
|
||||||
|
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hi"},
|
||||||
|
}
|
||||||
|
res, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if res.State != StateDispatched {
|
||||||
|
t.Fatalf("state=%s", res.State)
|
||||||
|
}
|
||||||
|
rows, err := db.Read.Query(`SELECT endpoint_id FROM deliveries ORDER BY endpoint_id`)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer func() { _ = rows.Close() }()
|
||||||
|
var got []string
|
||||||
|
for rows.Next() {
|
||||||
|
var id string
|
||||||
|
if err := rows.Scan(&id); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
got = append(got, id)
|
||||||
|
}
|
||||||
|
if len(got) != 2 || got[0] != "bob" || got[1] != "carol" {
|
||||||
|
t.Fatalf("recipients=%v", got)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("rate_limited", func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
lim := defaultTestLimits()
|
||||||
|
lim.RequestsPerSecond = 50
|
||||||
|
lim.RequestBurst = 2
|
||||||
|
app, db := openTestApp(t, lim)
|
||||||
|
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||||
|
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||||
|
ctx := context.Background()
|
||||||
|
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("r1", "bob")); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("r2", "bob")); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
_, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("r3", "bob"))
|
||||||
|
if protoCode(err) != protocol.CodeRateLimited {
|
||||||
|
t.Fatalf("want rate_limited got %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -0,0 +1,140 @@
|
|||||||
|
package auth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestPoolLimitsConcurrency(t *testing.T) {
|
||||||
|
pool := NewPoolSize(1)
|
||||||
|
errCh := make(chan error, 3)
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
wg.Add(3)
|
||||||
|
for i := 0; i < 3; i++ {
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
_, err := pool.Hash(context.Background(), PasswordLogin, "concurrency-test-password")
|
||||||
|
errCh <- err
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
close(errCh)
|
||||||
|
for err := range errCh {
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if pool.MaxActive() > 1 {
|
||||||
|
t.Fatalf("max active=%d want <=1", pool.MaxActive())
|
||||||
|
}
|
||||||
|
if pool.MaxActive() < 1 {
|
||||||
|
t.Fatal("expected at least one active hash")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPoolHashVerifyRoundTrip(t *testing.T) {
|
||||||
|
pool := NewPoolSize(2)
|
||||||
|
ctx := context.Background()
|
||||||
|
phc, err := pool.Hash(ctx, PasswordAdmin, "round-trip-password-1")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if !strings.HasPrefix(phc, "$argon2id$") {
|
||||||
|
t.Fatalf("phc=%q", phc)
|
||||||
|
}
|
||||||
|
ok, err := pool.Verify(ctx, PasswordAdmin, "round-trip-password-1", phc)
|
||||||
|
if err != nil || !ok {
|
||||||
|
t.Fatalf("ok=%v err=%v", ok, err)
|
||||||
|
}
|
||||||
|
ok, err = pool.Verify(ctx, PasswordAdmin, "wrong-password-xxxxx", phc)
|
||||||
|
if err != nil || ok {
|
||||||
|
t.Fatalf("mismatch ok=%v err=%v", ok, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSessionAndAPITokens(t *testing.T) {
|
||||||
|
s := NewSessionTokens()
|
||||||
|
tok, hash, err := s.Issue(context.Background())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if !strings.HasPrefix(tok, "nst_") {
|
||||||
|
t.Fatalf("tok=%q", tok)
|
||||||
|
}
|
||||||
|
if !s.LooksLikeSessionToken(tok) {
|
||||||
|
t.Fatal("LooksLikeSessionToken")
|
||||||
|
}
|
||||||
|
if !EqualHash(hash, s.HashToken(tok)) {
|
||||||
|
t.Fatal("hash mismatch")
|
||||||
|
}
|
||||||
|
|
||||||
|
a := NewAPITokens()
|
||||||
|
atok, ahash, err := a.Issue(context.Background())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if !strings.HasPrefix(atok, "nxm_") {
|
||||||
|
t.Fatalf("atok=%q", atok)
|
||||||
|
}
|
||||||
|
if !a.LooksLikeAPIToken(atok) || !EqualHash(ahash, a.HashToken(atok)) {
|
||||||
|
t.Fatal("api token hash")
|
||||||
|
}
|
||||||
|
if EqualHash(hash, ahash) {
|
||||||
|
t.Fatal("session and api hashes should differ")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoginLocksWindowAndExpiry(t *testing.T) {
|
||||||
|
locks := NewLoginLocks()
|
||||||
|
now := time.Date(2026, 9, 30, 12, 0, 0, 0, time.UTC)
|
||||||
|
locks.SetClock(func() time.Time { return now })
|
||||||
|
|
||||||
|
key := LockKey{Kind: LockLoginEndpointIP, EndpointID: "e1", IP: "1.2.3.4"}
|
||||||
|
for i := 0; i < 9; i++ {
|
||||||
|
locked, _ := locks.Fail(key)
|
||||||
|
if locked {
|
||||||
|
t.Fatalf("locked early at %d", i+1)
|
||||||
|
}
|
||||||
|
now = now.Add(time.Second)
|
||||||
|
}
|
||||||
|
locked, retry := locks.Fail(key)
|
||||||
|
if !locked || retry <= 0 {
|
||||||
|
t.Fatalf("want lock, locked=%v retry=%v", locked, retry)
|
||||||
|
}
|
||||||
|
locked, _ = locks.Check(key)
|
||||||
|
if !locked {
|
||||||
|
t.Fatal("check should be locked")
|
||||||
|
}
|
||||||
|
|
||||||
|
now = now.Add(5*time.Minute + time.Second)
|
||||||
|
locked, _ = locks.Check(key)
|
||||||
|
if locked {
|
||||||
|
t.Fatal("should unlock after window")
|
||||||
|
}
|
||||||
|
locked, _ = locks.Fail(key)
|
||||||
|
if locked {
|
||||||
|
t.Fatal("after expiry should not still be locked on first fail")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoginLocksClearEndpoint(t *testing.T) {
|
||||||
|
locks := NewLoginLocks()
|
||||||
|
k1 := LockKey{Kind: LockLoginEndpointIP, EndpointID: "e1", IP: "1.1.1.1"}
|
||||||
|
k2 := LockKey{Kind: LockLoginEndpoint, EndpointID: "e1"}
|
||||||
|
k3 := LockKey{Kind: LockLoginEndpointIP, EndpointID: "e2", IP: "1.1.1.1"}
|
||||||
|
for i := 0; i < 10; i++ {
|
||||||
|
locks.Fail(k1)
|
||||||
|
locks.Fail(k2)
|
||||||
|
}
|
||||||
|
locks.Fail(k3)
|
||||||
|
locks.ClearEndpoint("e1")
|
||||||
|
if locked, _ := locks.Check(k1); locked {
|
||||||
|
t.Fatal("e1 ip lock should clear")
|
||||||
|
}
|
||||||
|
if locked, _ := locks.Check(k2); locked {
|
||||||
|
t.Fatal("e1 total lock should clear")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,155 @@
|
|||||||
|
package auth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
type lockPolicy struct {
|
||||||
|
window time.Duration
|
||||||
|
maxFails int
|
||||||
|
lockFor time.Duration
|
||||||
|
}
|
||||||
|
|
||||||
|
func policyFor(kind LockKind) lockPolicy {
|
||||||
|
switch kind {
|
||||||
|
case LockLoginEndpoint, LockTalkTarget:
|
||||||
|
return lockPolicy{window: time.Hour, maxFails: 50, lockFor: time.Hour}
|
||||||
|
default:
|
||||||
|
// LockLoginEndpointIP、LockTalkPair、LockAdminIP、LockRegisterIP
|
||||||
|
return lockPolicy{window: 5 * time.Minute, maxFails: 10, lockFor: 5 * time.Minute}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type lockEntry struct {
|
||||||
|
fails []time.Time
|
||||||
|
lockedUntil time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
// MemoryLocks 是内存锁定计数器(重启清零)。
|
||||||
|
type MemoryLocks struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
entries map[string]*lockEntry
|
||||||
|
now func() time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewLoginLocks 创建默认锁定计数器。
|
||||||
|
func NewLoginLocks() *MemoryLocks {
|
||||||
|
return &MemoryLocks{
|
||||||
|
entries: make(map[string]*lockEntry),
|
||||||
|
now: time.Now,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetClock 注入时钟(测试到期解除)。
|
||||||
|
func (l *MemoryLocks) SetClock(now func() time.Time) {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
if now == nil {
|
||||||
|
l.now = time.Now
|
||||||
|
return
|
||||||
|
}
|
||||||
|
l.now = now
|
||||||
|
}
|
||||||
|
|
||||||
|
func lockMapKey(key LockKey) string {
|
||||||
|
return string(key.Kind) + "|" + key.EndpointID + "|" + key.PeerID + "|" + key.IP
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check 若当前已锁定返回 locked=true 与剩余时间。
|
||||||
|
func (l *MemoryLocks) Check(key LockKey) (bool, time.Duration) {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
now := l.now()
|
||||||
|
e := l.entries[lockMapKey(key)]
|
||||||
|
if e == nil {
|
||||||
|
return false, 0
|
||||||
|
}
|
||||||
|
if e.lockedUntil.After(now) {
|
||||||
|
return true, e.lockedUntil.Sub(now)
|
||||||
|
}
|
||||||
|
// 到期自动解除:清空锁定与窗口内失败(保留结构以便后续 Fail)。
|
||||||
|
if !e.lockedUntil.IsZero() && !e.lockedUntil.After(now) {
|
||||||
|
e.lockedUntil = time.Time{}
|
||||||
|
e.fails = nil
|
||||||
|
}
|
||||||
|
return false, 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fail 记录一次失败;若因此触发锁定,返回 locked=true。
|
||||||
|
func (l *MemoryLocks) Fail(key LockKey) (bool, time.Duration) {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
now := l.now()
|
||||||
|
k := lockMapKey(key)
|
||||||
|
e := l.entries[k]
|
||||||
|
if e == nil {
|
||||||
|
e = &lockEntry{}
|
||||||
|
l.entries[k] = e
|
||||||
|
}
|
||||||
|
if e.lockedUntil.After(now) {
|
||||||
|
return true, e.lockedUntil.Sub(now)
|
||||||
|
}
|
||||||
|
if !e.lockedUntil.IsZero() {
|
||||||
|
e.lockedUntil = time.Time{}
|
||||||
|
e.fails = nil
|
||||||
|
}
|
||||||
|
pol := policyFor(key.Kind)
|
||||||
|
cutoff := now.Add(-pol.window)
|
||||||
|
kept := e.fails[:0]
|
||||||
|
for _, t := range e.fails {
|
||||||
|
if t.After(cutoff) {
|
||||||
|
kept = append(kept, t)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
e.fails = append(kept, now)
|
||||||
|
if len(e.fails) >= pol.maxFails {
|
||||||
|
e.lockedUntil = now.Add(pol.lockFor)
|
||||||
|
e.fails = nil
|
||||||
|
return true, pol.lockFor
|
||||||
|
}
|
||||||
|
return false, 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// ClearEndpoint 清除某端编号相关的登录锁定(两种都清)。
|
||||||
|
func (l *MemoryLocks) ClearEndpoint(endpointID string) {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
for k, e := range l.entries {
|
||||||
|
// kind|endpoint|peer|ip
|
||||||
|
parts := splitLockKey(k)
|
||||||
|
if len(parts) != 4 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
kind, ep := LockKind(parts[0]), parts[1]
|
||||||
|
if ep != endpointID {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if kind == LockLoginEndpointIP || kind == LockLoginEndpoint {
|
||||||
|
delete(l.entries, k)
|
||||||
|
_ = e
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Clear 清除精确键。
|
||||||
|
func (l *MemoryLocks) Clear(key LockKey) {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
delete(l.entries, lockMapKey(key))
|
||||||
|
}
|
||||||
|
|
||||||
|
func splitLockKey(k string) []string {
|
||||||
|
out := make([]string, 0, 4)
|
||||||
|
start := 0
|
||||||
|
for i := 0; i < len(k); i++ {
|
||||||
|
if k[i] == '|' {
|
||||||
|
out = append(out, k[start:i])
|
||||||
|
start = i + 1
|
||||||
|
}
|
||||||
|
}
|
||||||
|
out = append(out, k[start:])
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
var _ LoginLocks = (*MemoryLocks)(nil)
|
||||||
@@ -0,0 +1,90 @@
|
|||||||
|
package auth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/subtle"
|
||||||
|
"encoding/base64"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"golang.org/x/crypto/argon2"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Argon2 参数(DEVELOPMENT 第 12 节 / OWASP 最低配置)。
|
||||||
|
const (
|
||||||
|
ArgonMemoryKiB = 19 * 1024 // 19 MiB
|
||||||
|
ArgonTime = 2
|
||||||
|
ArgonThreads = 1
|
||||||
|
ArgonKeyLen = 32
|
||||||
|
ArgonSaltLen = 16
|
||||||
|
)
|
||||||
|
|
||||||
|
// ErrInvalidPHC 表示 PHC 字符串无法解析或参数不支持。
|
||||||
|
var ErrInvalidPHC = errors.New("auth: invalid phc")
|
||||||
|
|
||||||
|
// HashPassword 用 argon2id 生成 PHC 字符串(不含并发池;池见 HashPool)。
|
||||||
|
func HashPassword(password string) (string, error) {
|
||||||
|
salt := make([]byte, ArgonSaltLen)
|
||||||
|
if _, err := rand.Read(salt); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
hash := argon2.IDKey([]byte(password), salt, ArgonTime, ArgonMemoryKiB, ArgonThreads, ArgonKeyLen)
|
||||||
|
return encodePHC(salt, hash), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// VerifyPassword 常量时间比较密码与 PHC;不匹配时 ok=false 且 err=nil。
|
||||||
|
func VerifyPassword(password, phc string) (bool, error) {
|
||||||
|
salt, hash, err := decodePHC(phc)
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
got := argon2.IDKey([]byte(password), salt, ArgonTime, ArgonMemoryKiB, ArgonThreads, ArgonKeyLen)
|
||||||
|
if subtle.ConstantTimeCompare(got, hash) == 1 {
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func encodePHC(salt, hash []byte) string {
|
||||||
|
return fmt.Sprintf(
|
||||||
|
"$argon2id$v=%d$m=%d,t=%d,p=%d$%s$%s",
|
||||||
|
argon2.Version,
|
||||||
|
ArgonMemoryKiB,
|
||||||
|
ArgonTime,
|
||||||
|
ArgonThreads,
|
||||||
|
base64.RawStdEncoding.EncodeToString(salt),
|
||||||
|
base64.RawStdEncoding.EncodeToString(hash),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
func decodePHC(phc string) (salt, hash []byte, err error) {
|
||||||
|
// $argon2id$v=19$m=19456,t=2,p=1$salt$hash
|
||||||
|
parts := strings.Split(phc, "$")
|
||||||
|
if len(parts) != 6 || parts[1] != "argon2id" {
|
||||||
|
return nil, nil, ErrInvalidPHC
|
||||||
|
}
|
||||||
|
var version int
|
||||||
|
if _, scanErr := fmt.Sscanf(parts[2], "v=%d", &version); scanErr != nil || version != argon2.Version {
|
||||||
|
return nil, nil, ErrInvalidPHC
|
||||||
|
}
|
||||||
|
var m, t, p int
|
||||||
|
if _, scanErr := fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &m, &t, &p); scanErr != nil {
|
||||||
|
return nil, nil, ErrInvalidPHC
|
||||||
|
}
|
||||||
|
if m != ArgonMemoryKiB || t != ArgonTime || p != ArgonThreads {
|
||||||
|
return nil, nil, ErrInvalidPHC
|
||||||
|
}
|
||||||
|
salt, err = base64.RawStdEncoding.DecodeString(parts[4])
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, ErrInvalidPHC
|
||||||
|
}
|
||||||
|
hash, err = base64.RawStdEncoding.DecodeString(parts[5])
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, ErrInvalidPHC
|
||||||
|
}
|
||||||
|
if len(salt) == 0 || len(hash) == 0 {
|
||||||
|
return nil, nil, ErrInvalidPHC
|
||||||
|
}
|
||||||
|
return salt, hash, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,82 @@
|
|||||||
|
package auth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"runtime"
|
||||||
|
"sync/atomic"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Pool 是 argon2id 并发池(DEVELOPMENT 第 12 节)。
|
||||||
|
type Pool struct {
|
||||||
|
sem chan struct{}
|
||||||
|
waiting atomic.Int64
|
||||||
|
active atomic.Int64
|
||||||
|
maxSeen atomic.Int64 // 测试用:观察到的最大并发
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewPool 创建大小为 CPU 核数的哈希池。
|
||||||
|
func NewPool() *Pool {
|
||||||
|
return NewPoolSize(runtime.NumCPU())
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewPoolSize 创建指定并发上限的哈希池(测试可传入 1)。
|
||||||
|
func NewPoolSize(n int) *Pool {
|
||||||
|
if n < 1 {
|
||||||
|
n = 1
|
||||||
|
}
|
||||||
|
return &Pool{sem: make(chan struct{}, n)}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Hash 在池内计算 PHC 格式哈希。
|
||||||
|
func (p *Pool) Hash(ctx context.Context, _ PasswordKind, password string) (string, error) {
|
||||||
|
if err := p.acquire(ctx); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
defer p.release()
|
||||||
|
return HashPassword(password)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify 在池内校验;常量时间比较。
|
||||||
|
func (p *Pool) Verify(ctx context.Context, _ PasswordKind, password, phc string) (bool, error) {
|
||||||
|
if err := p.acquire(ctx); err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
defer p.release()
|
||||||
|
return VerifyPassword(password, phc)
|
||||||
|
}
|
||||||
|
|
||||||
|
// QueueLen 返回等待获取池槽位的任务数。
|
||||||
|
func (p *Pool) QueueLen() int {
|
||||||
|
return int(p.waiting.Load())
|
||||||
|
}
|
||||||
|
|
||||||
|
// MaxActive 返回曾达到的最大并发哈希数(测试用)。
|
||||||
|
func (p *Pool) MaxActive() int {
|
||||||
|
return int(p.maxSeen.Load())
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Pool) acquire(ctx context.Context) error {
|
||||||
|
p.waiting.Add(1)
|
||||||
|
select {
|
||||||
|
case p.sem <- struct{}{}:
|
||||||
|
p.waiting.Add(-1)
|
||||||
|
cur := p.active.Add(1)
|
||||||
|
for {
|
||||||
|
old := p.maxSeen.Load()
|
||||||
|
if cur <= old || p.maxSeen.CompareAndSwap(old, cur) {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
case <-ctx.Done():
|
||||||
|
p.waiting.Add(-1)
|
||||||
|
return ctx.Err()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Pool) release() {
|
||||||
|
p.active.Add(-1)
|
||||||
|
<-p.sem
|
||||||
|
}
|
||||||
|
|
||||||
|
var _ HashPool = (*Pool)(nil)
|
||||||
@@ -0,0 +1,87 @@
|
|||||||
|
package auth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/sha256"
|
||||||
|
"crypto/subtle"
|
||||||
|
"encoding/base64"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
sessionTokenPrefix = "nst_"
|
||||||
|
apiTokenPrefix = "nxm_"
|
||||||
|
tokenRandomBytes = 32
|
||||||
|
)
|
||||||
|
|
||||||
|
// SessionTokenService 生成与哈希端会话令牌。
|
||||||
|
type SessionTokenService struct{}
|
||||||
|
|
||||||
|
// NewSessionTokens 创建会话令牌服务。
|
||||||
|
func NewSessionTokens() *SessionTokenService {
|
||||||
|
return &SessionTokenService{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Issue 生成 nst_ + 32 字节随机数的 base64url;返回明文与 SHA-256。
|
||||||
|
func (s *SessionTokenService) Issue(_ context.Context) (string, []byte, error) {
|
||||||
|
return issuePrefixedToken(sessionTokenPrefix)
|
||||||
|
}
|
||||||
|
|
||||||
|
// HashToken 对令牌做 SHA-256。
|
||||||
|
func (s *SessionTokenService) HashToken(token string) []byte {
|
||||||
|
sum := sha256.Sum256([]byte(token))
|
||||||
|
return sum[:]
|
||||||
|
}
|
||||||
|
|
||||||
|
// LooksLikeSessionToken 判断是否以 nst_ 开头。
|
||||||
|
func (s *SessionTokenService) LooksLikeSessionToken(credential string) bool {
|
||||||
|
return strings.HasPrefix(credential, sessionTokenPrefix)
|
||||||
|
}
|
||||||
|
|
||||||
|
// APITokenService 生成与哈希 API 令牌。
|
||||||
|
type APITokenService struct{}
|
||||||
|
|
||||||
|
// NewAPITokens 创建 API 令牌服务。
|
||||||
|
func NewAPITokens() *APITokenService {
|
||||||
|
return &APITokenService{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Issue 生成 nxm_ + 32 字节随机数的 base64url;返回明文与 SHA-256。
|
||||||
|
func (s *APITokenService) Issue(_ context.Context) (string, []byte, error) {
|
||||||
|
return issuePrefixedToken(apiTokenPrefix)
|
||||||
|
}
|
||||||
|
|
||||||
|
// HashToken 对令牌做 SHA-256。
|
||||||
|
func (s *APITokenService) HashToken(token string) []byte {
|
||||||
|
sum := sha256.Sum256([]byte(token))
|
||||||
|
return sum[:]
|
||||||
|
}
|
||||||
|
|
||||||
|
// LooksLikeAPIToken 判断是否以 nxm_ 开头。
|
||||||
|
func (s *APITokenService) LooksLikeAPIToken(credential string) bool {
|
||||||
|
return strings.HasPrefix(credential, apiTokenPrefix)
|
||||||
|
}
|
||||||
|
|
||||||
|
// EqualHash 常量时间比较两个哈希。
|
||||||
|
func EqualHash(a, b []byte) bool {
|
||||||
|
if len(a) != len(b) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return subtle.ConstantTimeCompare(a, b) == 1
|
||||||
|
}
|
||||||
|
|
||||||
|
func issuePrefixedToken(prefix string) (string, []byte, error) {
|
||||||
|
raw := make([]byte, tokenRandomBytes)
|
||||||
|
if _, err := rand.Read(raw); err != nil {
|
||||||
|
return "", nil, err
|
||||||
|
}
|
||||||
|
tok := prefix + base64.RawURLEncoding.EncodeToString(raw)
|
||||||
|
sum := sha256.Sum256([]byte(tok))
|
||||||
|
return tok, sum[:], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
_ SessionTokens = (*SessionTokenService)(nil)
|
||||||
|
_ APITokens = (*APITokenService)(nil)
|
||||||
|
)
|
||||||
@@ -0,0 +1,379 @@
|
|||||||
|
package broker
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/rand"
|
||||||
|
"encoding/hex"
|
||||||
|
"errors"
|
||||||
|
"log/slog"
|
||||||
|
"net"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||||
|
mqtt "github.com/mochi-mqtt/server/v2"
|
||||||
|
"github.com/mochi-mqtt/server/v2/packets"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
maxClients = 2000
|
||||||
|
maxPacketSize = 786432
|
||||||
|
uplinkQueueSize = 256
|
||||||
|
largeFrameBytes = 64 * 1024
|
||||||
|
largeFrameSlots = 64
|
||||||
|
packetOverheadBudget = 128 // 主题与 MQTT 包头预留
|
||||||
|
keepaliveMin = 10
|
||||||
|
keepaliveMax = 600
|
||||||
|
)
|
||||||
|
|
||||||
|
// ErrPayloadTooLarge 下行超过客户端 Maximum Packet Size(减包头预留)或 max_receive_bytes。
|
||||||
|
var ErrPayloadTooLarge = errors.New("broker: payload exceeds client limit")
|
||||||
|
|
||||||
|
// ErrNoConnection 目标端没有当前连接。
|
||||||
|
var ErrNoConnection = errors.New("broker: no active connection")
|
||||||
|
|
||||||
|
// AuthResult 是登录校验结论(N3 实现真实逻辑;N2 默认拒绝)。
|
||||||
|
type AuthResult struct {
|
||||||
|
OK bool
|
||||||
|
SessionToken string // 密码登录成功时由 N3 填写
|
||||||
|
}
|
||||||
|
|
||||||
|
// Authenticator 由 N3 实现;内部故障必须返回 error,不得当成密码错误。
|
||||||
|
type Authenticator interface {
|
||||||
|
Authenticate(ctx context.Context, endpointID string, password []byte, remoteIP string) (AuthResult, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// RejectAuthenticator 默认拒绝所有客户端(CONNACK 用户名密码错误)。
|
||||||
|
type RejectAuthenticator struct{}
|
||||||
|
|
||||||
|
func (RejectAuthenticator) Authenticate(context.Context, string, []byte, string) (AuthResult, error) {
|
||||||
|
return AuthResult{OK: false}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// AllowAuthenticator 测试用:允许任意编号。
|
||||||
|
type AllowAuthenticator struct{}
|
||||||
|
|
||||||
|
func (AllowAuthenticator) Authenticate(context.Context, string, []byte, string) (AuthResult, error) {
|
||||||
|
return AuthResult{OK: true}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Options 装配 broker。
|
||||||
|
type Options struct {
|
||||||
|
Authenticator Authenticator
|
||||||
|
Uplink port.UplinkHandler
|
||||||
|
Logger *slog.Logger
|
||||||
|
}
|
||||||
|
|
||||||
|
// Broker 内置 mochi,不自带监听端口。
|
||||||
|
type Broker struct {
|
||||||
|
server *mqtt.Server
|
||||||
|
auth Authenticator
|
||||||
|
uplink port.UplinkHandler
|
||||||
|
log *slog.Logger
|
||||||
|
|
||||||
|
hook *nixHook
|
||||||
|
|
||||||
|
connsMu sync.RWMutex
|
||||||
|
current map[string]*connState
|
||||||
|
byClient map[*mqtt.Client]*connState
|
||||||
|
|
||||||
|
queuesMu sync.Mutex
|
||||||
|
queues map[string]*uplinkQueue
|
||||||
|
|
||||||
|
largeSem chan struct{}
|
||||||
|
closed atomic.Bool
|
||||||
|
}
|
||||||
|
|
||||||
|
type connState struct {
|
||||||
|
connID port.ConnID
|
||||||
|
endpointID string
|
||||||
|
transport port.Transport
|
||||||
|
remoteIP string
|
||||||
|
client *mqtt.Client
|
||||||
|
maxPacketSize uint32
|
||||||
|
maxRecvBytes int
|
||||||
|
authOK bool
|
||||||
|
authErr error
|
||||||
|
sessionToken string
|
||||||
|
largeHeld int
|
||||||
|
mu sync.Mutex
|
||||||
|
}
|
||||||
|
|
||||||
|
// New 创建并 Serve mochi(无监听器)。
|
||||||
|
func New(opts Options) (*Broker, error) {
|
||||||
|
auth := opts.Authenticator
|
||||||
|
if auth == nil {
|
||||||
|
auth = RejectAuthenticator{}
|
||||||
|
}
|
||||||
|
uplink := opts.Uplink
|
||||||
|
if uplink == nil {
|
||||||
|
uplink = port.StubUplinkHandler{}
|
||||||
|
}
|
||||||
|
log := opts.Logger
|
||||||
|
if log == nil {
|
||||||
|
log = slog.Default()
|
||||||
|
}
|
||||||
|
|
||||||
|
caps := mqtt.NewDefaultServerCapabilities()
|
||||||
|
caps.MaximumClients = maxClients
|
||||||
|
caps.MaximumQos = 1
|
||||||
|
caps.MaximumPacketSize = maxPacketSize
|
||||||
|
caps.MaximumSessionExpiryInterval = 0
|
||||||
|
caps.ReceiveMaximum = 1024
|
||||||
|
caps.MaximumInflight = 1024
|
||||||
|
caps.MaximumClientWritesPending = 1024
|
||||||
|
caps.RetainAvailable = 0
|
||||||
|
caps.WildcardSubAvailable = 0
|
||||||
|
caps.SharedSubAvailable = 0
|
||||||
|
caps.TopicAliasMaximum = 0
|
||||||
|
caps.Compatibilities.ObscureNotAuthorized = true
|
||||||
|
|
||||||
|
srv := mqtt.New(&mqtt.Options{
|
||||||
|
InlineClient: true,
|
||||||
|
Capabilities: caps,
|
||||||
|
Logger: log,
|
||||||
|
})
|
||||||
|
|
||||||
|
b := &Broker{
|
||||||
|
server: srv,
|
||||||
|
auth: auth,
|
||||||
|
uplink: uplink,
|
||||||
|
log: log,
|
||||||
|
current: make(map[string]*connState),
|
||||||
|
byClient: make(map[*mqtt.Client]*connState),
|
||||||
|
queues: make(map[string]*uplinkQueue),
|
||||||
|
largeSem: make(chan struct{}, largeFrameSlots),
|
||||||
|
}
|
||||||
|
b.hook = &nixHook{b: b}
|
||||||
|
if err := srv.AddHook(b.hook, nil); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if err := srv.Serve(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return b, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Server 返回底层 mochi(测试用)。
|
||||||
|
func (b *Broker) Server() *mqtt.Server { return b.server }
|
||||||
|
|
||||||
|
// Close 关闭 broker。
|
||||||
|
func (b *Broker) Close() error {
|
||||||
|
if b.closed.Swap(true) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
b.queuesMu.Lock()
|
||||||
|
for _, q := range b.queues {
|
||||||
|
q.close()
|
||||||
|
}
|
||||||
|
b.queuesMu.Unlock()
|
||||||
|
return b.server.Close()
|
||||||
|
}
|
||||||
|
|
||||||
|
// AttachTCP 把裸 TCP/TLS 连接交给 mochi;阻塞到连接结束。
|
||||||
|
func (b *Broker) AttachTCP(conn net.Conn) error {
|
||||||
|
return b.server.EstablishConnection("tcp", conn)
|
||||||
|
}
|
||||||
|
|
||||||
|
// AttachWS 把 WebSocket NetConn 交给 mochi;阻塞到连接结束。
|
||||||
|
func (b *Broker) AttachWS(conn net.Conn) error {
|
||||||
|
return b.server.EstablishConnection("ws", conn)
|
||||||
|
}
|
||||||
|
|
||||||
|
// PublishDown 实现 port.Downlink。
|
||||||
|
func (b *Broker) PublishDown(ctx context.Context, endpointID string, connID port.ConnID, payload []byte, opts port.PublishOpts) error {
|
||||||
|
if b.closed.Load() {
|
||||||
|
return errors.New("broker: closed")
|
||||||
|
}
|
||||||
|
st := b.lookupConn(endpointID, connID)
|
||||||
|
if st == nil {
|
||||||
|
return ErrNoConnection
|
||||||
|
}
|
||||||
|
|
||||||
|
limit := effectivePayloadLimit(st.maxPacketSize, st.maxRecvBytes)
|
||||||
|
if limit > 0 && len(payload) > limit {
|
||||||
|
return ErrPayloadTooLarge
|
||||||
|
}
|
||||||
|
|
||||||
|
qos := opts.QoS
|
||||||
|
if qos > 1 {
|
||||||
|
qos = 1
|
||||||
|
}
|
||||||
|
topic := downTopic(endpointID)
|
||||||
|
large := len(payload) > largeFrameBytes
|
||||||
|
|
||||||
|
if large {
|
||||||
|
select {
|
||||||
|
case b.largeSem <- struct{}{}:
|
||||||
|
case <-ctx.Done():
|
||||||
|
return ctx.Err()
|
||||||
|
}
|
||||||
|
st.mu.Lock()
|
||||||
|
st.largeHeld++
|
||||||
|
st.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := b.server.Publish(topic, payload, false, qos); err != nil {
|
||||||
|
if large {
|
||||||
|
b.releaseOneLarge(st)
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if large && qos == 0 {
|
||||||
|
b.releaseOneLarge(st)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Broker) releaseOneLarge(st *connState) {
|
||||||
|
st.mu.Lock()
|
||||||
|
if st.largeHeld > 0 {
|
||||||
|
st.largeHeld--
|
||||||
|
st.mu.Unlock()
|
||||||
|
select {
|
||||||
|
case <-b.largeSem:
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
st.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Broker) releaseAllLarge(st *connState) {
|
||||||
|
st.mu.Lock()
|
||||||
|
n := st.largeHeld
|
||||||
|
st.largeHeld = 0
|
||||||
|
st.mu.Unlock()
|
||||||
|
for i := 0; i < n; i++ {
|
||||||
|
select {
|
||||||
|
case <-b.largeSem:
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Disconnect 实现 port.ConnControl。
|
||||||
|
func (b *Broker) Disconnect(_ context.Context, endpointID string, connID port.ConnID, reason port.DisconnectReason) error {
|
||||||
|
st := b.lookupConn(endpointID, connID)
|
||||||
|
if st == nil {
|
||||||
|
return ErrNoConnection
|
||||||
|
}
|
||||||
|
code := packets.CodeDisconnect
|
||||||
|
switch reason {
|
||||||
|
case port.DisconnectTakenOver:
|
||||||
|
code = packets.ErrSessionTakenOver
|
||||||
|
case port.DisconnectKicked, port.DisconnectFatal:
|
||||||
|
code = packets.ErrAdministrativeAction
|
||||||
|
}
|
||||||
|
return b.server.DisconnectClient(st.client, code)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Broker) lookupConn(endpointID string, connID port.ConnID) *connState {
|
||||||
|
b.connsMu.RLock()
|
||||||
|
defer b.connsMu.RUnlock()
|
||||||
|
if connID != "" {
|
||||||
|
for _, st := range b.byClient {
|
||||||
|
if st.endpointID == endpointID && st.connID == connID {
|
||||||
|
return st
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return b.current[endpointID]
|
||||||
|
}
|
||||||
|
|
||||||
|
func downTopic(endpointID string) string {
|
||||||
|
return "nix/c/" + endpointID + "/down"
|
||||||
|
}
|
||||||
|
|
||||||
|
func upTopic(endpointID string) string {
|
||||||
|
return "nix/c/" + endpointID + "/up"
|
||||||
|
}
|
||||||
|
|
||||||
|
func effectivePayloadLimit(maxPacketSize uint32, maxRecvBytes int) int {
|
||||||
|
limit := 0
|
||||||
|
if maxPacketSize > 0 {
|
||||||
|
if maxPacketSize > packetOverheadBudget {
|
||||||
|
limit = int(maxPacketSize) - packetOverheadBudget
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if maxRecvBytes > 0 {
|
||||||
|
if limit == 0 || maxRecvBytes < limit {
|
||||||
|
limit = maxRecvBytes
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return limit
|
||||||
|
}
|
||||||
|
|
||||||
|
func randomConnID() port.ConnID {
|
||||||
|
var b [16]byte
|
||||||
|
_, _ = rand.Read(b[:])
|
||||||
|
return port.ConnID(hex.EncodeToString(b[:]))
|
||||||
|
}
|
||||||
|
|
||||||
|
func transportOf(cl *mqtt.Client) port.Transport {
|
||||||
|
if cl != nil && cl.Net.Listener == "ws" {
|
||||||
|
return port.TransportWS
|
||||||
|
}
|
||||||
|
return port.TransportTCP
|
||||||
|
}
|
||||||
|
|
||||||
|
func remoteIPOf(cl *mqtt.Client) string {
|
||||||
|
if cl == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
addr := cl.Net.Remote
|
||||||
|
if addr == "" && cl.Net.Conn != nil && cl.Net.Conn.RemoteAddr() != nil {
|
||||||
|
addr = cl.Net.Conn.RemoteAddr().String()
|
||||||
|
}
|
||||||
|
host, _, err := net.SplitHostPort(addr)
|
||||||
|
if err != nil {
|
||||||
|
return addr
|
||||||
|
}
|
||||||
|
return host
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetMaxReceiveBytes 供 N3 握手后设置;0 表示不限。
|
||||||
|
func (b *Broker) SetMaxReceiveBytes(endpointID string, connID port.ConnID, n int) {
|
||||||
|
st := b.lookupConn(endpointID, connID)
|
||||||
|
if st == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
st.mu.Lock()
|
||||||
|
st.maxRecvBytes = n
|
||||||
|
st.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// ConnInfoOf 返回连接信息(测试/N3)。
|
||||||
|
func (b *Broker) ConnInfoOf(endpointID string) (port.ConnInfo, bool) {
|
||||||
|
b.connsMu.RLock()
|
||||||
|
st := b.current[endpointID]
|
||||||
|
b.connsMu.RUnlock()
|
||||||
|
if st == nil {
|
||||||
|
return port.ConnInfo{}, false
|
||||||
|
}
|
||||||
|
return port.ConnInfo{
|
||||||
|
ConnID: st.connID,
|
||||||
|
EndpointID: st.endpointID,
|
||||||
|
Transport: st.transport,
|
||||||
|
RemoteIP: st.remoteIP,
|
||||||
|
SessionToken: st.sessionToken,
|
||||||
|
MaxPacketSize: st.maxPacketSize,
|
||||||
|
}, true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *Broker) enqueueUplink(endpointID string, conn port.ConnInfo, payload []byte) {
|
||||||
|
b.queuesMu.Lock()
|
||||||
|
q, ok := b.queues[endpointID]
|
||||||
|
if !ok {
|
||||||
|
q = newUplinkQueue(b, endpointID)
|
||||||
|
b.queues[endpointID] = q
|
||||||
|
}
|
||||||
|
b.queuesMu.Unlock()
|
||||||
|
q.push(uplinkItem{conn: conn, payload: payload})
|
||||||
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
_ port.Downlink = (*Broker)(nil)
|
||||||
|
_ port.ConnControl = (*Broker)(nil)
|
||||||
|
)
|
||||||
@@ -0,0 +1,259 @@
|
|||||||
|
package broker
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||||
|
"github.com/coder/websocket"
|
||||||
|
"github.com/mochi-mqtt/server/v2/packets"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestWSCrossOriginAllowed(t *testing.T) {
|
||||||
|
b, err := New(Options{Authenticator: AllowAuthenticator{}})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer func() { _ = b.Close() }()
|
||||||
|
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
mux.Handle("/mqtt", b.WSHandler(nil))
|
||||||
|
srv := httptest.NewServer(mux)
|
||||||
|
defer srv.Close()
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
c, _, err := websocket.Dial(ctx, "ws"+srv.URL[len("http"):]+"/mqtt", &websocket.DialOptions{
|
||||||
|
HTTPHeader: http.Header{"Origin": []string{"https://other.example"}},
|
||||||
|
Subprotocols: []string{"mqtt"},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("cross-origin dial: %v", err)
|
||||||
|
}
|
||||||
|
defer func() { _ = c.Close(websocket.StatusNormalClosure, "") }()
|
||||||
|
if c.Subprotocol() != "mqtt" {
|
||||||
|
t.Fatalf("subprotocol=%q", c.Subprotocol())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWSWrongSubprotocolClosed(t *testing.T) {
|
||||||
|
b, err := New(Options{Authenticator: AllowAuthenticator{}})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer func() { _ = b.Close() }()
|
||||||
|
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
mux.Handle("/mqtt", b.WSHandler(nil))
|
||||||
|
srv := httptest.NewServer(mux)
|
||||||
|
defer srv.Close()
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
c, _, err := websocket.Dial(ctx, "ws"+srv.URL[len("http"):]+"/mqtt", &websocket.DialOptions{
|
||||||
|
HTTPHeader: http.Header{"Origin": []string{"https://other.example"}},
|
||||||
|
Subprotocols: []string{"not-mqtt"},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
// 有的实现在握手阶段就失败;也算关闭
|
||||||
|
return
|
||||||
|
}
|
||||||
|
defer func() { _ = c.Close(websocket.StatusNormalClosure, "") }()
|
||||||
|
|
||||||
|
// 服务端应立刻关掉;后续读写会失败
|
||||||
|
c.SetReadLimit(16)
|
||||||
|
_, _, readErr := c.Read(ctx)
|
||||||
|
if readErr == nil {
|
||||||
|
t.Fatal("expected connection closed for wrong subprotocol")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPublishDownExceedsClientMax(t *testing.T) {
|
||||||
|
b, err := New(Options{Authenticator: AllowAuthenticator{}})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer func() { _ = b.Close() }()
|
||||||
|
|
||||||
|
clientDone := make(chan struct{})
|
||||||
|
r, w := net.Pipe()
|
||||||
|
go func() {
|
||||||
|
defer close(clientDone)
|
||||||
|
_ = b.AttachTCP(r)
|
||||||
|
}()
|
||||||
|
|
||||||
|
endpoint := "ep-limit"
|
||||||
|
connectAndSubscribe(t, w, endpoint, 200) // MaximumPacketSize=200 → payload limit 72
|
||||||
|
|
||||||
|
// 等会话建立
|
||||||
|
deadline := time.Now().Add(3 * time.Second)
|
||||||
|
for {
|
||||||
|
if _, ok := b.ConnInfoOf(endpoint); ok {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if time.Now().After(deadline) {
|
||||||
|
t.Fatal("session not established")
|
||||||
|
}
|
||||||
|
time.Sleep(10 * time.Millisecond)
|
||||||
|
}
|
||||||
|
|
||||||
|
big := bytes.Repeat([]byte("x"), 100) // > 200-128
|
||||||
|
pubErr := b.PublishDown(context.Background(), endpoint, "", big, port.PublishOpts{QoS: 1})
|
||||||
|
if !errors.Is(pubErr, ErrPayloadTooLarge) {
|
||||||
|
t.Fatalf("PublishDown err=%v want ErrPayloadTooLarge", pubErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 合法大小应成功
|
||||||
|
small := []byte(`{"v":1,"type":"resp"}`)
|
||||||
|
if err := b.PublishDown(context.Background(), endpoint, "", small, port.PublishOpts{QoS: 0}); err != nil {
|
||||||
|
t.Fatalf("small publish: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
_ = w.Close()
|
||||||
|
select {
|
||||||
|
case <-clientDone:
|
||||||
|
case <-time.After(3 * time.Second):
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInternalAuthErrorDoesNotReturnBadPassword(t *testing.T) {
|
||||||
|
auth := &errAuthenticator{err: context.DeadlineExceeded}
|
||||||
|
b, err := New(Options{Authenticator: auth})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer func() { _ = b.Close() }()
|
||||||
|
|
||||||
|
r, w := net.Pipe()
|
||||||
|
errCh := make(chan error, 1)
|
||||||
|
go func() { errCh <- b.AttachTCP(r) }()
|
||||||
|
|
||||||
|
writeConnect(t, w, "ep-err", 30, 0)
|
||||||
|
// 不应收到 CONNACK(内部错误直接断开)
|
||||||
|
_ = w.SetReadDeadline(time.Now().Add(500 * time.Millisecond))
|
||||||
|
buf := make([]byte, 64)
|
||||||
|
n, readErr := w.Read(buf)
|
||||||
|
if readErr == nil && n > 0 {
|
||||||
|
// 若收到包,不能是 bad username/password CONNACK (reason 0x86)
|
||||||
|
if n >= 2 && buf[0]>>4 == packets.Connack {
|
||||||
|
t.Fatalf("unexpected connack on internal error: %x", buf[:n])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
_ = w.Close()
|
||||||
|
select {
|
||||||
|
case <-errCh:
|
||||||
|
case <-time.After(2 * time.Second):
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRejectUnknownByDefault(t *testing.T) {
|
||||||
|
b, err := New(Options{}) // RejectAuthenticator
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer func() { _ = b.Close() }()
|
||||||
|
|
||||||
|
r, w := net.Pipe()
|
||||||
|
go func() { _ = b.AttachTCP(r) }()
|
||||||
|
writeConnect(t, w, "ep-unknown", 30, 0)
|
||||||
|
_ = w.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||||
|
buf := make([]byte, 128)
|
||||||
|
n, err := io.ReadAtLeast(w, buf, 2)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if buf[0]>>4 != packets.Connack {
|
||||||
|
t.Fatalf("want connack, got %x", buf[:n])
|
||||||
|
}
|
||||||
|
_ = w.Close()
|
||||||
|
}
|
||||||
|
|
||||||
|
type errAuthenticator struct {
|
||||||
|
err error
|
||||||
|
mu sync.Mutex
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *errAuthenticator) Authenticate(context.Context, string, []byte, string) (AuthResult, error) {
|
||||||
|
a.mu.Lock()
|
||||||
|
defer a.mu.Unlock()
|
||||||
|
return AuthResult{}, a.err
|
||||||
|
}
|
||||||
|
|
||||||
|
func connectAndSubscribe(t *testing.T, w net.Conn, endpoint string, maxPacket uint32) {
|
||||||
|
t.Helper()
|
||||||
|
writeConnect(t, w, endpoint, 30, maxPacket)
|
||||||
|
// read CONNACK
|
||||||
|
_ = w.SetReadDeadline(time.Now().Add(3 * time.Second))
|
||||||
|
buf := make([]byte, 256)
|
||||||
|
n, err := io.ReadAtLeast(w, buf, 2)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if buf[0]>>4 != packets.Connack {
|
||||||
|
t.Fatalf("want connack got %x", buf[:n])
|
||||||
|
}
|
||||||
|
writeSubscribe(t, w, downTopic(endpoint))
|
||||||
|
// read SUBACK
|
||||||
|
n, err = io.ReadAtLeast(w, buf, 2)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if buf[0]>>4 != packets.Suback {
|
||||||
|
t.Fatalf("want suback got %x", buf[:n])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeConnect(t *testing.T, w net.Conn, endpoint string, keepalive uint16, maxPacket uint32) {
|
||||||
|
t.Helper()
|
||||||
|
pk := packets.Packet{
|
||||||
|
FixedHeader: packets.FixedHeader{Type: packets.Connect},
|
||||||
|
ProtocolVersion: 5,
|
||||||
|
Connect: packets.ConnectParams{
|
||||||
|
ProtocolName: []byte("MQTT"),
|
||||||
|
Clean: true,
|
||||||
|
ClientIdentifier: endpoint,
|
||||||
|
Keepalive: keepalive,
|
||||||
|
UsernameFlag: true,
|
||||||
|
Username: []byte(endpoint),
|
||||||
|
PasswordFlag: true,
|
||||||
|
Password: []byte("test"),
|
||||||
|
},
|
||||||
|
Properties: packets.Properties{
|
||||||
|
MaximumPacketSize: maxPacket,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
var buf bytes.Buffer
|
||||||
|
if err := pk.ConnectEncode(&buf); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := w.Write(buf.Bytes()); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeSubscribe(t *testing.T, w net.Conn, topic string) {
|
||||||
|
t.Helper()
|
||||||
|
pk := packets.Packet{
|
||||||
|
FixedHeader: packets.FixedHeader{Type: packets.Subscribe, Qos: 1},
|
||||||
|
ProtocolVersion: 5,
|
||||||
|
PacketID: 1,
|
||||||
|
Filters: packets.Subscriptions{
|
||||||
|
{Filter: topic, Qos: 1},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
var buf bytes.Buffer
|
||||||
|
if err := pk.SubscribeEncode(&buf); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := w.Write(buf.Bytes()); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,196 @@
|
|||||||
|
package broker
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||||
|
mqtt "github.com/mochi-mqtt/server/v2"
|
||||||
|
"github.com/mochi-mqtt/server/v2/packets"
|
||||||
|
)
|
||||||
|
|
||||||
|
type nixHook struct {
|
||||||
|
mqtt.HookBase
|
||||||
|
b *Broker
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *nixHook) ID() string { return "nixmsg" }
|
||||||
|
|
||||||
|
func (h *nixHook) Provides(b byte) bool {
|
||||||
|
return bytes.Contains([]byte{
|
||||||
|
mqtt.OnConnect,
|
||||||
|
mqtt.OnConnectAuthenticate,
|
||||||
|
mqtt.OnACLCheck,
|
||||||
|
mqtt.OnPublish,
|
||||||
|
mqtt.OnPublishDropped,
|
||||||
|
mqtt.OnSessionEstablished,
|
||||||
|
mqtt.OnDisconnect,
|
||||||
|
mqtt.OnQosComplete,
|
||||||
|
}, []byte{b})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *nixHook) OnConnect(cl *mqtt.Client, pk packets.Packet) error {
|
||||||
|
endpointID := string(pk.Connect.Username)
|
||||||
|
if endpointID == "" {
|
||||||
|
endpointID = pk.Connect.ClientIdentifier
|
||||||
|
}
|
||||||
|
remoteIP := remoteIPOf(cl)
|
||||||
|
|
||||||
|
st := &connState{
|
||||||
|
connID: randomConnID(),
|
||||||
|
endpointID: endpointID,
|
||||||
|
transport: transportOf(cl),
|
||||||
|
remoteIP: remoteIP,
|
||||||
|
client: cl,
|
||||||
|
maxPacketSize: pk.Properties.MaximumPacketSize,
|
||||||
|
}
|
||||||
|
|
||||||
|
// 心跳校正:超出 10–600 秒就改写 Keepalive 并设 ServerKeepalive
|
||||||
|
ka := pk.Connect.Keepalive
|
||||||
|
if ka < keepaliveMin || ka > keepaliveMax {
|
||||||
|
if ka < keepaliveMin {
|
||||||
|
ka = keepaliveMin
|
||||||
|
}
|
||||||
|
if ka > keepaliveMax {
|
||||||
|
ka = keepaliveMax
|
||||||
|
}
|
||||||
|
cl.State.Keepalive = ka
|
||||||
|
cl.State.ServerKeepalive = true
|
||||||
|
}
|
||||||
|
|
||||||
|
res, err := h.b.auth.Authenticate(context.Background(), endpointID, pk.Connect.Password, remoteIP)
|
||||||
|
if err != nil {
|
||||||
|
st.authErr = err
|
||||||
|
h.rememberPending(cl, st)
|
||||||
|
return err // mochi 不回 CONNACK,直接断开
|
||||||
|
}
|
||||||
|
st.authOK = res.OK
|
||||||
|
st.sessionToken = res.SessionToken
|
||||||
|
h.rememberPending(cl, st)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *nixHook) rememberPending(cl *mqtt.Client, st *connState) {
|
||||||
|
h.b.connsMu.Lock()
|
||||||
|
h.b.byClient[cl] = st
|
||||||
|
h.b.connsMu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *nixHook) OnConnectAuthenticate(cl *mqtt.Client, _ packets.Packet) bool {
|
||||||
|
h.b.connsMu.RLock()
|
||||||
|
st := h.b.byClient[cl]
|
||||||
|
h.b.connsMu.RUnlock()
|
||||||
|
if st == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
// 内部故障已在 OnConnect 返回 error;此处只反映业务上的拒绝
|
||||||
|
return st.authOK
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *nixHook) OnACLCheck(cl *mqtt.Client, topic string, write bool) bool {
|
||||||
|
h.b.connsMu.RLock()
|
||||||
|
st := h.b.byClient[cl]
|
||||||
|
h.b.connsMu.RUnlock()
|
||||||
|
if st == nil || st.endpointID == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
up := upTopic(st.endpointID)
|
||||||
|
down := downTopic(st.endpointID)
|
||||||
|
if write {
|
||||||
|
return topic == up
|
||||||
|
}
|
||||||
|
return topic == down
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *nixHook) OnPublish(cl *mqtt.Client, pk packets.Packet) (packets.Packet, error) {
|
||||||
|
h.b.connsMu.RLock()
|
||||||
|
st := h.b.byClient[cl]
|
||||||
|
h.b.connsMu.RUnlock()
|
||||||
|
if st == nil {
|
||||||
|
return pk, packets.CodeSuccessIgnore
|
||||||
|
}
|
||||||
|
payload := append([]byte(nil), pk.Payload...)
|
||||||
|
info := port.ConnInfo{
|
||||||
|
ConnID: st.connID,
|
||||||
|
EndpointID: st.endpointID,
|
||||||
|
Transport: st.transport,
|
||||||
|
RemoteIP: st.remoteIP,
|
||||||
|
SessionToken: st.sessionToken,
|
||||||
|
MaxPacketSize: st.maxPacketSize,
|
||||||
|
}
|
||||||
|
h.b.enqueueUplink(st.endpointID, info, payload)
|
||||||
|
return pk, packets.CodeSuccessIgnore
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *nixHook) OnPublishDropped(cl *mqtt.Client, pk packets.Packet) {
|
||||||
|
h.b.log.Debug("publish dropped", "client", cl.ID, "topic", pk.TopicName, "size", len(pk.Payload))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *nixHook) OnSessionEstablished(cl *mqtt.Client, _ packets.Packet) {
|
||||||
|
h.b.connsMu.Lock()
|
||||||
|
st := h.b.byClient[cl]
|
||||||
|
if st != nil {
|
||||||
|
h.b.current[st.endpointID] = st
|
||||||
|
}
|
||||||
|
h.b.connsMu.Unlock()
|
||||||
|
if st == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
info := port.ConnInfo{
|
||||||
|
ConnID: st.connID,
|
||||||
|
EndpointID: st.endpointID,
|
||||||
|
Transport: st.transport,
|
||||||
|
RemoteIP: st.remoteIP,
|
||||||
|
SessionToken: st.sessionToken,
|
||||||
|
MaxPacketSize: st.maxPacketSize,
|
||||||
|
}
|
||||||
|
_ = h.b.uplink.OnSessionEstablished(context.Background(), info)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *nixHook) OnDisconnect(cl *mqtt.Client, err error, _ bool) {
|
||||||
|
h.b.connsMu.Lock()
|
||||||
|
st := h.b.byClient[cl]
|
||||||
|
delete(h.b.byClient, cl)
|
||||||
|
if st != nil && h.b.current[st.endpointID] == st {
|
||||||
|
delete(h.b.current, st.endpointID)
|
||||||
|
}
|
||||||
|
h.b.connsMu.Unlock()
|
||||||
|
if st == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
h.b.releaseAllLarge(st)
|
||||||
|
|
||||||
|
reason := port.DisconnectNormal
|
||||||
|
if err != nil {
|
||||||
|
if code, ok := err.(packets.Code); ok {
|
||||||
|
switch code.Code {
|
||||||
|
case packets.ErrSessionTakenOver.Code:
|
||||||
|
reason = port.DisconnectTakenOver
|
||||||
|
case packets.ErrAdministrativeAction.Code:
|
||||||
|
reason = port.DisconnectKicked
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
info := port.ConnInfo{
|
||||||
|
ConnID: st.connID,
|
||||||
|
EndpointID: st.endpointID,
|
||||||
|
Transport: st.transport,
|
||||||
|
RemoteIP: st.remoteIP,
|
||||||
|
SessionToken: st.sessionToken,
|
||||||
|
MaxPacketSize: st.maxPacketSize,
|
||||||
|
}
|
||||||
|
h.b.uplink.OnDisconnect(context.Background(), info, reason)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *nixHook) OnQosComplete(cl *mqtt.Client, pk packets.Packet) {
|
||||||
|
if len(pk.Payload) <= largeFrameBytes {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
h.b.connsMu.RLock()
|
||||||
|
st := h.b.byClient[cl]
|
||||||
|
h.b.connsMu.RUnlock()
|
||||||
|
if st == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
h.b.releaseOneLarge(st)
|
||||||
|
}
|
||||||
@@ -0,0 +1,45 @@
|
|||||||
|
package broker
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||||
|
)
|
||||||
|
|
||||||
|
type uplinkItem struct {
|
||||||
|
conn port.ConnInfo
|
||||||
|
payload []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// uplinkQueue 每端串行队列,长度 256,满了堵住 OnPublish(背压)。
|
||||||
|
type uplinkQueue struct {
|
||||||
|
b *Broker
|
||||||
|
endpointID string
|
||||||
|
ch chan uplinkItem
|
||||||
|
once sync.Once
|
||||||
|
}
|
||||||
|
|
||||||
|
func newUplinkQueue(b *Broker, endpointID string) *uplinkQueue {
|
||||||
|
q := &uplinkQueue{
|
||||||
|
b: b,
|
||||||
|
endpointID: endpointID,
|
||||||
|
ch: make(chan uplinkItem, uplinkQueueSize),
|
||||||
|
}
|
||||||
|
go q.loop()
|
||||||
|
return q
|
||||||
|
}
|
||||||
|
|
||||||
|
func (q *uplinkQueue) push(item uplinkItem) {
|
||||||
|
q.ch <- item // 满则阻塞读循环,形成背压
|
||||||
|
}
|
||||||
|
|
||||||
|
func (q *uplinkQueue) close() {
|
||||||
|
q.once.Do(func() { close(q.ch) })
|
||||||
|
}
|
||||||
|
|
||||||
|
func (q *uplinkQueue) loop() {
|
||||||
|
for item := range q.ch {
|
||||||
|
_ = q.b.uplink.HandleUplink(context.Background(), item.conn, item.payload)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,41 @@
|
|||||||
|
package broker
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
|
||||||
|
"git.asio.asia/nixevol/NixMsg/internal/listener"
|
||||||
|
"github.com/coder/websocket"
|
||||||
|
)
|
||||||
|
|
||||||
|
// WSHandler 返回 /mqtt 的 WebSocket 升级处理。
|
||||||
|
// Accept 时 InsecureSkipVerify=true;之后检查 Subprotocol==mqtt。
|
||||||
|
// NetConn 使用 Background 派生的 context,不用请求 Context。
|
||||||
|
func (b *Broker) WSHandler(proxies *listener.ProxySet) http.Handler {
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
c, err := websocket.Accept(w, r, &websocket.AcceptOptions{
|
||||||
|
Subprotocols: []string{"mqtt"},
|
||||||
|
InsecureSkipVerify: true,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if c.Subprotocol() != "mqtt" {
|
||||||
|
_ = c.Close(websocket.StatusPolicyViolation, "subprotocol must be mqtt")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
nc := websocket.NetConn(ctx, c, websocket.MessageBinary)
|
||||||
|
if proxies != nil {
|
||||||
|
ip := proxies.ClientIP(r)
|
||||||
|
if ip != "" {
|
||||||
|
nc = listener.WithRemoteAddr(nc, &net.TCPAddr{IP: net.ParseIP(ip)})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
_ = b.AttachWS(nc)
|
||||||
|
})
|
||||||
|
}
|
||||||
+142
-5
@@ -2,11 +2,16 @@ package config
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"net"
|
||||||
"os"
|
"os"
|
||||||
|
"strings"
|
||||||
|
|
||||||
"go.yaml.in/yaml/v3"
|
"go.yaml.in/yaml/v3"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// MaxBodyBytesCap 是 max_body_bytes 的硬上限(DEVELOPMENT 11.1)。
|
||||||
|
const MaxBodyBytesCap = 262144
|
||||||
|
|
||||||
// Config 对应 DEVELOPMENT 第 11.1 节的配置文件。
|
// Config 对应 DEVELOPMENT 第 11.1 节的配置文件。
|
||||||
type Config struct {
|
type Config struct {
|
||||||
Listen string `yaml:"listen"`
|
Listen string `yaml:"listen"`
|
||||||
@@ -65,7 +70,7 @@ func Default() Config {
|
|||||||
TrustedProxies: nil,
|
TrustedProxies: nil,
|
||||||
DataDir: "./data",
|
DataDir: "./data",
|
||||||
Limits: LimitsConfig{
|
Limits: LimitsConfig{
|
||||||
MaxBodyBytes: 262144,
|
MaxBodyBytes: MaxBodyBytesCap,
|
||||||
MaxMetaBytes: 4096,
|
MaxMetaBytes: 4096,
|
||||||
MaxFrameBytes: 786432,
|
MaxFrameBytes: 786432,
|
||||||
MaxTTLSeconds: 2592000,
|
MaxTTLSeconds: 2592000,
|
||||||
@@ -101,16 +106,148 @@ func Load(path string) (Config, error) {
|
|||||||
if err := yaml.Unmarshal(data, &cfg); err != nil {
|
if err := yaml.Unmarshal(data, &cfg); err != nil {
|
||||||
return Config{}, fmt.Errorf("parse config: %w", err)
|
return Config{}, fmt.Errorf("parse config: %w", err)
|
||||||
}
|
}
|
||||||
|
applyEmptyDefaults(&cfg)
|
||||||
|
return cfg, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func applyEmptyDefaults(cfg *Config) {
|
||||||
|
def := Default()
|
||||||
if cfg.Listen == "" {
|
if cfg.Listen == "" {
|
||||||
cfg.Listen = Default().Listen
|
cfg.Listen = def.Listen
|
||||||
}
|
}
|
||||||
if cfg.DataDir == "" {
|
if cfg.DataDir == "" {
|
||||||
cfg.DataDir = Default().DataDir
|
cfg.DataDir = def.DataDir
|
||||||
}
|
}
|
||||||
if cfg.SQLiteSynchronous == "" {
|
if cfg.SQLiteSynchronous == "" {
|
||||||
cfg.SQLiteSynchronous = Default().SQLiteSynchronous
|
cfg.SQLiteSynchronous = def.SQLiteSynchronous
|
||||||
}
|
}
|
||||||
return cfg, nil
|
if cfg.Log.Level == "" {
|
||||||
|
cfg.Log.Level = def.Log.Level
|
||||||
|
}
|
||||||
|
if cfg.Limits.MaxBodyBytes == 0 {
|
||||||
|
cfg.Limits.MaxBodyBytes = def.Limits.MaxBodyBytes
|
||||||
|
}
|
||||||
|
if cfg.Limits.MaxMetaBytes == 0 {
|
||||||
|
cfg.Limits.MaxMetaBytes = def.Limits.MaxMetaBytes
|
||||||
|
}
|
||||||
|
if cfg.Limits.MaxFrameBytes == 0 {
|
||||||
|
cfg.Limits.MaxFrameBytes = def.Limits.MaxFrameBytes
|
||||||
|
}
|
||||||
|
if cfg.Limits.MaxTTLSeconds == 0 {
|
||||||
|
cfg.Limits.MaxTTLSeconds = def.Limits.MaxTTLSeconds
|
||||||
|
}
|
||||||
|
if cfg.Limits.MaxScheduleSeconds == 0 {
|
||||||
|
cfg.Limits.MaxScheduleSeconds = def.Limits.MaxScheduleSeconds
|
||||||
|
}
|
||||||
|
if cfg.Limits.MaxGroupMembers == 0 {
|
||||||
|
cfg.Limits.MaxGroupMembers = def.Limits.MaxGroupMembers
|
||||||
|
}
|
||||||
|
if cfg.Limits.GraceSeconds == 0 {
|
||||||
|
cfg.Limits.GraceSeconds = def.Limits.GraceSeconds
|
||||||
|
}
|
||||||
|
if cfg.Limits.AckTimeoutSeconds == 0 {
|
||||||
|
cfg.Limits.AckTimeoutSeconds = def.Limits.AckTimeoutSeconds
|
||||||
|
}
|
||||||
|
if cfg.Limits.DeliveryWindow == 0 {
|
||||||
|
cfg.Limits.DeliveryWindow = def.Limits.DeliveryWindow
|
||||||
|
}
|
||||||
|
if cfg.Limits.ReceiptWindow == 0 {
|
||||||
|
cfg.Limits.ReceiptWindow = def.Limits.ReceiptWindow
|
||||||
|
}
|
||||||
|
if cfg.Limits.RequestsPerSecond == 0 {
|
||||||
|
cfg.Limits.RequestsPerSecond = def.Limits.RequestsPerSecond
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate 校验 DEVELOPMENT 11.1 全部字段;拒绝 max_body_bytes > 262144。
|
||||||
|
func (c Config) Validate() error {
|
||||||
|
var errs []string
|
||||||
|
if strings.TrimSpace(c.Listen) == "" {
|
||||||
|
errs = append(errs, "listen is required")
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(c.DataDir) == "" {
|
||||||
|
errs = append(errs, "data_dir is required")
|
||||||
|
}
|
||||||
|
sync := strings.ToUpper(strings.TrimSpace(c.SQLiteSynchronous))
|
||||||
|
if sync != "FULL" && sync != "NORMAL" {
|
||||||
|
errs = append(errs, "sqlite_synchronous must be FULL or NORMAL")
|
||||||
|
}
|
||||||
|
switch strings.ToLower(strings.TrimSpace(c.Log.Level)) {
|
||||||
|
case "debug", "info", "warn", "warning", "error":
|
||||||
|
default:
|
||||||
|
errs = append(errs, "log.level must be debug, info, warn, or error")
|
||||||
|
}
|
||||||
|
if c.Limits.MaxBodyBytes <= 0 {
|
||||||
|
errs = append(errs, "limits.max_body_bytes must be > 0")
|
||||||
|
}
|
||||||
|
if c.Limits.MaxBodyBytes > MaxBodyBytesCap {
|
||||||
|
errs = append(errs, fmt.Sprintf("limits.max_body_bytes must be <= %d", MaxBodyBytesCap))
|
||||||
|
}
|
||||||
|
if c.Limits.MaxMetaBytes <= 0 {
|
||||||
|
errs = append(errs, "limits.max_meta_bytes must be > 0")
|
||||||
|
}
|
||||||
|
if c.Limits.MaxFrameBytes <= 0 {
|
||||||
|
errs = append(errs, "limits.max_frame_bytes must be > 0")
|
||||||
|
}
|
||||||
|
if c.Limits.MaxFrameBytes < c.Limits.MaxBodyBytes {
|
||||||
|
errs = append(errs, "limits.max_frame_bytes must be >= limits.max_body_bytes")
|
||||||
|
}
|
||||||
|
if c.Limits.MaxTTLSeconds <= 0 {
|
||||||
|
errs = append(errs, "limits.max_ttl_seconds must be > 0")
|
||||||
|
}
|
||||||
|
if c.Limits.MaxScheduleSeconds <= 0 {
|
||||||
|
errs = append(errs, "limits.max_schedule_seconds must be > 0")
|
||||||
|
}
|
||||||
|
if c.Limits.MaxGroupMembers <= 0 {
|
||||||
|
errs = append(errs, "limits.max_group_members must be > 0")
|
||||||
|
}
|
||||||
|
if c.Limits.GraceSeconds < 0 {
|
||||||
|
errs = append(errs, "limits.grace_seconds must be >= 0")
|
||||||
|
}
|
||||||
|
if c.Limits.AckTimeoutSeconds <= 0 {
|
||||||
|
errs = append(errs, "limits.ack_timeout_seconds must be > 0")
|
||||||
|
}
|
||||||
|
if c.Limits.DeliveryWindow <= 0 {
|
||||||
|
errs = append(errs, "limits.delivery_window must be > 0")
|
||||||
|
}
|
||||||
|
if c.Limits.ReceiptWindow <= 0 {
|
||||||
|
errs = append(errs, "limits.receipt_window must be > 0")
|
||||||
|
}
|
||||||
|
if c.Limits.RequestsPerSecond <= 0 {
|
||||||
|
errs = append(errs, "limits.requests_per_second must be > 0")
|
||||||
|
}
|
||||||
|
if c.Limits.MaxPendingPerSender < 0 {
|
||||||
|
errs = append(errs, "limits.max_pending_per_sender must be >= 0")
|
||||||
|
}
|
||||||
|
if c.Limits.MaxPendingPerReceiver < 0 {
|
||||||
|
errs = append(errs, "limits.max_pending_per_receiver must be >= 0")
|
||||||
|
}
|
||||||
|
if c.SessionIdleDays < 0 {
|
||||||
|
errs = append(errs, "session_idle_days must be >= 0")
|
||||||
|
}
|
||||||
|
if c.RecordRetentionDays < 0 {
|
||||||
|
errs = append(errs, "record_retention_days must be >= 0")
|
||||||
|
}
|
||||||
|
if c.IdempotencyHours < 0 {
|
||||||
|
errs = append(errs, "idempotency_hours must be >= 0")
|
||||||
|
}
|
||||||
|
if c.ReceiptRetentionDays < 0 {
|
||||||
|
errs = append(errs, "receipt_retention_days must be >= 0")
|
||||||
|
}
|
||||||
|
cert := strings.TrimSpace(c.TLS.CertFile)
|
||||||
|
key := strings.TrimSpace(c.TLS.KeyFile)
|
||||||
|
if (cert == "") != (key == "") {
|
||||||
|
errs = append(errs, "tls.cert_file and tls.key_file must both be set or both empty")
|
||||||
|
}
|
||||||
|
for _, cidr := range c.TrustedProxies {
|
||||||
|
if _, _, err := net.ParseCIDR(strings.TrimSpace(cidr)); err != nil {
|
||||||
|
errs = append(errs, fmt.Sprintf("trusted_proxies entry %q is not a valid CIDR", cidr))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(errs) > 0 {
|
||||||
|
return fmt.Errorf("invalid config: %s", strings.Join(errs, "; "))
|
||||||
|
}
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// PathFromEnv 返回 NIXMSG_CONFIG 或默认 ./config.yaml。
|
// PathFromEnv 返回 NIXMSG_CONFIG 或默认 ./config.yaml。
|
||||||
|
|||||||
@@ -0,0 +1,62 @@
|
|||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestValidateDefaultOK(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
cfg := Default()
|
||||||
|
if err := cfg.Validate(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestValidateRejectsMaxBodyBytesTooLarge(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
cfg := Default()
|
||||||
|
cfg.Limits.MaxBodyBytes = MaxBodyBytesCap + 1
|
||||||
|
err := cfg.Validate()
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected error")
|
||||||
|
}
|
||||||
|
if !strings.Contains(err.Error(), "max_body_bytes") {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadAndValidate(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
dir := t.TempDir()
|
||||||
|
path := filepath.Join(dir, "config.yaml")
|
||||||
|
body := "listen: \"127.0.0.1:0\"\ndata_dir: \"" + filepath.ToSlash(dir) + "\"\n"
|
||||||
|
if err := os.WriteFile(path, []byte(body), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
cfg, err := Load(path)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := cfg.Validate(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if cfg.Listen != "127.0.0.1:0" {
|
||||||
|
t.Fatalf("listen=%q", cfg.Listen)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestValidateTrustedProxies(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
cfg := Default()
|
||||||
|
cfg.TrustedProxies = []string{"not-a-cidr"}
|
||||||
|
if err := cfg.Validate(); err == nil {
|
||||||
|
t.Fatal("expected error")
|
||||||
|
}
|
||||||
|
cfg.TrustedProxies = []string{"127.0.0.1/32"}
|
||||||
|
if err := cfg.Validate(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,62 @@
|
|||||||
|
package listener
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"sync"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ChanListener 是从通道取连接的 net.Listener,交给同一个 http.Server。
|
||||||
|
type ChanListener struct {
|
||||||
|
addr net.Addr
|
||||||
|
ch chan net.Conn
|
||||||
|
closed chan struct{}
|
||||||
|
once sync.Once
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewChanListener 创建缓冲通道监听器;addr 仅用于 Addr()。
|
||||||
|
func NewChanListener(addr net.Addr, buf int) *ChanListener {
|
||||||
|
if buf < 1 {
|
||||||
|
buf = 64
|
||||||
|
}
|
||||||
|
if addr == nil {
|
||||||
|
addr = &net.TCPAddr{IP: net.IPv4zero, Port: 0}
|
||||||
|
}
|
||||||
|
return &ChanListener{
|
||||||
|
addr: addr,
|
||||||
|
ch: make(chan net.Conn, buf),
|
||||||
|
closed: make(chan struct{}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Addr 返回构造时给出的地址。
|
||||||
|
func (l *ChanListener) Addr() net.Addr { return l.addr }
|
||||||
|
|
||||||
|
// Accept 阻塞直到有连接或关闭。
|
||||||
|
func (l *ChanListener) Accept() (net.Conn, error) {
|
||||||
|
select {
|
||||||
|
case <-l.closed:
|
||||||
|
return nil, net.ErrClosed
|
||||||
|
case c, ok := <-l.ch:
|
||||||
|
if !ok {
|
||||||
|
return nil, net.ErrClosed
|
||||||
|
}
|
||||||
|
return c, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close 关闭监听器并唤醒 Accept。
|
||||||
|
func (l *ChanListener) Close() error {
|
||||||
|
l.once.Do(func() {
|
||||||
|
close(l.closed)
|
||||||
|
})
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Enqueue 把识别为 HTTP 的连接交给 http.Server;已关闭时丢弃并关闭连接。
|
||||||
|
func (l *ChanListener) Enqueue(c net.Conn) {
|
||||||
|
select {
|
||||||
|
case <-l.closed:
|
||||||
|
_ = c.Close()
|
||||||
|
case l.ch <- c:
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,76 @@
|
|||||||
|
package listener
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bufio"
|
||||||
|
"net"
|
||||||
|
"strconv"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
const firstByteTimeout = 10 * time.Second
|
||||||
|
|
||||||
|
// bufferedConn 把已读字节放回连接,供后续 TLS/HTTP/MQTT 继续读。
|
||||||
|
type bufferedConn struct {
|
||||||
|
net.Conn
|
||||||
|
r *bufio.Reader
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *bufferedConn) Read(p []byte) (int, error) {
|
||||||
|
return c.r.Read(p)
|
||||||
|
}
|
||||||
|
|
||||||
|
func wrapBuffered(c net.Conn) *bufferedConn {
|
||||||
|
if bc, ok := c.(*bufferedConn); ok {
|
||||||
|
return bc
|
||||||
|
}
|
||||||
|
return &bufferedConn{Conn: c, r: bufio.NewReader(c)}
|
||||||
|
}
|
||||||
|
|
||||||
|
// peekFirstByte 在超时内读首字节并 Unread,返回仍可读完整流的连接。
|
||||||
|
func peekFirstByte(c net.Conn) (net.Conn, byte, error) {
|
||||||
|
bc := wrapBuffered(c)
|
||||||
|
_ = bc.SetReadDeadline(time.Now().Add(firstByteTimeout))
|
||||||
|
b, err := bc.r.ReadByte()
|
||||||
|
_ = bc.SetReadDeadline(time.Time{})
|
||||||
|
if err != nil {
|
||||||
|
return nil, 0, err
|
||||||
|
}
|
||||||
|
if err := bc.r.UnreadByte(); err != nil {
|
||||||
|
return nil, 0, err
|
||||||
|
}
|
||||||
|
return bc, b, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// addrConn 只改 RemoteAddr,用于受信任代理后的真实 IP。
|
||||||
|
type addrConn struct {
|
||||||
|
net.Conn
|
||||||
|
remote net.Addr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *addrConn) RemoteAddr() net.Addr {
|
||||||
|
if c.remote != nil {
|
||||||
|
return c.remote
|
||||||
|
}
|
||||||
|
return c.Conn.RemoteAddr()
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithRemoteAddr 包装连接,使 RemoteAddr 返回指定地址(通常是解析出的客户端 IP)。
|
||||||
|
func WithRemoteAddr(c net.Conn, remote net.Addr) net.Conn {
|
||||||
|
if remote == nil {
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
return &addrConn{Conn: c, remote: remote}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TCPAddrFromIPPort 把 "ip:port" 或纯 IP 转成 *net.TCPAddr。
|
||||||
|
func TCPAddrFromIPPort(ipPort string) *net.TCPAddr {
|
||||||
|
if ipPort == "" {
|
||||||
|
return &net.TCPAddr{}
|
||||||
|
}
|
||||||
|
host, portStr, err := net.SplitHostPort(ipPort)
|
||||||
|
if err != nil {
|
||||||
|
return &net.TCPAddr{IP: net.ParseIP(ipPort)}
|
||||||
|
}
|
||||||
|
port, _ := strconv.Atoi(portStr)
|
||||||
|
return &net.TCPAddr{IP: net.ParseIP(host), Port: port}
|
||||||
|
}
|
||||||
@@ -0,0 +1,398 @@
|
|||||||
|
package listener
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/ecdsa"
|
||||||
|
"crypto/elliptic"
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/tls"
|
||||||
|
"crypto/x509"
|
||||||
|
"crypto/x509/pkix"
|
||||||
|
"encoding/pem"
|
||||||
|
"io"
|
||||||
|
"math/big"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestIdentifyPlainHTTPAndMQTT(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
gotMQTT := make(chan net.Conn, 1)
|
||||||
|
mux := NewMux(RoleShared, Handlers{
|
||||||
|
Healthz: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
_, _ = w.Write([]byte("ok"))
|
||||||
|
}),
|
||||||
|
})
|
||||||
|
s, err := New(Options{
|
||||||
|
Listen: "127.0.0.1:0",
|
||||||
|
DataDir: dir,
|
||||||
|
ClientHandler: mux,
|
||||||
|
AllowPlaintext: true,
|
||||||
|
OnMQTT: func(c net.Conn) {
|
||||||
|
gotMQTT <- c
|
||||||
|
buf := make([]byte, 1)
|
||||||
|
_, _ = c.Read(buf)
|
||||||
|
_ = c.Close()
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
if startErr := s.Start(ctx); startErr != nil {
|
||||||
|
t.Fatal(startErr)
|
||||||
|
}
|
||||||
|
defer func() { _ = s.Close() }()
|
||||||
|
|
||||||
|
addr := s.ListenAddr()
|
||||||
|
if addr == "" {
|
||||||
|
t.Fatal("empty listen addr")
|
||||||
|
}
|
||||||
|
b, err := os.ReadFile(filepath.Join(dir, "listen.addr"))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if string(b) != addr+"\n" {
|
||||||
|
t.Fatalf("listen.addr=%q want %q", b, addr+"\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := http.Get("http://" + addr + "/healthz")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
_ = resp.Body.Close()
|
||||||
|
if resp.StatusCode != 200 || string(body) != "ok" {
|
||||||
|
t.Fatalf("healthz: %d %q", resp.StatusCode, body)
|
||||||
|
}
|
||||||
|
|
||||||
|
c, err := net.DialTimeout("tcp", addr, 2*time.Second)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
_, _ = c.Write([]byte{0x10, 0x00})
|
||||||
|
select {
|
||||||
|
case mc := <-gotMQTT:
|
||||||
|
_ = mc.Close()
|
||||||
|
case <-time.After(3 * time.Second):
|
||||||
|
t.Fatal("mqtt not delivered")
|
||||||
|
}
|
||||||
|
_ = c.Close()
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIdentifyTLSHTTPAndMQTT(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
certPath, keyPath := writeTestCert(t, dir, "old")
|
||||||
|
gotMQTT := make(chan net.Conn, 1)
|
||||||
|
mux := NewMux(RoleShared, Handlers{
|
||||||
|
Healthz: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||||
|
_, _ = w.Write([]byte("tls-ok"))
|
||||||
|
}),
|
||||||
|
})
|
||||||
|
s, err := New(Options{
|
||||||
|
Listen: "127.0.0.1:0",
|
||||||
|
DataDir: dir,
|
||||||
|
CertFile: certPath,
|
||||||
|
KeyFile: keyPath,
|
||||||
|
AllowPlaintext: false,
|
||||||
|
ClientHandler: mux,
|
||||||
|
OnMQTT: func(c net.Conn) {
|
||||||
|
gotMQTT <- c
|
||||||
|
_ = c.Close()
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
if startErr := s.Start(ctx); startErr != nil {
|
||||||
|
t.Fatal(startErr)
|
||||||
|
}
|
||||||
|
defer func() { _ = s.Close() }()
|
||||||
|
|
||||||
|
addr := s.ListenAddr()
|
||||||
|
tlsCfg := &tls.Config{InsecureSkipVerify: true}
|
||||||
|
|
||||||
|
// TLS + HTTP
|
||||||
|
tr := &http.Transport{TLSClientConfig: tlsCfg}
|
||||||
|
client := &http.Client{Transport: tr, Timeout: 5 * time.Second}
|
||||||
|
resp, err := client.Get("https://" + addr + "/healthz")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
_ = resp.Body.Close()
|
||||||
|
if string(body) != "tls-ok" {
|
||||||
|
t.Fatalf("body=%q", body)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TLS + MQTT (0x10 after handshake)
|
||||||
|
raw, err := tls.Dial("tcp", addr, tlsCfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
_, _ = raw.Write([]byte{0x10})
|
||||||
|
select {
|
||||||
|
case mc := <-gotMQTT:
|
||||||
|
_ = mc.Close()
|
||||||
|
case <-time.After(3 * time.Second):
|
||||||
|
t.Fatal("tls mqtt not delivered")
|
||||||
|
}
|
||||||
|
_ = raw.Close()
|
||||||
|
|
||||||
|
// ALPN mqtt 客户端仍能握手(服务端不设 NextProtos)
|
||||||
|
alpn, err := tls.Dial("tcp", addr, &tls.Config{
|
||||||
|
InsecureSkipVerify: true,
|
||||||
|
NextProtos: []string{"mqtt"},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("alpn mqtt handshake: %v", err)
|
||||||
|
}
|
||||||
|
_ = alpn.Close()
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCertReloadUsesNewCert(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
certPath, keyPath := writeTestCert(t, dir, "v1")
|
||||||
|
cr, err := NewCertReloader(certPath, keyPath, nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer cr.Close()
|
||||||
|
old := cr.Certificate()
|
||||||
|
if old == nil {
|
||||||
|
t.Fatal("nil cert")
|
||||||
|
}
|
||||||
|
|
||||||
|
time.Sleep(20 * time.Millisecond) // 保证 mtime 变化
|
||||||
|
certPath2, keyPath2 := writeTestCert(t, dir, "v2")
|
||||||
|
// 覆盖原路径
|
||||||
|
data, _ := os.ReadFile(certPath2)
|
||||||
|
_ = os.WriteFile(certPath, data, 0o644)
|
||||||
|
data, _ = os.ReadFile(keyPath2)
|
||||||
|
_ = os.WriteFile(keyPath, data, 0o644)
|
||||||
|
|
||||||
|
if reloadErr := cr.ReloadNow(); reloadErr != nil {
|
||||||
|
t.Fatal(reloadErr)
|
||||||
|
}
|
||||||
|
neu := cr.Certificate()
|
||||||
|
if neu == nil || neu == old {
|
||||||
|
t.Fatal("certificate not reloaded")
|
||||||
|
}
|
||||||
|
|
||||||
|
// 完整服务:重载后新连接用新证书(用 Leaf CN 区分)
|
||||||
|
mux := NewMux(RoleShared, Handlers{})
|
||||||
|
s, err := New(Options{
|
||||||
|
Listen: "127.0.0.1:0",
|
||||||
|
DataDir: dir,
|
||||||
|
CertFile: certPath,
|
||||||
|
KeyFile: keyPath,
|
||||||
|
ClientHandler: mux,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
if startErr := s.Start(ctx); startErr != nil {
|
||||||
|
t.Fatal(startErr)
|
||||||
|
}
|
||||||
|
defer func() { _ = s.Close() }()
|
||||||
|
|
||||||
|
// 再换一版证书
|
||||||
|
time.Sleep(20 * time.Millisecond)
|
||||||
|
c3, k3 := writeTestCert(t, dir, "v3")
|
||||||
|
data, _ = os.ReadFile(c3)
|
||||||
|
_ = os.WriteFile(certPath, data, 0o644)
|
||||||
|
data, _ = os.ReadFile(k3)
|
||||||
|
_ = os.WriteFile(keyPath, data, 0o644)
|
||||||
|
if reloadErr := s.certs.ReloadNow(); reloadErr != nil {
|
||||||
|
t.Fatal(reloadErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
conn, err := tls.Dial("tcp", s.ListenAddr(), &tls.Config{InsecureSkipVerify: true})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer func() { _ = conn.Close() }()
|
||||||
|
state := conn.ConnectionState()
|
||||||
|
if len(state.PeerCertificates) == 0 {
|
||||||
|
t.Fatal("no peer cert")
|
||||||
|
}
|
||||||
|
if cn := state.PeerCertificates[0].Subject.CommonName; cn != "v3" {
|
||||||
|
t.Fatalf("cn=%q want v3", cn)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAdminSeparateMQTTClosedAndAdmin404OnListen(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
clientMux := NewMux(RoleClient, Handlers{
|
||||||
|
Healthz: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||||
|
_, _ = w.Write([]byte("client"))
|
||||||
|
}),
|
||||||
|
})
|
||||||
|
adminMux := NewMux(RoleAdmin, Handlers{
|
||||||
|
AdminAPI: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||||
|
_, _ = w.Write([]byte("admin"))
|
||||||
|
}),
|
||||||
|
Healthz: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||||
|
_, _ = w.Write([]byte("admin-health"))
|
||||||
|
}),
|
||||||
|
})
|
||||||
|
mqttSeen := make(chan struct{}, 1)
|
||||||
|
s, err := New(Options{
|
||||||
|
Listen: "127.0.0.1:0",
|
||||||
|
AdminListen: "127.0.0.1:0",
|
||||||
|
DataDir: dir,
|
||||||
|
AllowPlaintext: true,
|
||||||
|
ClientHandler: clientMux,
|
||||||
|
AdminHandler: adminMux,
|
||||||
|
OnMQTT: func(c net.Conn) {
|
||||||
|
mqttSeen <- struct{}{}
|
||||||
|
_ = c.Close()
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
if startErr := s.Start(ctx); startErr != nil {
|
||||||
|
t.Fatal(startErr)
|
||||||
|
}
|
||||||
|
defer func() { _ = s.Close() }()
|
||||||
|
|
||||||
|
// listen 上 /api/admin/ 404
|
||||||
|
resp, err := http.Get("http://" + s.ListenAddr() + "/api/admin/x")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
_ = resp.Body.Close()
|
||||||
|
if resp.StatusCode != 404 {
|
||||||
|
t.Fatalf("listen admin status=%d", resp.StatusCode)
|
||||||
|
}
|
||||||
|
|
||||||
|
// admin 上 API 可用
|
||||||
|
resp, err = http.Get("http://" + s.AdminAddr() + "/api/admin/x")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
_ = resp.Body.Close()
|
||||||
|
if string(body) != "admin" {
|
||||||
|
t.Fatalf("admin body=%q", body)
|
||||||
|
}
|
||||||
|
|
||||||
|
// admin_listen 上裸 MQTT 关闭,不回调
|
||||||
|
ac, err := net.DialTimeout("tcp", s.AdminAddr(), 2*time.Second)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
_, _ = ac.Write([]byte{0x10})
|
||||||
|
time.Sleep(200 * time.Millisecond)
|
||||||
|
buf := make([]byte, 1)
|
||||||
|
_ = ac.SetReadDeadline(time.Now().Add(500 * time.Millisecond))
|
||||||
|
_, readErr := ac.Read(buf)
|
||||||
|
_ = ac.Close()
|
||||||
|
if readErr == nil {
|
||||||
|
t.Fatal("expected admin mqtt connection closed")
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case <-mqttSeen:
|
||||||
|
t.Fatal("mqtt should not be accepted on admin")
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := os.ReadFile(filepath.Join(dir, "admin.addr")); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTrustedProxyClientIP(t *testing.T) {
|
||||||
|
ps, err := ParseTrustedProxies([]string{"10.0.0.0/8", "192.168.1.1"})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
r := &http.Request{
|
||||||
|
RemoteAddr: "10.1.2.3:1234",
|
||||||
|
Header: http.Header{"X-Forwarded-For": []string{"1.1.1.1, 10.9.9.9"}},
|
||||||
|
}
|
||||||
|
if got := ps.ClientIP(r); got != "1.1.1.1" {
|
||||||
|
t.Fatalf("got %q", got)
|
||||||
|
}
|
||||||
|
// 非代理来源忽略头
|
||||||
|
r2 := &http.Request{
|
||||||
|
RemoteAddr: "8.8.8.8:9",
|
||||||
|
Header: http.Header{"X-Forwarded-For": []string{"1.1.1.1"}},
|
||||||
|
}
|
||||||
|
if got := ps.ClientIP(r2); got != "8.8.8.8" {
|
||||||
|
t.Fatalf("got %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFirstByteTimeout(t *testing.T) {
|
||||||
|
c1, c2 := net.Pipe()
|
||||||
|
defer func() { _ = c1.Close() }()
|
||||||
|
defer func() { _ = c2.Close() }()
|
||||||
|
done := make(chan error, 1)
|
||||||
|
go func() {
|
||||||
|
_, _, err := peekFirstByte(c2)
|
||||||
|
done <- err
|
||||||
|
}()
|
||||||
|
select {
|
||||||
|
case err := <-done:
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected timeout error")
|
||||||
|
}
|
||||||
|
case <-time.After(12 * time.Second):
|
||||||
|
t.Fatal("peek did not time out")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeTestCert(t *testing.T, dir, cn string) (certPath, keyPath string) {
|
||||||
|
t.Helper()
|
||||||
|
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
tmpl := &x509.Certificate{
|
||||||
|
SerialNumber: big.NewInt(time.Now().UnixNano()),
|
||||||
|
Subject: pkix.Name{CommonName: cn},
|
||||||
|
NotBefore: time.Now().Add(-time.Hour),
|
||||||
|
NotAfter: time.Now().Add(24 * time.Hour),
|
||||||
|
KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment,
|
||||||
|
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
|
||||||
|
DNSNames: []string{"localhost"},
|
||||||
|
IPAddresses: []net.IP{net.ParseIP("127.0.0.1")},
|
||||||
|
}
|
||||||
|
der, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, &key.PublicKey, key)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
certPath = filepath.Join(dir, cn+".crt")
|
||||||
|
keyPath = filepath.Join(dir, cn+".key")
|
||||||
|
certOut, err := os.Create(certPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
_ = pem.Encode(certOut, &pem.Block{Type: "CERTIFICATE", Bytes: der})
|
||||||
|
_ = certOut.Close()
|
||||||
|
keyOut, err := os.Create(keyPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
b, err := x509.MarshalECPrivateKey(key)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
_ = pem.Encode(keyOut, &pem.Block{Type: "EC PRIVATE KEY", Bytes: b})
|
||||||
|
_ = keyOut.Close()
|
||||||
|
return certPath, keyPath
|
||||||
|
}
|
||||||
@@ -0,0 +1,101 @@
|
|||||||
|
package listener
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ProxySet 保存受信任代理地址段。
|
||||||
|
type ProxySet struct {
|
||||||
|
nets []*net.IPNet
|
||||||
|
}
|
||||||
|
|
||||||
|
// ParseTrustedProxies 解析 CIDR 或单 IP 列表。
|
||||||
|
func ParseTrustedProxies(cidrs []string) (*ProxySet, error) {
|
||||||
|
ps := &ProxySet{}
|
||||||
|
for _, s := range cidrs {
|
||||||
|
s = strings.TrimSpace(s)
|
||||||
|
if s == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if !strings.Contains(s, "/") {
|
||||||
|
ip := net.ParseIP(s)
|
||||||
|
if ip == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if ip.To4() != nil {
|
||||||
|
s += "/32"
|
||||||
|
} else {
|
||||||
|
s += "/128"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
_, n, err := net.ParseCIDR(s)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
ps.nets = append(ps.nets, n)
|
||||||
|
}
|
||||||
|
return ps, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Contains 判断 IP 是否在受信任段内。
|
||||||
|
func (ps *ProxySet) Contains(ip net.IP) bool {
|
||||||
|
if ps == nil || ip == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, n := range ps.nets {
|
||||||
|
if n.Contains(ip) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// ClientIP 按 DEVELOPMENT 4.5:来自受信任代理时,取 X-Forwarded-For 从右往左第一个不在段内的 IP。
|
||||||
|
// 非代理来源忽略转发头,返回 RemoteAddr 的 IP。
|
||||||
|
func (ps *ProxySet) ClientIP(r *http.Request) string {
|
||||||
|
remoteIP := ipFromAddr(r.RemoteAddr)
|
||||||
|
if ps == nil || remoteIP == nil || !ps.Contains(remoteIP) {
|
||||||
|
if remoteIP != nil {
|
||||||
|
return remoteIP.String()
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
xff := r.Header.Get("X-Forwarded-For")
|
||||||
|
if xff == "" {
|
||||||
|
return remoteIP.String()
|
||||||
|
}
|
||||||
|
parts := strings.Split(xff, ",")
|
||||||
|
for i := len(parts) - 1; i >= 0; i-- {
|
||||||
|
ipStr := strings.TrimSpace(parts[i])
|
||||||
|
ip := net.ParseIP(ipStr)
|
||||||
|
if ip == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if !ps.Contains(ip) {
|
||||||
|
return ip.String()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return remoteIP.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsHTTPS 来自受信任代理时按 X-Forwarded-Proto 判断。
|
||||||
|
func (ps *ProxySet) IsHTTPS(r *http.Request) bool {
|
||||||
|
if r.TLS != nil {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
remoteIP := ipFromAddr(r.RemoteAddr)
|
||||||
|
if ps == nil || remoteIP == nil || !ps.Contains(remoteIP) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return strings.EqualFold(r.Header.Get("X-Forwarded-Proto"), "https")
|
||||||
|
}
|
||||||
|
|
||||||
|
func ipFromAddr(remoteAddr string) net.IP {
|
||||||
|
host, _, err := net.SplitHostPort(remoteAddr)
|
||||||
|
if err != nil {
|
||||||
|
return net.ParseIP(remoteAddr)
|
||||||
|
}
|
||||||
|
return net.ParseIP(host)
|
||||||
|
}
|
||||||
@@ -0,0 +1,124 @@
|
|||||||
|
package listener
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// RouteRole 区分监听用途,决定哪些路径可用。
|
||||||
|
type RouteRole int
|
||||||
|
|
||||||
|
const (
|
||||||
|
// RoleShared listen 与 admin 共用同一端口。
|
||||||
|
RoleShared RouteRole = iota
|
||||||
|
// RoleClient 仅端接入(admin_listen 已分离)。
|
||||||
|
RoleClient
|
||||||
|
// RoleAdmin 仅后台。
|
||||||
|
RoleAdmin
|
||||||
|
)
|
||||||
|
|
||||||
|
// Handlers 由上层注入各路径处理函数;未设置的路径返回 404。
|
||||||
|
type Handlers struct {
|
||||||
|
MQTT http.Handler // /mqtt
|
||||||
|
ClientAPI http.Handler // /api/client/
|
||||||
|
AdminAPI http.Handler // /api/admin/
|
||||||
|
Metrics http.Handler // /metrics
|
||||||
|
Static http.Handler // 后台静态页
|
||||||
|
Healthz http.Handler // /healthz
|
||||||
|
Readyz http.Handler // /readyz
|
||||||
|
// MetricsToken 共用端口时校验 Authorization: Bearer;空则 /metrics 404。
|
||||||
|
MetricsToken string
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewMux 按角色装配 HTTP 路由骨架。
|
||||||
|
func NewMux(role RouteRole, h Handlers) http.Handler {
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
|
||||||
|
healthz := h.Healthz
|
||||||
|
if healthz == nil {
|
||||||
|
healthz = http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
_, _ = w.Write([]byte("ok"))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
readyz := h.Readyz
|
||||||
|
if readyz == nil {
|
||||||
|
readyz = http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
_, _ = w.Write([]byte("ok"))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
mux.Handle("GET /healthz", healthz)
|
||||||
|
mux.Handle("GET /readyz", readyz)
|
||||||
|
|
||||||
|
switch role {
|
||||||
|
case RoleClient:
|
||||||
|
if h.MQTT != nil {
|
||||||
|
mux.Handle("/mqtt", h.MQTT)
|
||||||
|
}
|
||||||
|
if h.ClientAPI != nil {
|
||||||
|
mux.Handle("/api/client/", h.ClientAPI)
|
||||||
|
}
|
||||||
|
// 后台路径在端端口一律 404
|
||||||
|
mux.Handle("/api/admin/", http.NotFoundHandler())
|
||||||
|
mux.Handle("/metrics", http.NotFoundHandler())
|
||||||
|
mux.Handle("/", http.NotFoundHandler())
|
||||||
|
|
||||||
|
case RoleAdmin:
|
||||||
|
if h.AdminAPI != nil {
|
||||||
|
mux.Handle("/api/admin/", h.AdminAPI)
|
||||||
|
}
|
||||||
|
if h.Metrics != nil {
|
||||||
|
mux.Handle("GET /metrics", h.Metrics)
|
||||||
|
} else {
|
||||||
|
mux.Handle("GET /metrics", http.NotFoundHandler())
|
||||||
|
}
|
||||||
|
mux.Handle("/mqtt", http.NotFoundHandler())
|
||||||
|
mux.Handle("/api/client/", http.NotFoundHandler())
|
||||||
|
if h.Static != nil {
|
||||||
|
mux.Handle("/", h.Static)
|
||||||
|
} else {
|
||||||
|
mux.Handle("/", http.NotFoundHandler())
|
||||||
|
}
|
||||||
|
|
||||||
|
default: // RoleShared
|
||||||
|
if h.MQTT != nil {
|
||||||
|
mux.Handle("/mqtt", h.MQTT)
|
||||||
|
}
|
||||||
|
if h.ClientAPI != nil {
|
||||||
|
mux.Handle("/api/client/", h.ClientAPI)
|
||||||
|
}
|
||||||
|
if h.AdminAPI != nil {
|
||||||
|
mux.Handle("/api/admin/", h.AdminAPI)
|
||||||
|
}
|
||||||
|
mux.Handle("GET /metrics", metricsGate(h.MetricsToken, h.Metrics))
|
||||||
|
if h.Static != nil {
|
||||||
|
mux.Handle("/", h.Static)
|
||||||
|
} else {
|
||||||
|
mux.Handle("/", http.NotFoundHandler())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return mux
|
||||||
|
}
|
||||||
|
|
||||||
|
func metricsGate(token string, next http.Handler) http.Handler {
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if token == "" {
|
||||||
|
http.NotFound(w, r)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
auth := r.Header.Get("Authorization")
|
||||||
|
const prefix = "Bearer "
|
||||||
|
if !strings.HasPrefix(auth, prefix) || auth[len(prefix):] != token {
|
||||||
|
w.WriteHeader(http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if next == nil {
|
||||||
|
http.NotFound(w, r)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
next.ServeHTTP(w, r)
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -0,0 +1,338 @@
|
|||||||
|
package listener
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/tls"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Kind 识别结果。
|
||||||
|
type Kind int
|
||||||
|
|
||||||
|
const (
|
||||||
|
KindHTTP Kind = iota
|
||||||
|
KindMQTT
|
||||||
|
KindClosed
|
||||||
|
)
|
||||||
|
|
||||||
|
// Options 控制双端口识别与分流。
|
||||||
|
type Options struct {
|
||||||
|
Listen string
|
||||||
|
AdminListen string
|
||||||
|
DataDir string
|
||||||
|
CertFile string
|
||||||
|
KeyFile string
|
||||||
|
AllowPlaintext bool
|
||||||
|
TrustedProxies []string
|
||||||
|
// ClientHandler / AdminHandler 分别为端端口与后台端口的 HTTP 处理;AdminListen 为空时只用 ClientHandler。
|
||||||
|
ClientHandler http.Handler
|
||||||
|
AdminHandler http.Handler
|
||||||
|
// OnMQTT 在 listen 上识别到裸 MQTT(含 TLS 后)时调用;应阻塞到连接结束。
|
||||||
|
OnMQTT func(conn net.Conn)
|
||||||
|
Logger *slog.Logger
|
||||||
|
}
|
||||||
|
|
||||||
|
// Server 一个或两个 TCP 监听上的协议识别与分流。
|
||||||
|
type Server struct {
|
||||||
|
opts Options
|
||||||
|
log *slog.Logger
|
||||||
|
proxies *ProxySet
|
||||||
|
certs *CertReloader
|
||||||
|
|
||||||
|
clientHTTP *ChanListener
|
||||||
|
adminHTTP *ChanListener
|
||||||
|
clientSrv *http.Server
|
||||||
|
adminSrv *http.Server
|
||||||
|
|
||||||
|
clientLn net.Listener
|
||||||
|
adminLn net.Listener
|
||||||
|
|
||||||
|
listenAddr string
|
||||||
|
adminAddr string
|
||||||
|
writeListenAddr bool
|
||||||
|
writeAdminAddr bool
|
||||||
|
|
||||||
|
wg sync.WaitGroup
|
||||||
|
closed chan struct{}
|
||||||
|
closeOnce sync.Once
|
||||||
|
}
|
||||||
|
|
||||||
|
// New 校验选项并准备证书;不开始监听。
|
||||||
|
func New(opts Options) (*Server, error) {
|
||||||
|
if opts.Listen == "" {
|
||||||
|
return nil, errors.New("listener: listen is required")
|
||||||
|
}
|
||||||
|
if opts.ClientHandler == nil {
|
||||||
|
return nil, errors.New("listener: ClientHandler is required")
|
||||||
|
}
|
||||||
|
log := opts.Logger
|
||||||
|
if log == nil {
|
||||||
|
log = slog.Default()
|
||||||
|
}
|
||||||
|
ps, err := ParseTrustedProxies(opts.TrustedProxies)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("trusted_proxies: %w", err)
|
||||||
|
}
|
||||||
|
s := &Server{
|
||||||
|
opts: opts,
|
||||||
|
log: log,
|
||||||
|
proxies: ps,
|
||||||
|
closed: make(chan struct{}),
|
||||||
|
}
|
||||||
|
hasCert := opts.CertFile != "" && opts.KeyFile != ""
|
||||||
|
if hasCert {
|
||||||
|
cr, err := NewCertReloader(opts.CertFile, opts.KeyFile, log)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("tls: %w", err)
|
||||||
|
}
|
||||||
|
s.certs = cr
|
||||||
|
} else {
|
||||||
|
log.Warn("tls not configured; plaintext only")
|
||||||
|
}
|
||||||
|
s.writeListenAddr = isPortZero(opts.Listen)
|
||||||
|
s.writeAdminAddr = opts.AdminListen != "" && isPortZero(opts.AdminListen)
|
||||||
|
return s, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Start 开始监听并分流;非阻塞,关闭用 Close。
|
||||||
|
func (s *Server) Start(ctx context.Context) error {
|
||||||
|
ln, err := net.Listen("tcp", s.opts.Listen)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("listen %s: %w", s.opts.Listen, err)
|
||||||
|
}
|
||||||
|
s.clientLn = ln
|
||||||
|
s.listenAddr = ln.Addr().String()
|
||||||
|
if s.writeListenAddr {
|
||||||
|
if err := writeAddrFile(s.opts.DataDir, "listen.addr", s.listenAddr); err != nil {
|
||||||
|
_ = ln.Close()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
s.clientHTTP = NewChanListener(ln.Addr(), 128)
|
||||||
|
s.clientSrv = &http.Server{
|
||||||
|
Handler: s.opts.ClientHandler,
|
||||||
|
ReadHeaderTimeout: 10 * time.Second,
|
||||||
|
}
|
||||||
|
s.wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer s.wg.Done()
|
||||||
|
err := s.clientSrv.Serve(s.clientHTTP)
|
||||||
|
if err != nil && !errors.Is(err, http.ErrServerClosed) {
|
||||||
|
s.log.Error("client http serve", "err", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
s.wg.Add(1)
|
||||||
|
go s.acceptLoop(ln, false)
|
||||||
|
|
||||||
|
if s.opts.AdminListen != "" {
|
||||||
|
aln, err := net.Listen("tcp", s.opts.AdminListen)
|
||||||
|
if err != nil {
|
||||||
|
_ = s.Close()
|
||||||
|
return fmt.Errorf("admin_listen %s: %w", s.opts.AdminListen, err)
|
||||||
|
}
|
||||||
|
s.adminLn = aln
|
||||||
|
s.adminAddr = aln.Addr().String()
|
||||||
|
if s.writeAdminAddr {
|
||||||
|
if err := writeAddrFile(s.opts.DataDir, "admin.addr", s.adminAddr); err != nil {
|
||||||
|
_ = s.Close()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
adminHandler := s.opts.AdminHandler
|
||||||
|
if adminHandler == nil {
|
||||||
|
adminHandler = http.NotFoundHandler()
|
||||||
|
}
|
||||||
|
s.adminHTTP = NewChanListener(aln.Addr(), 64)
|
||||||
|
s.adminSrv = &http.Server{
|
||||||
|
Handler: adminHandler,
|
||||||
|
ReadHeaderTimeout: 10 * time.Second,
|
||||||
|
}
|
||||||
|
s.wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer s.wg.Done()
|
||||||
|
err := s.adminSrv.Serve(s.adminHTTP)
|
||||||
|
if err != nil && !errors.Is(err, http.ErrServerClosed) {
|
||||||
|
s.log.Error("admin http serve", "err", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
s.wg.Add(1)
|
||||||
|
go s.acceptLoop(aln, true)
|
||||||
|
}
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
<-ctx.Done()
|
||||||
|
_ = s.Close()
|
||||||
|
}()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ListenAddr 返回端监听实际地址。
|
||||||
|
func (s *Server) ListenAddr() string { return s.listenAddr }
|
||||||
|
|
||||||
|
// AdminAddr 返回后台监听实际地址(未分离时为空)。
|
||||||
|
func (s *Server) AdminAddr() string { return s.adminAddr }
|
||||||
|
|
||||||
|
// Proxies 返回受信任代理集合。
|
||||||
|
func (s *Server) Proxies() *ProxySet { return s.proxies }
|
||||||
|
|
||||||
|
// TLSConfig 返回当前 TLS 配置(未配置证书时为 nil)。
|
||||||
|
func (s *Server) TLSConfig() *tls.Config {
|
||||||
|
if s.certs == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return s.certs.TLSConfig()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close 停止接受并关闭 HTTP。
|
||||||
|
func (s *Server) Close() error {
|
||||||
|
var first error
|
||||||
|
s.closeOnce.Do(func() {
|
||||||
|
close(s.closed)
|
||||||
|
if s.clientLn != nil {
|
||||||
|
_ = s.clientLn.Close()
|
||||||
|
}
|
||||||
|
if s.adminLn != nil {
|
||||||
|
_ = s.adminLn.Close()
|
||||||
|
}
|
||||||
|
shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
if s.clientSrv != nil {
|
||||||
|
if err := s.clientSrv.Shutdown(shutdownCtx); err != nil && first == nil {
|
||||||
|
first = err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if s.adminSrv != nil {
|
||||||
|
if err := s.adminSrv.Shutdown(shutdownCtx); err != nil && first == nil {
|
||||||
|
first = err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if s.clientHTTP != nil {
|
||||||
|
_ = s.clientHTTP.Close()
|
||||||
|
}
|
||||||
|
if s.adminHTTP != nil {
|
||||||
|
_ = s.adminHTTP.Close()
|
||||||
|
}
|
||||||
|
if s.certs != nil {
|
||||||
|
s.certs.Close()
|
||||||
|
}
|
||||||
|
})
|
||||||
|
s.wg.Wait()
|
||||||
|
return first
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) acceptLoop(ln net.Listener, isAdmin bool) {
|
||||||
|
defer s.wg.Done()
|
||||||
|
for {
|
||||||
|
c, err := ln.Accept()
|
||||||
|
if err != nil {
|
||||||
|
select {
|
||||||
|
case <-s.closed:
|
||||||
|
return
|
||||||
|
default:
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
s.wg.Add(1)
|
||||||
|
go func(conn net.Conn) {
|
||||||
|
defer s.wg.Done()
|
||||||
|
s.handleConn(conn, isAdmin)
|
||||||
|
}(c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) handleConn(conn net.Conn, isAdmin bool) {
|
||||||
|
kind, out, err := s.classify(conn, isAdmin, false)
|
||||||
|
if err != nil || kind == KindClosed {
|
||||||
|
_ = conn.Close()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
switch kind {
|
||||||
|
case KindHTTP:
|
||||||
|
httpLn := s.clientHTTP
|
||||||
|
if isAdmin {
|
||||||
|
httpLn = s.adminHTTP
|
||||||
|
}
|
||||||
|
if httpLn == nil {
|
||||||
|
_ = out.Close()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
httpLn.Enqueue(out)
|
||||||
|
case KindMQTT:
|
||||||
|
if s.opts.OnMQTT != nil {
|
||||||
|
s.opts.OnMQTT(out)
|
||||||
|
} else {
|
||||||
|
_ = out.Close()
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
_ = out.Close()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// classify 读首字节分流;afterTLS 表示已在 TLS 内层再识别。
|
||||||
|
func (s *Server) classify(conn net.Conn, isAdmin, afterTLS bool) (Kind, net.Conn, error) {
|
||||||
|
c, b, err := peekFirstByte(conn)
|
||||||
|
if err != nil {
|
||||||
|
return KindClosed, nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
hasCert := s.certs != nil
|
||||||
|
allowPlain := s.opts.AllowPlaintext || !hasCert
|
||||||
|
|
||||||
|
if b == 0x16 {
|
||||||
|
if !hasCert {
|
||||||
|
return KindClosed, nil, errors.New("tls client hello but no certificate")
|
||||||
|
}
|
||||||
|
if afterTLS {
|
||||||
|
return KindClosed, nil, errors.New("nested tls")
|
||||||
|
}
|
||||||
|
tlsConn := tls.Server(c, s.certs.TLSConfig())
|
||||||
|
if err := tlsConn.Handshake(); err != nil {
|
||||||
|
return KindClosed, nil, err
|
||||||
|
}
|
||||||
|
return s.classify(tlsConn, isAdmin, true)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !allowPlain && !afterTLS {
|
||||||
|
// 配了证书且未允许明文:非 TLS 首字节直接关
|
||||||
|
return KindClosed, nil, errors.New("plaintext not allowed")
|
||||||
|
}
|
||||||
|
|
||||||
|
if b >= 'A' && b <= 'Z' {
|
||||||
|
return KindHTTP, c, nil
|
||||||
|
}
|
||||||
|
if b == 0x10 {
|
||||||
|
if isAdmin {
|
||||||
|
return KindClosed, nil, errors.New("mqtt not allowed on admin_listen")
|
||||||
|
}
|
||||||
|
return KindMQTT, c, nil
|
||||||
|
}
|
||||||
|
return KindClosed, nil, errors.New("unknown first byte")
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeAddrFile(dataDir, name, addr string) error {
|
||||||
|
if dataDir == "" {
|
||||||
|
return errors.New("data_dir required to write addr file")
|
||||||
|
}
|
||||||
|
if err := os.MkdirAll(dataDir, 0o755); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return os.WriteFile(filepath.Join(dataDir, name), []byte(addr+"\n"), 0o644)
|
||||||
|
}
|
||||||
|
|
||||||
|
func isPortZero(addr string) bool {
|
||||||
|
_, port, err := net.SplitHostPort(addr)
|
||||||
|
if err != nil {
|
||||||
|
return strings.HasSuffix(addr, ":0") || addr == ":0"
|
||||||
|
}
|
||||||
|
return port == "0"
|
||||||
|
}
|
||||||
@@ -0,0 +1,117 @@
|
|||||||
|
package listener
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/tls"
|
||||||
|
"log/slog"
|
||||||
|
"os"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// CertReloader 按文件修改时间每小时重载证书;失败继续用旧证书。
|
||||||
|
type CertReloader struct {
|
||||||
|
certFile string
|
||||||
|
keyFile string
|
||||||
|
log *slog.Logger
|
||||||
|
|
||||||
|
mu sync.RWMutex
|
||||||
|
cert *tls.Certificate
|
||||||
|
certMod time.Time
|
||||||
|
keyMod time.Time
|
||||||
|
stop chan struct{}
|
||||||
|
stopOnce sync.Once
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewCertReloader 立即加载一次证书。
|
||||||
|
func NewCertReloader(certFile, keyFile string, log *slog.Logger) (*CertReloader, error) {
|
||||||
|
if log == nil {
|
||||||
|
log = slog.Default()
|
||||||
|
}
|
||||||
|
r := &CertReloader{
|
||||||
|
certFile: certFile,
|
||||||
|
keyFile: keyFile,
|
||||||
|
log: log,
|
||||||
|
stop: make(chan struct{}),
|
||||||
|
}
|
||||||
|
if err := r.reload(true); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
go r.loop()
|
||||||
|
return r, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetCertificate 供 tls.Config.GetCertificate 使用。
|
||||||
|
func (r *CertReloader) GetCertificate(*tls.ClientHelloInfo) (*tls.Certificate, error) {
|
||||||
|
r.mu.RLock()
|
||||||
|
defer r.mu.RUnlock()
|
||||||
|
return r.cert, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Certificate 返回当前证书(测试用)。
|
||||||
|
func (r *CertReloader) Certificate() *tls.Certificate {
|
||||||
|
r.mu.RLock()
|
||||||
|
defer r.mu.RUnlock()
|
||||||
|
return r.cert
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close 停止重载循环。
|
||||||
|
func (r *CertReloader) Close() {
|
||||||
|
r.stopOnce.Do(func() { close(r.stop) })
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *CertReloader) loop() {
|
||||||
|
t := time.NewTicker(time.Hour)
|
||||||
|
defer t.Stop()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-r.stop:
|
||||||
|
return
|
||||||
|
case <-t.C:
|
||||||
|
if err := r.reload(false); err != nil {
|
||||||
|
r.log.Error("tls cert reload failed, keeping old cert", "err", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ReloadNow 立即按 mtime 检查并重载(测试用)。
|
||||||
|
func (r *CertReloader) ReloadNow() error {
|
||||||
|
return r.reload(false)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *CertReloader) reload(force bool) error {
|
||||||
|
certInfo, err := os.Stat(r.certFile)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
keyInfo, err := os.Stat(r.keyFile)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
r.mu.RLock()
|
||||||
|
same := !force && certInfo.ModTime().Equal(r.certMod) && keyInfo.ModTime().Equal(r.keyMod) && r.cert != nil
|
||||||
|
r.mu.RUnlock()
|
||||||
|
if same {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
cert, err := tls.LoadX509KeyPair(r.certFile, r.keyFile)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
r.mu.Lock()
|
||||||
|
r.cert = &cert
|
||||||
|
r.certMod = certInfo.ModTime()
|
||||||
|
r.keyMod = keyInfo.ModTime()
|
||||||
|
r.mu.Unlock()
|
||||||
|
r.log.Info("tls certificate loaded", "cert", r.certFile)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// TLSConfig 构造不设 NextProtos 的服务端 TLS 配置。
|
||||||
|
func (r *CertReloader) TLSConfig() *tls.Config {
|
||||||
|
return &tls.Config{
|
||||||
|
GetCertificate: r.GetCertificate,
|
||||||
|
MinVersion: tls.VersionTLS12,
|
||||||
|
// 故意不设 NextProtos,以便声明 ALPN mqtt 的客户端仍能握手。
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,115 @@
|
|||||||
|
package metrics
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/prometheus/client_golang/prometheus"
|
||||||
|
"github.com/prometheus/client_golang/prometheus/promhttp"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Registry 持有 NixMsg 指标与独立注册表(避免污染默认全局注册表)。
|
||||||
|
type Registry struct {
|
||||||
|
reg *prometheus.Registry
|
||||||
|
|
||||||
|
Connections *prometheus.GaugeVec
|
||||||
|
EndpointsTotal prometheus.Gauge
|
||||||
|
DeliveriesPending prometheus.Gauge
|
||||||
|
MessagesScheduled prometheus.Gauge
|
||||||
|
DispatchToPushSeconds prometheus.Observer
|
||||||
|
AckSeconds prometheus.Observer
|
||||||
|
WriteQueueLength prometheus.Gauge
|
||||||
|
WriteCommitSeconds prometheus.Observer
|
||||||
|
PasswordHashQueue prometheus.Gauge
|
||||||
|
ErrorsTotal *prometheus.CounterVec
|
||||||
|
}
|
||||||
|
|
||||||
|
// New 按 DEVELOPMENT 4.3 注册指标。
|
||||||
|
func New() *Registry {
|
||||||
|
reg := prometheus.NewRegistry()
|
||||||
|
r := &Registry{reg: reg}
|
||||||
|
|
||||||
|
r.Connections = prometheus.NewGaugeVec(prometheus.GaugeOpts{
|
||||||
|
Name: "nixmsg_connections",
|
||||||
|
Help: "Current online connections by transport (ws or tcp).",
|
||||||
|
}, []string{"transport"})
|
||||||
|
|
||||||
|
r.EndpointsTotal = prometheus.NewGauge(prometheus.GaugeOpts{
|
||||||
|
Name: "nixmsg_endpoints",
|
||||||
|
Help: "Total number of endpoints.",
|
||||||
|
})
|
||||||
|
|
||||||
|
r.DeliveriesPending = prometheus.NewGauge(prometheus.GaugeOpts{
|
||||||
|
Name: "nixmsg_deliveries_pending",
|
||||||
|
Help: "Number of pending deliveries.",
|
||||||
|
})
|
||||||
|
|
||||||
|
r.MessagesScheduled = prometheus.NewGauge(prometheus.GaugeOpts{
|
||||||
|
Name: "nixmsg_messages_scheduled",
|
||||||
|
Help: "Number of scheduled messages not yet dispatched.",
|
||||||
|
})
|
||||||
|
|
||||||
|
dispatchHist := prometheus.NewHistogram(prometheus.HistogramOpts{
|
||||||
|
Name: "nixmsg_dispatch_to_push_duration_seconds",
|
||||||
|
Help: "Latency from message due time to push.",
|
||||||
|
Buckets: prometheus.DefBuckets,
|
||||||
|
})
|
||||||
|
r.DispatchToPushSeconds = dispatchHist
|
||||||
|
|
||||||
|
ackHist := prometheus.NewHistogram(prometheus.HistogramOpts{
|
||||||
|
Name: "nixmsg_ack_duration_seconds",
|
||||||
|
Help: "Latency from push to ack.",
|
||||||
|
Buckets: prometheus.DefBuckets,
|
||||||
|
})
|
||||||
|
r.AckSeconds = ackHist
|
||||||
|
|
||||||
|
r.WriteQueueLength = prometheus.NewGauge(prometheus.GaugeOpts{
|
||||||
|
Name: "nixmsg_write_queue_length",
|
||||||
|
Help: "Number of write operations waiting or in the current batch.",
|
||||||
|
})
|
||||||
|
|
||||||
|
commitHist := prometheus.NewHistogram(prometheus.HistogramOpts{
|
||||||
|
Name: "nixmsg_write_batch_commit_duration_seconds",
|
||||||
|
Help: "Duration of each merged write-queue commit.",
|
||||||
|
Buckets: prometheus.DefBuckets,
|
||||||
|
})
|
||||||
|
r.WriteCommitSeconds = commitHist
|
||||||
|
|
||||||
|
r.PasswordHashQueue = prometheus.NewGauge(prometheus.GaugeOpts{
|
||||||
|
Name: "nixmsg_password_hash_queue_length",
|
||||||
|
Help: "Number of password hash jobs waiting for a pool slot.",
|
||||||
|
})
|
||||||
|
|
||||||
|
r.ErrorsTotal = prometheus.NewCounterVec(prometheus.CounterOpts{
|
||||||
|
Name: "nixmsg_errors_total",
|
||||||
|
Help: "Count of application error codes.",
|
||||||
|
}, []string{"code"})
|
||||||
|
|
||||||
|
reg.MustRegister(
|
||||||
|
r.Connections,
|
||||||
|
r.EndpointsTotal,
|
||||||
|
r.DeliveriesPending,
|
||||||
|
r.MessagesScheduled,
|
||||||
|
dispatchHist,
|
||||||
|
ackHist,
|
||||||
|
r.WriteQueueLength,
|
||||||
|
commitHist,
|
||||||
|
r.PasswordHashQueue,
|
||||||
|
r.ErrorsTotal,
|
||||||
|
)
|
||||||
|
|
||||||
|
// 初始化 transport 标签,便于空载抓取也能看到序列。
|
||||||
|
r.Connections.WithLabelValues("ws").Set(0)
|
||||||
|
r.Connections.WithLabelValues("tcp").Set(0)
|
||||||
|
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
|
||||||
|
// Handler 返回 Prometheus 文本格式的 /metrics 处理函数(访问规则由 A 线加)。
|
||||||
|
func (r *Registry) Handler() http.Handler {
|
||||||
|
return promhttp.HandlerFor(r.reg, promhttp.HandlerOpts{})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Gatherer 暴露底层 Gatherer(测试用)。
|
||||||
|
func (r *Registry) Gatherer() prometheus.Gatherer {
|
||||||
|
return r.reg
|
||||||
|
}
|
||||||
@@ -0,0 +1,58 @@
|
|||||||
|
package metrics
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestMetricsHandlerExposesText(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
m := New()
|
||||||
|
m.Connections.WithLabelValues("ws").Set(3)
|
||||||
|
m.Connections.WithLabelValues("tcp").Set(2)
|
||||||
|
m.EndpointsTotal.Set(10)
|
||||||
|
m.DeliveriesPending.Set(4)
|
||||||
|
m.MessagesScheduled.Set(1)
|
||||||
|
m.DispatchToPushSeconds.Observe(0.05)
|
||||||
|
m.AckSeconds.Observe(0.02)
|
||||||
|
m.WriteQueueLength.Set(7)
|
||||||
|
m.WriteCommitSeconds.Observe(0.001)
|
||||||
|
m.PasswordHashQueue.Set(2)
|
||||||
|
m.ErrorsTotal.WithLabelValues("busy").Inc()
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/metrics", nil)
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
m.Handler().ServeHTTP(rec, req)
|
||||||
|
resp := rec.Result()
|
||||||
|
defer func() { _ = resp.Body.Close() }()
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
t.Fatalf("status=%d", resp.StatusCode)
|
||||||
|
}
|
||||||
|
body, err := io.ReadAll(resp.Body)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
text := string(body)
|
||||||
|
for _, want := range []string{
|
||||||
|
"nixmsg_connections",
|
||||||
|
`transport="ws"`,
|
||||||
|
`transport="tcp"`,
|
||||||
|
"nixmsg_endpoints",
|
||||||
|
"nixmsg_deliveries_pending",
|
||||||
|
"nixmsg_messages_scheduled",
|
||||||
|
"nixmsg_dispatch_to_push_duration_seconds",
|
||||||
|
"nixmsg_ack_duration_seconds",
|
||||||
|
"nixmsg_write_queue_length",
|
||||||
|
"nixmsg_write_batch_commit_duration_seconds",
|
||||||
|
"nixmsg_password_hash_queue_length",
|
||||||
|
"nixmsg_errors_total",
|
||||||
|
`code="busy"`,
|
||||||
|
} {
|
||||||
|
if !strings.Contains(text, want) {
|
||||||
|
t.Fatalf("metrics text missing %q\n%s", want, text)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,47 @@
|
|||||||
|
package store
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
const settingAdminPasswordHash = "admin_password_hash"
|
||||||
|
|
||||||
|
// HasAdminPassword 检查 settings 中是否已有管理员密码哈希。
|
||||||
|
func HasAdminPassword(ctx context.Context, db *sql.DB) (bool, error) {
|
||||||
|
var value string
|
||||||
|
err := db.QueryRowContext(ctx, `SELECT value FROM settings WHERE key = ?`, settingAdminPasswordHash).Scan(&value)
|
||||||
|
if err == sql.ErrNoRows {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
return value != "", nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetAdminPasswordHash 写入或覆盖管理员密码哈希。
|
||||||
|
func SetAdminPasswordHash(ctx context.Context, db *sql.DB, phc string) error {
|
||||||
|
now := time.Now().UnixMilli()
|
||||||
|
_, err := db.ExecContext(ctx, `
|
||||||
|
INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)
|
||||||
|
ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = excluded.updated_at
|
||||||
|
`, settingAdminPasswordHash, phc, now)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// VacuumInto 对打开的写连接执行 VACUUM INTO(可用于运行中备份)。
|
||||||
|
func VacuumInto(ctx context.Context, db *sql.DB, outPath string) error {
|
||||||
|
if outPath == "" {
|
||||||
|
return fmt.Errorf("store: vacuum into path empty")
|
||||||
|
}
|
||||||
|
escaped := strings.ReplaceAll(outPath, "'", "''")
|
||||||
|
_, err := db.ExecContext(ctx, "VACUUM INTO '"+escaped+"'")
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("vacuum into: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
+254
-20
@@ -4,29 +4,57 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"errors"
|
"errors"
|
||||||
|
"fmt"
|
||||||
"sync"
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
// ErrQueueClosed 表示写入队列已关闭。
|
// ErrQueueClosed 表示写入队列已关闭。
|
||||||
var ErrQueueClosed = errors.New("store: write queue closed")
|
var ErrQueueClosed = errors.New("store: write queue closed")
|
||||||
|
|
||||||
|
// ErrBusy 表示写库基础设施失败(可映射为协议 busy;/readyz 应失败)。
|
||||||
|
var ErrBusy = errors.New("store: busy")
|
||||||
|
|
||||||
|
const (
|
||||||
|
maxBatchOps = 256
|
||||||
|
batchWait = 2 * time.Millisecond
|
||||||
|
queueBuffSize = 1024
|
||||||
|
)
|
||||||
|
|
||||||
// WriteFunc 在单个写事务中执行的操作。
|
// WriteFunc 在单个写事务中执行的操作。
|
||||||
type WriteFunc func(tx *sql.Tx) error
|
type WriteFunc func(tx *sql.Tx) error
|
||||||
|
|
||||||
// Queue 写入队列:提交一个写操作并拿到结果。
|
type writeJob struct {
|
||||||
//
|
ctx context.Context
|
||||||
// 本任务(T0.3)实现为互斥串行的一操作一事务,不做合并;DEVELOPMENT 7.2
|
fn WriteFunc
|
||||||
// 要求的写 goroutine 合并提交(最多 256 个或凑满 2ms、SAVEPOINT 隔离失败)留给 P2。
|
res chan error
|
||||||
|
}
|
||||||
|
|
||||||
|
// Queue 写入队列:写 goroutine 合并提交(最多 256 个或凑满 2ms),每操作用 SAVEPOINT 隔离。
|
||||||
type Queue struct {
|
type Queue struct {
|
||||||
db *sql.DB
|
db *sql.DB
|
||||||
|
|
||||||
mu sync.Mutex
|
ch chan writeJob
|
||||||
closed bool
|
done chan struct{}
|
||||||
|
closed atomic.Bool
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
|
ready bool
|
||||||
|
lastWriteErr error
|
||||||
|
pending int
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewQueue 创建简单写入队列(一操作一事务)。
|
// NewQueue 创建合并写入队列并启动写 goroutine。
|
||||||
func NewQueue(db *sql.DB) *Queue {
|
func NewQueue(db *sql.DB) *Queue {
|
||||||
return &Queue{db: db}
|
q := &Queue{
|
||||||
|
db: db,
|
||||||
|
ch: make(chan writeJob, queueBuffSize),
|
||||||
|
done: make(chan struct{}),
|
||||||
|
ready: true,
|
||||||
|
}
|
||||||
|
go q.loop()
|
||||||
|
return q
|
||||||
}
|
}
|
||||||
|
|
||||||
// Do 提交写操作并等待提交结果。
|
// Do 提交写操作并等待提交结果。
|
||||||
@@ -34,29 +62,235 @@ func (q *Queue) Do(ctx context.Context, fn WriteFunc) error {
|
|||||||
if fn == nil {
|
if fn == nil {
|
||||||
return errors.New("store: nil write func")
|
return errors.New("store: nil write func")
|
||||||
}
|
}
|
||||||
q.mu.Lock()
|
if q.closed.Load() {
|
||||||
defer q.mu.Unlock()
|
|
||||||
if q.closed {
|
|
||||||
return ErrQueueClosed
|
return ErrQueueClosed
|
||||||
}
|
}
|
||||||
if err := ctx.Err(); err != nil {
|
if err := ctx.Err(); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
tx, err := q.db.BeginTx(ctx, nil)
|
job := writeJob{ctx: ctx, fn: fn, res: make(chan error, 1)}
|
||||||
if err != nil {
|
q.mu.Lock()
|
||||||
return err
|
q.pending++
|
||||||
|
q.mu.Unlock()
|
||||||
|
select {
|
||||||
|
case q.ch <- job:
|
||||||
|
case <-ctx.Done():
|
||||||
|
q.mu.Lock()
|
||||||
|
q.pending--
|
||||||
|
q.mu.Unlock()
|
||||||
|
return ctx.Err()
|
||||||
|
case <-q.done:
|
||||||
|
q.mu.Lock()
|
||||||
|
q.pending--
|
||||||
|
q.mu.Unlock()
|
||||||
|
return ErrQueueClosed
|
||||||
}
|
}
|
||||||
if err := fn(tx); err != nil {
|
select {
|
||||||
_ = tx.Rollback()
|
case err := <-job.res:
|
||||||
return err
|
return err
|
||||||
|
case <-ctx.Done():
|
||||||
|
// 操作可能仍在队列中执行;结果通道仍会被写端关闭式填入。
|
||||||
|
select {
|
||||||
|
case err := <-job.res:
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return ctx.Err()
|
||||||
|
case <-time.After(30 * time.Second):
|
||||||
|
return ctx.Err()
|
||||||
|
}
|
||||||
}
|
}
|
||||||
return tx.Commit()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close 关闭队列,之后 Do 返回 ErrQueueClosed。
|
func (q *Queue) loop() {
|
||||||
func (q *Queue) Close() error {
|
defer close(q.done)
|
||||||
|
for {
|
||||||
|
job, ok := <-q.ch
|
||||||
|
if !ok {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
batch := []writeJob{job}
|
||||||
|
timer := time.NewTimer(batchWait)
|
||||||
|
collect:
|
||||||
|
for len(batch) < maxBatchOps {
|
||||||
|
select {
|
||||||
|
case j, ok := <-q.ch:
|
||||||
|
if !ok {
|
||||||
|
break collect
|
||||||
|
}
|
||||||
|
batch = append(batch, j)
|
||||||
|
case <-timer.C:
|
||||||
|
break collect
|
||||||
|
}
|
||||||
|
}
|
||||||
|
timer.Stop()
|
||||||
|
q.runBatch(batch)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (q *Queue) runBatch(batch []writeJob) {
|
||||||
|
defer func() {
|
||||||
|
q.mu.Lock()
|
||||||
|
q.pending -= len(batch)
|
||||||
|
if q.pending < 0 {
|
||||||
|
q.pending = 0
|
||||||
|
}
|
||||||
|
q.mu.Unlock()
|
||||||
|
}()
|
||||||
|
|
||||||
|
// 过滤已取消的任务。
|
||||||
|
active := make([]writeJob, 0, len(batch))
|
||||||
|
for _, j := range batch {
|
||||||
|
if err := j.ctx.Err(); err != nil {
|
||||||
|
j.res <- err
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
active = append(active, j)
|
||||||
|
}
|
||||||
|
if len(active) == 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
tx, err := q.db.BeginTx(context.Background(), nil)
|
||||||
|
if err != nil {
|
||||||
|
q.markBusy(err)
|
||||||
|
for _, j := range active {
|
||||||
|
j.res <- errors.Join(ErrBusy, err)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
type outcome struct {
|
||||||
|
job writeJob
|
||||||
|
opErr error
|
||||||
|
success bool
|
||||||
|
}
|
||||||
|
outcomes := make([]outcome, 0, len(active))
|
||||||
|
for i, j := range active {
|
||||||
|
sp := fmt.Sprintf("sp_%d", i)
|
||||||
|
if _, err := tx.Exec("SAVEPOINT " + sp); err != nil {
|
||||||
|
_ = tx.Rollback()
|
||||||
|
q.markBusy(err)
|
||||||
|
for _, o := range outcomes {
|
||||||
|
if o.success {
|
||||||
|
o.job.res <- errors.Join(ErrBusy, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, rest := range active[i:] {
|
||||||
|
rest.res <- errors.Join(ErrBusy, err)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
opErr := j.fn(tx)
|
||||||
|
if opErr != nil {
|
||||||
|
if _, rbErr := tx.Exec("ROLLBACK TO " + sp); rbErr != nil {
|
||||||
|
_ = tx.Rollback()
|
||||||
|
q.markBusy(rbErr)
|
||||||
|
for _, o := range outcomes {
|
||||||
|
if o.success {
|
||||||
|
o.job.res <- errors.Join(ErrBusy, rbErr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
j.res <- opErr
|
||||||
|
for _, rest := range active[i+1:] {
|
||||||
|
rest.res <- errors.Join(ErrBusy, rbErr)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
_, _ = tx.Exec("RELEASE " + sp)
|
||||||
|
outcomes = append(outcomes, outcome{job: j, opErr: opErr, success: false})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if _, err := tx.Exec("RELEASE " + sp); err != nil {
|
||||||
|
_ = tx.Rollback()
|
||||||
|
q.markBusy(err)
|
||||||
|
for _, o := range outcomes {
|
||||||
|
if o.success {
|
||||||
|
o.job.res <- errors.Join(ErrBusy, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
j.res <- errors.Join(ErrBusy, err)
|
||||||
|
for _, rest := range active[i+1:] {
|
||||||
|
rest.res <- errors.Join(ErrBusy, err)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
outcomes = append(outcomes, outcome{job: j, success: true})
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := tx.Commit(); err != nil {
|
||||||
|
q.markBusy(err)
|
||||||
|
for _, o := range outcomes {
|
||||||
|
if o.success {
|
||||||
|
o.job.res <- errors.Join(ErrBusy, err)
|
||||||
|
} else {
|
||||||
|
o.job.res <- o.opErr
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
for _, o := range outcomes {
|
||||||
|
if o.success {
|
||||||
|
o.job.res <- nil
|
||||||
|
} else {
|
||||||
|
o.job.res <- o.opErr
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (q *Queue) markBusy(err error) {
|
||||||
q.mu.Lock()
|
q.mu.Lock()
|
||||||
defer q.mu.Unlock()
|
defer q.mu.Unlock()
|
||||||
q.closed = true
|
q.ready = false
|
||||||
|
q.lastWriteErr = err
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsReady 写库是否仍可用(写失败后为 false)。
|
||||||
|
func (q *Queue) IsReady() bool {
|
||||||
|
q.mu.Lock()
|
||||||
|
defer q.mu.Unlock()
|
||||||
|
return q.ready
|
||||||
|
}
|
||||||
|
|
||||||
|
// LastWriteError 返回最近一次基础设施写失败。
|
||||||
|
func (q *Queue) LastWriteError() error {
|
||||||
|
q.mu.Lock()
|
||||||
|
defer q.mu.Unlock()
|
||||||
|
return q.lastWriteErr
|
||||||
|
}
|
||||||
|
|
||||||
|
// Len 返回排队与批次中的写操作数。
|
||||||
|
func (q *Queue) Len() int {
|
||||||
|
q.mu.Lock()
|
||||||
|
defer q.mu.Unlock()
|
||||||
|
return q.pending
|
||||||
|
}
|
||||||
|
|
||||||
|
// Drain 等待已入队操作完成,或 ctx 取消。关闭后也会等到 loop 退出。
|
||||||
|
func (q *Queue) Drain(ctx context.Context) error {
|
||||||
|
ticker := time.NewTicker(5 * time.Millisecond)
|
||||||
|
defer ticker.Stop()
|
||||||
|
for {
|
||||||
|
q.mu.Lock()
|
||||||
|
n := q.pending
|
||||||
|
q.mu.Unlock()
|
||||||
|
if n == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return ctx.Err()
|
||||||
|
case <-ticker.C:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close 关闭队列:不再接受新任务,并等待写 goroutine 处理完已入队任务后退出。
|
||||||
|
func (q *Queue) Close() error {
|
||||||
|
if q.closed.Swap(true) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
close(q.ch)
|
||||||
|
<-q.done
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,175 @@
|
|||||||
|
package store
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestWriteQueueSavepointIsolatesFailure(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
dir := t.TempDir()
|
||||||
|
db, err := Open(dir, "FULL")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer func() { _ = db.Close() }()
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
wg.Add(2)
|
||||||
|
errOK := make(chan error, 1)
|
||||||
|
errBad := make(chan error, 1)
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
errOK <- db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||||
|
_, e := tx.Exec(
|
||||||
|
`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)`,
|
||||||
|
"k_ok", "1", time.Now().UnixMilli(),
|
||||||
|
)
|
||||||
|
return e
|
||||||
|
})
|
||||||
|
}()
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
// 稍等,尽量与成功操作进同一批。
|
||||||
|
time.Sleep(500 * time.Microsecond)
|
||||||
|
errBad <- db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||||
|
return errors.New("forced op failure")
|
||||||
|
})
|
||||||
|
}()
|
||||||
|
wg.Wait()
|
||||||
|
|
||||||
|
if e := <-errOK; e != nil {
|
||||||
|
t.Fatalf("ok op: %v", e)
|
||||||
|
}
|
||||||
|
if e := <-errBad; e == nil || e.Error() != "forced op failure" {
|
||||||
|
t.Fatalf("bad op: %v", e)
|
||||||
|
}
|
||||||
|
|
||||||
|
var value string
|
||||||
|
if scanErr := db.Read.QueryRow(`SELECT value FROM settings WHERE key = ?`, "k_ok").Scan(&value); scanErr != nil {
|
||||||
|
t.Fatalf("ok row missing: %v", scanErr)
|
||||||
|
}
|
||||||
|
if value != "1" {
|
||||||
|
t.Fatalf("value=%q", value)
|
||||||
|
}
|
||||||
|
if !db.Queue.IsReady() {
|
||||||
|
t.Fatal("op failure should not mark queue busy")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteQueueBatchCommit(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
dir := t.TempDir()
|
||||||
|
db, err := Open(dir, "FULL")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer func() { _ = db.Close() }()
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
const n = 32
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
wg.Add(n)
|
||||||
|
errs := make([]error, n)
|
||||||
|
for i := 0; i < n; i++ {
|
||||||
|
i := i
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
key := fmt.Sprintf("batch_%d", i)
|
||||||
|
errs[i] = db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||||
|
_, e := tx.Exec(
|
||||||
|
`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)`,
|
||||||
|
key, "1", time.Now().UnixMilli(),
|
||||||
|
)
|
||||||
|
return e
|
||||||
|
})
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
for i, e := range errs {
|
||||||
|
if e != nil {
|
||||||
|
t.Fatalf("op %d: %v", i, e)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
var count int
|
||||||
|
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM settings WHERE key LIKE 'batch_%'`).Scan(&count); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if count != n {
|
||||||
|
t.Fatalf("count=%d want %d", count, n)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteQueueNoBacklogAt200PerSec(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
dir := t.TempDir()
|
||||||
|
db, err := Open(dir, "FULL")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer func() { _ = db.Close() }()
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
const total = 200
|
||||||
|
start := time.Now()
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
wg.Add(total)
|
||||||
|
for i := 0; i < total; i++ {
|
||||||
|
i := i
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
key := fmt.Sprintf("load_%d", i)
|
||||||
|
if e := db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||||
|
_, e := tx.Exec(
|
||||||
|
`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)`,
|
||||||
|
key, "1", time.Now().UnixMilli(),
|
||||||
|
)
|
||||||
|
return e
|
||||||
|
}); e != nil {
|
||||||
|
t.Errorf("op %d: %v", i, e)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
elapsed := time.Since(start)
|
||||||
|
if db.Queue.Len() != 0 {
|
||||||
|
t.Fatalf("queue backlog len=%d", db.Queue.Len())
|
||||||
|
}
|
||||||
|
// 200 条并发写入应在数秒内完成(合并提交);过长则合并未生效。
|
||||||
|
if elapsed > 5*time.Second {
|
||||||
|
t.Fatalf("200 writes took %s, queue likely not merging well", elapsed)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func BenchmarkWriteQueue200PerSec(b *testing.B) {
|
||||||
|
dir := b.TempDir()
|
||||||
|
db, err := Open(dir, "FULL")
|
||||||
|
if err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
defer func() { _ = db.Close() }()
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
key := fmt.Sprintf("bench_%d", i)
|
||||||
|
if err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||||
|
_, e := tx.Exec(
|
||||||
|
`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)
|
||||||
|
ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = excluded.updated_at`,
|
||||||
|
key, "1", time.Now().UnixMilli(),
|
||||||
|
)
|
||||||
|
return e
|
||||||
|
}); err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,26 @@
|
|||||||
|
package store
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Ready 检查读写库可用且写入队列未因磁盘类失败进入 busy。
|
||||||
|
func (d *DB) Ready(ctx context.Context) error {
|
||||||
|
if d == nil || d.Write == nil || d.Read == nil {
|
||||||
|
return errors.New("store: not open")
|
||||||
|
}
|
||||||
|
if err := d.Write.PingContext(ctx); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := d.Read.PingContext(ctx); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if d.Queue != nil && !d.Queue.IsReady() {
|
||||||
|
if err := d.Queue.LastWriteError(); err != nil {
|
||||||
|
return errors.Join(ErrBusy, err)
|
||||||
|
}
|
||||||
|
return ErrBusy
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user