Files
yms-daemon/internal/daemonserver/server.go
T
2026-08-17 10:10:14 +08:00

245 lines
11 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// Package daemonserver 通过固定的本地 Unix Socket 接受命令行客户端请求,并将其分派给事务编排器执行。
// 服务端负责监听并校验请求、把请求映射到后端更新或重启动作、将执行进度流式写回客户端,
// 最终返回携带事务终态与错误信息的结果。
package daemonserver
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net"
"os"
"path/filepath"
"strings"
"sync"
"time"
"yms-daemon/internal/backendupdate"
"yms-daemon/internal/daemonapi"
"yms-daemon/internal/transaction"
)
// maximumRequestBytes 单次请求体的最大字节数,超过该上限的请求会被拒绝,以防内存被异常输入耗尽。
const maximumRequestBytes = 1 << 20
// backendUpdater 抽象了后端更新与重启的具体实现,是 Server 与事务编排器之间的解耦接口。
// 它由 backendupdate 包在运行时注入,便于在测试中用轻量替身替换。
type backendUpdater interface {
// UpdateRepack 使用重新打包的 ZIP 制品执行后端更新,file 为制品绝对路径。
UpdateRepack(context.Context, string, backendupdate.ProgressReporter) (transaction.Transaction, error)
// UpdateNativeJAR 使用原生后端 JAR 文件执行后端更新,file 为制品绝对路径。
UpdateNativeJAR(context.Context, string, backendupdate.ProgressReporter) (transaction.Transaction, error)
// UpdateContainerImage 使用容器镜像执行后端更新,imageReference 为镜像引用,startLog 控制是否输出启动日志。
UpdateContainerImage(context.Context, string, bool, backendupdate.ProgressReporter) (transaction.Transaction, error)
// Restart 重启后端服务而不更换制品。
Restart(context.Context, backendupdate.ProgressReporter) (transaction.Transaction, error)
}
// Server 持有本地 Unix Socket 并把收到的请求分派给事务编排器。
// 它负责连接的生命周期管理、请求解码与校验、执行进度回写以及最终结果的返回。
type Server struct {
// socketPath 守护进程监听的 Unix Socket 绝对路径。
socketPath string
// updater 后端更新与重启动作的实际执行者。
updater backendUpdater
// logger 用于记录服务运行与请求处理过程中的日志。
logger *slog.Logger
}
// New 构造一个 Server 实例并校验入参。
// socketPath 必须是规范化的绝对路径,updater 不能为空,logger 为空时回退为默认日志器。
// 当路径非法或 updater 为空时返回 error。
func New(socketPath string, updater backendUpdater, logger *slog.Logger) (*Server, error) {
if !filepath.IsAbs(socketPath) || filepath.Clean(socketPath) != socketPath {
return nil, errors.New("daemon Unix Socket path must be a clean absolute path")
}
if updater == nil {
return nil, errors.New("backend updater is required")
}
if logger == nil {
logger = slog.Default()
}
return &Server{socketPath: socketPath, updater: updater, logger: logger}, nil
}
// Serve 监听 Unix Socket 直到 ctx 被取消,返回前会等待所有连接处理完成并清理 Socket 文件。
// 它先准备 Socket 目录并移除残留的旧 Socket,然后以 0o600 权限监听;
// 每个被接受的连接都由独立的 goroutine 处理,连接在响应写完后由 handle 关闭。
// 当监听或准备 Socket 路径失败时返回 errorctx 取消后正常返回 nil。
func (s *Server) Serve(ctx context.Context) error {
if err := prepareSocketPath(s.socketPath); err != nil {
return err
}
listener, err := net.Listen("unix", s.socketPath)
if err != nil {
return fmt.Errorf("listen on daemon Unix Socket %s: %w", s.socketPath, err)
}
if err := os.Chmod(s.socketPath, 0o600); err != nil {
_ = listener.Close()
return fmt.Errorf("set daemon Unix Socket permissions: %w", err)
}
defer func() {
_ = listener.Close()
_ = os.Remove(s.socketPath)
}()
s.logger.InfoContext(ctx, "daemon server listening", "socket", s.socketPath)
var connections sync.WaitGroup
defer connections.Wait()
go func() {
<-ctx.Done()
_ = listener.Close()
}()
for {
connection, err := listener.Accept()
if err != nil {
if ctx.Err() != nil {
return nil
}
return fmt.Errorf("accept daemon Unix Socket connection: %w", err)
}
connections.Add(1)
go func() {
defer connections.Done()
s.handle(ctx, connection)
}()
}
}
// handle 处理单个已接受的连接:读取并校验请求、执行对应操作、把进度与结果写回客户端。
// 它限制请求读取时限为 10 秒,仅接受 service 为 backend 的请求;
// 对于 update 操作会按 inputType 分派到不同更新实现,对 restart 操作校验其不携带多余输入。
// 执行期间产生的进度通过 report 回调实时写回,一旦写回失败便停止后续进度发送,
// 最终无论成功失败都会写入一条 ResponseResult 响应并记录日志。
func (s *Server) handle(ctx context.Context, connection net.Conn) {
defer connection.Close()
_ = connection.SetReadDeadline(time.Now().Add(10 * time.Second))
request, err := decodeRequest(connection)
_ = connection.SetReadDeadline(time.Time{})
if err != nil {
_ = s.writeResponse(connection, daemonapi.Response{Kind: daemonapi.ResponseResult, Error: err.Error()})
return
}
if request.Service != "backend" {
_ = s.writeResponse(connection, daemonapi.Response{Kind: daemonapi.ResponseResult, Error: "service must be backend"})
return
}
progressWritable := true
report := func(progress backendupdate.Progress) {
if !progressWritable {
return
}
err := s.writeResponse(connection, daemonapi.Response{
Kind: daemonapi.ResponseProgress,
TransactionID: progress.TransactionID,
State: string(progress.State),
Message: progress.Message,
})
if err != nil {
progressWritable = false
s.logger.WarnContext(ctx, "write daemon update progress", "error", err)
}
}
var record transaction.Transaction
var updateErr error
switch request.Operation {
case daemonapi.OperationUpdate:
switch request.InputType {
case daemonapi.InputTypeRepackZIP:
if request.File == "" || request.ImageReference != "" {
_ = s.writeResponse(connection, daemonapi.Response{Kind: daemonapi.ResponseResult, Error: "repack-zip requires file and does not accept imageReference"})
return
}
record, updateErr = s.updater.UpdateRepack(ctx, request.File, report)
case daemonapi.InputTypeNativeJAR:
if request.File == "" || request.ImageReference != "" {
_ = s.writeResponse(connection, daemonapi.Response{Kind: daemonapi.ResponseResult, Error: "native-jar requires file and does not accept imageReference"})
return
}
record, updateErr = s.updater.UpdateNativeJAR(ctx, request.File, report)
case daemonapi.InputTypeContainerImage:
if request.File != "" || request.ImageReference == "" {
_ = s.writeResponse(connection, daemonapi.Response{Kind: daemonapi.ResponseResult, Error: "container-image requires imageReference and does not accept file"})
return
}
record, updateErr = s.updater.UpdateContainerImage(ctx, request.ImageReference, request.StartLog, report)
default:
_ = s.writeResponse(connection, daemonapi.Response{Kind: daemonapi.ResponseResult, Error: "inputType must be repack-zip, native-jar, or container-image"})
return
}
case daemonapi.OperationRestart:
if request.InputType != "" || request.File != "" || request.ImageReference != "" {
_ = s.writeResponse(connection, daemonapi.Response{Kind: daemonapi.ResponseResult, Error: "restart does not accept inputType or file"})
return
}
record, updateErr = s.updater.Restart(ctx, report)
default:
_ = s.writeResponse(connection, daemonapi.Response{Kind: daemonapi.ResponseResult, Error: "operation must be update or restart"})
return
}
response := daemonapi.Response{Kind: daemonapi.ResponseResult, TransactionID: record.ID, State: string(record.State)}
if updateErr != nil {
response.Error = updateErr.Error()
s.logger.ErrorContext(ctx, "backend operation failed", "operation", request.Operation, "transaction_id", record.ID, "state", record.State, "error", updateErr)
} else {
s.logger.InfoContext(ctx, "backend operation completed", "operation", request.Operation, "transaction_id", record.ID, "state", record.State)
}
if err := s.writeResponse(connection, response); err != nil {
s.logger.ErrorContext(ctx, "write daemon operation result", "operation", request.Operation, "transaction_id", record.ID, "error", err)
}
}
// writeResponse 把一条响应以换行分隔的 JSON 编码写入连接,失败时返回底层写入错误。
func (s *Server) writeResponse(connection net.Conn, response daemonapi.Response) error {
return json.NewEncoder(connection).Encode(response)
}
// decodeRequest 从 reader 中解码并校验一条客户端请求。
// 它限制请求体大小不超过 maximumRequestBytes,禁止未知字段,并要求请求只能包含一个 JSON 值,
// 且所有字符串字段不得带有前后空白字符;任一条不满足时返回 error。
func decodeRequest(reader io.Reader) (daemonapi.Request, error) {
limited := io.LimitReader(reader, maximumRequestBytes+1)
decoder := json.NewDecoder(limited)
decoder.DisallowUnknownFields()
var request daemonapi.Request
if err := decoder.Decode(&request); err != nil {
return daemonapi.Request{}, fmt.Errorf("decode daemon request: %w", err)
}
var trailing any
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
if err == nil {
return daemonapi.Request{}, errors.New("daemon request contains multiple JSON values")
}
return daemonapi.Request{}, fmt.Errorf("decode daemon request trailing content: %w", err)
}
if strings.TrimSpace(request.Operation) != request.Operation || strings.TrimSpace(request.Service) != request.Service || strings.TrimSpace(request.InputType) != request.InputType || strings.TrimSpace(request.File) != request.File || strings.TrimSpace(request.ImageReference) != request.ImageReference {
return daemonapi.Request{}, errors.New("daemon request fields must not contain surrounding whitespace")
}
return request, nil
}
// prepareSocketPath 在监听前准备 Socket 路径:创建父目录,若目标位置已被占用则校验其是否为 Socket,
// 若是残留 Socket 则移除,若被非 Socket 文件占用则返回 error。
func prepareSocketPath(socketPath string) error {
if err := os.MkdirAll(filepath.Dir(socketPath), 0o755); err != nil {
return fmt.Errorf("create daemon Unix Socket directory: %w", err)
}
info, err := os.Lstat(socketPath)
if errors.Is(err, os.ErrNotExist) {
return nil
}
if err != nil {
return fmt.Errorf("inspect daemon Unix Socket path: %w", err)
}
if info.Mode()&os.ModeSocket == 0 {
return fmt.Errorf("daemon Unix Socket path is occupied by a non-socket file: %s", socketPath)
}
if err := os.Remove(socketPath); err != nil {
return fmt.Errorf("remove stale daemon Unix Socket: %w", err)
}
return nil
}