423 lines
12 KiB
Go
423 lines
12 KiB
Go
package main
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"io/fs"
|
|
"log/slog"
|
|
"net"
|
|
"net/http"
|
|
"os"
|
|
"os/signal"
|
|
"path/filepath"
|
|
"strings"
|
|
"syscall"
|
|
"time"
|
|
|
|
"git.asio.asia/nixevol/NixMsg/internal/admin"
|
|
"git.asio.asia/nixevol/NixMsg/internal/app/group"
|
|
"git.asio.asia/nixevol/NixMsg/internal/app/identity"
|
|
"git.asio.asia/nixevol/NixMsg/internal/app/message"
|
|
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
|
"git.asio.asia/nixevol/NixMsg/internal/app/presence"
|
|
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
|
"git.asio.asia/nixevol/NixMsg/internal/broker"
|
|
"git.asio.asia/nixevol/NixMsg/internal/config"
|
|
"git.asio.asia/nixevol/NixMsg/internal/httpx"
|
|
"git.asio.asia/nixevol/NixMsg/internal/listener"
|
|
"git.asio.asia/nixevol/NixMsg/internal/metrics"
|
|
"git.asio.asia/nixevol/NixMsg/internal/store"
|
|
"git.asio.asia/nixevol/NixMsg/web"
|
|
)
|
|
|
|
func cmdServe(_ []string) error {
|
|
cfgPath := config.PathFromEnv()
|
|
cfg, err := config.Load(cfgPath)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if err := cfg.Validate(); err != nil {
|
|
return err
|
|
}
|
|
setupJSONLogger(cfg.Log)
|
|
ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
|
|
defer stop()
|
|
return runServe(withServeStop(ctx, stop), cfg)
|
|
}
|
|
|
|
type serveStopKey struct{}
|
|
|
|
func withServeStop(ctx context.Context, stop context.CancelFunc) context.Context {
|
|
return context.WithValue(ctx, serveStopKey{}, stop)
|
|
}
|
|
|
|
func invokeServeStop(ctx context.Context) {
|
|
stop, _ := ctx.Value(serveStopKey{}).(context.CancelFunc)
|
|
if stop != nil {
|
|
stop()
|
|
}
|
|
}
|
|
|
|
func runServe(ctx context.Context, cfg config.Config) error {
|
|
if err := os.MkdirAll(cfg.DataDir, 0o755); err != nil {
|
|
return fmt.Errorf("mkdir data_dir: %w", err)
|
|
}
|
|
|
|
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 password not initialized; run: nixmsg admin init")
|
|
}
|
|
|
|
hashPool := auth.NewPool()
|
|
sessionTokens := auth.NewSessionTokens()
|
|
apiTokens := auth.NewAPITokens()
|
|
loginLocks := auth.NewLoginLocks()
|
|
memConns := message.NewMemoryConns()
|
|
msgLim := message.LimitsFromFullConfig(cfg)
|
|
metricsReg := metrics.New()
|
|
db.Queue.OnBatchCommit = func(d time.Duration) {
|
|
metricsReg.WriteCommitSeconds.Observe(d.Seconds())
|
|
}
|
|
|
|
login := broker.NewLogin(broker.LoginOptions{
|
|
DB: db,
|
|
Pool: hashPool,
|
|
Tokens: sessionTokens,
|
|
Locks: loginLocks,
|
|
IdleDays: cfg.SessionIdleDays,
|
|
})
|
|
|
|
// 先用空下行建 message,拿到 App 指针;broker 建好后再 WithDownlink 重建并回填 uplink.msg。
|
|
msgApp := message.New(db, msgLim, hashPool,
|
|
message.WithLocks(loginLocks),
|
|
message.WithConnRegistry(memConns),
|
|
message.WithMetrics(metricsReg),
|
|
)
|
|
uplink := &appUplink{
|
|
msg: msgApp,
|
|
conns: memConns,
|
|
log: slog.Default(),
|
|
metrics: metricsReg,
|
|
}
|
|
sess := broker.NewSession(broker.SessionOptions{
|
|
Login: login,
|
|
Inner: uplink,
|
|
Limits: broker.HelloLimits{
|
|
MaxBodyBytes: cfg.Limits.MaxBodyBytes,
|
|
MaxMetaBytes: cfg.Limits.MaxMetaBytes,
|
|
MaxFrameBytes: cfg.Limits.MaxFrameBytes,
|
|
MaxTTLSeconds: int64(cfg.Limits.MaxTTLSeconds),
|
|
MaxScheduleSeconds: int64(cfg.Limits.MaxScheduleSeconds),
|
|
AckTimeoutSeconds: int64(cfg.Limits.AckTimeoutSeconds),
|
|
ServerVersion: Version,
|
|
},
|
|
Logger: slog.Default(),
|
|
})
|
|
|
|
brk, err := broker.New(broker.Options{
|
|
Authenticator: login,
|
|
Uplink: sess,
|
|
Logger: slog.Default(),
|
|
Metrics: metricsReg,
|
|
OnPublishDropped: func(dropCtx context.Context, endpointID string, connID port.ConnID, payload []byte) {
|
|
if dropErr := msgApp.OnPublishDropped(dropCtx, endpointID, connID, payload); dropErr != nil {
|
|
slog.Error("on publish dropped", "endpoint", endpointID, "err", dropErr)
|
|
}
|
|
},
|
|
})
|
|
if err != nil {
|
|
return fmt.Errorf("broker: %w", err)
|
|
}
|
|
defer func() { _ = brk.Close() }()
|
|
sess.Attach(brk)
|
|
|
|
msgApp = message.New(db, msgLim, hashPool,
|
|
message.WithLocks(loginLocks),
|
|
message.WithConnRegistry(memConns),
|
|
message.WithDownlink(brk),
|
|
message.WithMetrics(metricsReg),
|
|
)
|
|
uplink.msg = msgApp
|
|
uplink.down = brk
|
|
|
|
presApp := presence.New(presence.Config{
|
|
DB: db,
|
|
Downlink: brk,
|
|
Conns: &presenceConnTable{conns: memConns},
|
|
})
|
|
sess.SetPresence(presApp)
|
|
uplink.presence = presApp
|
|
|
|
trustedNets := httpx.ParseCIDRs(cfg.TrustedProxies)
|
|
idApp := identity.New(identity.Config{
|
|
DB: db,
|
|
Hash: hashPool,
|
|
Locks: loginLocks,
|
|
Sessions: sessionTokens,
|
|
MaxScheduleSeconds: int64(cfg.Limits.MaxScheduleSeconds),
|
|
Logger: slog.Default(),
|
|
ConnControl: nil, // B-04:踢线走 Session 钩子,避免 identity 20ms 异步 Disconnect
|
|
Downlink: brk,
|
|
ClientIP: func(r *http.Request) string {
|
|
return httpx.ClientIP(r, trustedNets)
|
|
},
|
|
})
|
|
uplink.identity = idApp
|
|
|
|
groupApp := group.New(group.Config{
|
|
DB: db,
|
|
Talk: idApp,
|
|
Online: presApp,
|
|
Downlink: brk,
|
|
MaxGroupMembers: cfg.Limits.MaxGroupMembers,
|
|
})
|
|
uplink.groups = groupApp
|
|
|
|
if recoverErr := msgApp.RecoverOnStart(ctx); recoverErr != nil {
|
|
return fmt.Errorf("message recover: %w", recoverErr)
|
|
}
|
|
|
|
adminHandler := admin.New(admin.Deps{
|
|
DB: db,
|
|
Hash: hashPool,
|
|
Tokens: apiTokens,
|
|
Locks: loginLocks,
|
|
Logger: slog.Default(),
|
|
AuditLogger: slog.New(slog.NewJSONHandler(os.Stdout, &slog.HandlerOptions{Level: slog.LevelInfo})),
|
|
TrustedProxies: trustedNets,
|
|
Identity: idApp,
|
|
Groups: groupApp,
|
|
Config: cfg,
|
|
Version: Version,
|
|
// Kick:只断开,令牌不变,SDK 重连(PRD 踢下线)。
|
|
KickEndpoint: func(kickCtx context.Context, endpointID string) (bool, error) {
|
|
if _, found := brk.ConnInfoOf(endpointID); !found {
|
|
return false, nil
|
|
}
|
|
if kickErr := sess.Kick(kickCtx, endpointID); kickErr != nil {
|
|
return false, kickErr
|
|
}
|
|
return true, nil
|
|
},
|
|
// 停用/删除/重置:先 fatal 再断开(DEVELOPMENT 6.8)。
|
|
DisableKick: func(kickCtx context.Context, endpointID string) (bool, error) {
|
|
if _, found := brk.ConnInfoOf(endpointID); !found {
|
|
return false, nil
|
|
}
|
|
if disableErr := sess.Disable(kickCtx, endpointID); disableErr != nil {
|
|
return false, disableErr
|
|
}
|
|
return true, nil
|
|
},
|
|
DeleteKick: func(kickCtx context.Context, endpointID string) (bool, error) {
|
|
if _, found := brk.ConnInfoOf(endpointID); !found {
|
|
return false, nil
|
|
}
|
|
if deleteErr := sess.Deleted(kickCtx, endpointID); deleteErr != nil {
|
|
return false, deleteErr
|
|
}
|
|
return true, nil
|
|
},
|
|
PasswordResetKick: func(kickCtx context.Context, endpointID string) (bool, error) {
|
|
if _, found := brk.ConnInfoOf(endpointID); !found {
|
|
return false, nil
|
|
}
|
|
if resetErr := sess.ResetPassword(kickCtx, endpointID); resetErr != nil {
|
|
return false, resetErr
|
|
}
|
|
return true, nil
|
|
},
|
|
})
|
|
|
|
buildHandlers := func(proxies *listener.ProxySet) listener.Handlers {
|
|
return listener.Handlers{
|
|
MQTT: brk.WSHandler(proxies),
|
|
ClientAPI: idApp.Handler(),
|
|
AdminAPI: adminHandler,
|
|
Metrics: metricsReg.Handler(),
|
|
MetricsToken: cfg.Metrics.Token,
|
|
Static: staticFileHandler(),
|
|
Healthz: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
w.WriteHeader(http.StatusOK)
|
|
_, _ = w.Write([]byte("ok"))
|
|
}),
|
|
Readyz: http.HandlerFunc(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"))
|
|
}),
|
|
}
|
|
}
|
|
|
|
proxies, err := listener.ParseTrustedProxies(cfg.TrustedProxies)
|
|
if err != nil {
|
|
return fmt.Errorf("trusted_proxies: %w", err)
|
|
}
|
|
handlers := buildHandlers(proxies)
|
|
shared := cfg.AdminListen == ""
|
|
var clientHandler, adminHTTP http.Handler
|
|
if shared {
|
|
clientHandler = listener.NewMux(listener.RoleShared, handlers)
|
|
} else {
|
|
clientHandler = listener.NewMux(listener.RoleClient, handlers)
|
|
adminHTTP = listener.NewMux(listener.RoleAdmin, handlers)
|
|
}
|
|
|
|
lnOpts := listener.Options{
|
|
Listen: cfg.Listen,
|
|
AdminListen: cfg.AdminListen,
|
|
DataDir: cfg.DataDir,
|
|
CertFile: cfg.TLS.CertFile,
|
|
KeyFile: cfg.TLS.KeyFile,
|
|
AllowPlaintext: cfg.TLS.AllowPlaintext,
|
|
TrustedProxies: cfg.TrustedProxies,
|
|
ClientHandler: clientHandler,
|
|
AdminHandler: adminHTTP,
|
|
OnMQTT: func(conn net.Conn) {
|
|
_ = brk.AttachTCP(conn)
|
|
},
|
|
Logger: slog.Default(),
|
|
}
|
|
lnSrv, err := listener.New(lnOpts)
|
|
if err != nil {
|
|
return fmt.Errorf("listener: %w", err)
|
|
}
|
|
|
|
if err := lnSrv.Start(ctx); err != nil {
|
|
return err
|
|
}
|
|
defer func() { _ = lnSrv.Close() }()
|
|
|
|
if err := writeListenAddr(cfg.DataDir, lnSrv.ListenAddr()); err != nil {
|
|
return err
|
|
}
|
|
if cfg.AdminListen != "" && lnSrv.AdminAddr() != "" {
|
|
if err := os.WriteFile(filepath.Join(cfg.DataDir, "admin.addr"), []byte(lnSrv.AdminAddr()+"\n"), 0o644); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
|
|
loopCtx, loopCancel := context.WithCancel(context.Background())
|
|
defer loopCancel()
|
|
loopsDone := make(chan struct{})
|
|
go func() {
|
|
defer close(loopsDone)
|
|
messageLoops(loopCtx, msgApp, db, hashPool, metricsReg)
|
|
}()
|
|
|
|
<-ctx.Done()
|
|
invokeServeStop(ctx)
|
|
shutdownDeadline := time.Now().Add(30 * time.Second)
|
|
_ = lnSrv.StopAccept()
|
|
loopCancel()
|
|
<-loopsDone
|
|
|
|
drainBudget := 10 * time.Second
|
|
drainStart := time.Now()
|
|
drainCtx, drainCancel := context.WithTimeout(context.Background(), drainBudget)
|
|
if drainErr := db.Queue.Drain(drainCtx); drainErr != nil && !errors.Is(drainErr, context.DeadlineExceeded) {
|
|
slog.Error("write queue drain", "err", drainErr)
|
|
}
|
|
drainCancel()
|
|
|
|
shutCtx, shutCancel := context.WithTimeout(context.Background(), 5*time.Second)
|
|
if shutErr := brk.Shutdown(shutCtx); shutErr != nil && !errors.Is(shutErr, context.DeadlineExceeded) && !errors.Is(shutErr, context.Canceled) {
|
|
slog.Error("broker shutdown", "err", shutErr)
|
|
}
|
|
shutCancel()
|
|
|
|
secondDrain := drainBudget - time.Since(drainStart)
|
|
if secondDrain < time.Second {
|
|
secondDrain = time.Until(shutdownDeadline)
|
|
}
|
|
if secondDrain < time.Second {
|
|
secondDrain = time.Second
|
|
}
|
|
drain2, drain2Cancel := context.WithTimeout(context.Background(), secondDrain)
|
|
if drainErr := db.Queue.Drain(drain2); drainErr != nil && !errors.Is(drainErr, context.DeadlineExceeded) {
|
|
slog.Error("write queue drain", "err", drainErr)
|
|
}
|
|
drain2Cancel()
|
|
|
|
waitRemain := time.Until(shutdownDeadline)
|
|
if waitRemain < 2*time.Second {
|
|
waitRemain = 2 * time.Second
|
|
}
|
|
waitCtx, waitCancel := context.WithTimeout(context.Background(), waitRemain)
|
|
_ = lnSrv.Wait(waitCtx)
|
|
waitCancel()
|
|
return nil
|
|
}
|
|
|
|
func messageLoops(ctx context.Context, msgApp *message.App, db *store.DB, hashPool auth.HashPool, met *metrics.Registry) {
|
|
msgApp.StartLoops(ctx)
|
|
defer msgApp.WaitLoops()
|
|
t := time.NewTicker(15 * time.Second)
|
|
defer t.Stop()
|
|
sample := func() {
|
|
opCtx, cancel := context.WithTimeout(ctx, 5*time.Second)
|
|
defer cancel()
|
|
if err := metrics.SampleStoreGauges(opCtx, met, db.Read); err != nil {
|
|
slog.Debug("sample store gauges", "err", err)
|
|
}
|
|
metrics.SampleQueues(met, db.Queue.Len(), hashPool.QueueLen())
|
|
}
|
|
sample()
|
|
for {
|
|
select {
|
|
case <-ctx.Done():
|
|
return
|
|
case <-t.C:
|
|
sample()
|
|
}
|
|
}
|
|
}
|
|
|
|
func staticFileHandler() http.Handler {
|
|
root := web.Dist()
|
|
for _, prefix := range []string{"dist", "stub"} {
|
|
sub, err := fs.Sub(root, prefix)
|
|
if err != nil {
|
|
continue
|
|
}
|
|
if f, err := sub.Open("index.html"); err == nil {
|
|
_ = f.Close()
|
|
return httpx.SPA(sub)
|
|
}
|
|
}
|
|
return httpx.SPA(root)
|
|
}
|
|
|
|
func writeListenAddr(dataDir, addr string) error {
|
|
path := filepath.Join(dataDir, "listen.addr")
|
|
return os.WriteFile(path, []byte(addr+"\n"), 0o644)
|
|
}
|
|
|
|
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))
|
|
}
|