Files
yms-daemon/internal/daemonserver/server.go
T

215 lines
7.8 KiB
Go

// Package daemonserver accepts local CLI requests over the frozen 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"
)
const maximumRequestBytes = 1 << 20
type backendUpdater interface {
UpdateRepack(context.Context, string, backendupdate.ProgressReporter) (transaction.Transaction, error)
UpdateNativeJAR(context.Context, string, backendupdate.ProgressReporter) (transaction.Transaction, error)
UpdateContainerImage(context.Context, string, bool, backendupdate.ProgressReporter) (transaction.Transaction, error)
Restart(context.Context, backendupdate.ProgressReporter) (transaction.Transaction, error)
}
// Server owns the local Unix Socket and dispatches requests to the transaction orchestrator.
type Server struct {
socketPath string
updater backendUpdater
logger *slog.Logger
}
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 listens until ctx is canceled. Each accepted request owns its response connection.
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)
}()
}
}
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)
}
}
func (s *Server) writeResponse(connection net.Conn, response daemonapi.Response) error {
return json.NewEncoder(connection).Encode(response)
}
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
}
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
}