main
go 206 lines 4.54 KB
Raw
1 //go:build windows
2
3 package service
4
5 import (
6 "context"
7 "errors"
8 "syscall"
9 "time"
10
11 "golang.org/x/sys/windows"
12 "golang.org/x/sys/windows/svc"
13 "golang.org/x/sys/windows/svc/mgr"
14 )
15
16 func Install(ctx context.Context, def Definition) error {
17 m, err := mgr.Connect()
18 if err != nil {
19 return err
20 }
21 defer m.Disconnect()
22
23 cfg := mgr.Config{
24 StartType: mgr.StartAutomatic,
25 ErrorControl: mgr.ErrorNormal,
26 DisplayName: def.DisplayName,
27 Description: def.Description,
28 DelayedAutoStart: true,
29 }
30 s, err := m.OpenService(def.Name)
31 if err == nil {
32 defer s.Close()
33 existing, cfgErr := s.Config()
34 if cfgErr != nil {
35 return cfgErr
36 }
37 cfg = existing
38 cfg.StartType = mgr.StartAutomatic
39 cfg.ErrorControl = mgr.ErrorNormal
40 cfg.DisplayName = def.DisplayName
41 cfg.Description = def.Description
42 cfg.DelayedAutoStart = true
43 cfg.BinaryPathName = windowsCommandLine(def)
44 if err := s.UpdateConfig(cfg); err != nil {
45 return err
46 }
47 return configureWindowsRecovery(s)
48 }
49 if !errors.Is(err, windows.ERROR_SERVICE_DOES_NOT_EXIST) {
50 return err
51 }
52 s, err = m.CreateService(def.Name, def.Executable, cfg, def.Args...)
53 if err != nil {
54 return err
55 }
56 defer s.Close()
57 if err := configureWindowsRecovery(s); err != nil {
58 return err
59 }
60 return ctx.Err()
61 }
62
63 func Start(ctx context.Context, name string) error {
64 s, err := openService(name)
65 if err != nil {
66 return err
67 }
68 defer s.Close()
69 err = s.Start()
70 if err != nil && !errors.Is(err, windows.ERROR_SERVICE_ALREADY_RUNNING) {
71 return err
72 }
73 return waitWindowsService(ctx, s, svc.Running)
74 }
75
76 func Stop(ctx context.Context, name string) error {
77 s, err := openService(name)
78 if err != nil {
79 if errors.Is(err, windows.ERROR_SERVICE_DOES_NOT_EXIST) {
80 return nil
81 }
82 return err
83 }
84 defer s.Close()
85
86 status, err := s.Query()
87 if err == nil && status.State != svc.Stopped {
88 _, _ = s.Control(svc.Stop)
89 return waitWindowsService(ctx, s, svc.Stopped)
90 }
91 return ctx.Err()
92 }
93
94 func StopDisable(ctx context.Context, name string) error {
95 s, err := openService(name)
96 if err != nil {
97 if errors.Is(err, windows.ERROR_SERVICE_DOES_NOT_EXIST) {
98 return nil
99 }
100 return err
101 }
102 defer s.Close()
103
104 cfg, err := s.Config()
105 if err == nil {
106 cfg.StartType = mgr.StartDisabled
107 _ = s.UpdateConfig(cfg)
108 }
109 status, err := s.Query()
110 if err == nil && status.State != svc.Stopped {
111 _, _ = s.Control(svc.Stop)
112 return waitWindowsService(ctx, s, svc.Stopped)
113 }
114 return ctx.Err()
115 }
116
117 func openService(name string) (*mgr.Service, error) {
118 m, err := mgr.Connect()
119 if err != nil {
120 return nil, err
121 }
122 s, err := m.OpenService(name)
123 _ = m.Disconnect()
124 return s, err
125 }
126
127 func waitWindowsService(ctx context.Context, s *mgr.Service, want svc.State) error {
128 ticker := time.NewTicker(300 * time.Millisecond)
129 defer ticker.Stop()
130 for {
131 status, err := s.Query()
132 if err != nil {
133 return err
134 }
135 if status.State == want {
136 return nil
137 }
138 select {
139 case <-ctx.Done():
140 return ctx.Err()
141 case <-ticker.C:
142 }
143 }
144 }
145
146 func windowsCommandLine(def Definition) string {
147 line := syscall.EscapeArg(def.Executable)
148 for _, arg := range def.Args {
149 line += " " + syscall.EscapeArg(arg)
150 }
151 return line
152 }
153
154 func configureWindowsRecovery(s *mgr.Service) error {
155 if err := s.SetRecoveryActions([]mgr.RecoveryAction{
156 {Type: mgr.ServiceRestart, Delay: 5 * time.Second},
157 {Type: mgr.ServiceRestart, Delay: 10 * time.Second},
158 {Type: mgr.ServiceRestart, Delay: 30 * time.Second},
159 }, 60); err != nil {
160 return err
161 }
162 return s.SetRecoveryActionsOnNonCrashFailures(false)
163 }
164
165 func Run(ctx context.Context, name string, run func(context.Context) error) error {
166 inService, err := svc.IsWindowsService()
167 if err != nil || !inService {
168 return run(ctx)
169 }
170 return svc.Run(name, windowsServiceHandler{ctx: ctx, run: run})
171 }
172
173 type windowsServiceHandler struct {
174 ctx context.Context
175 run func(context.Context) error
176 }
177
178 func (h windowsServiceHandler) Execute(args []string, requests <-chan svc.ChangeRequest, status chan<- svc.Status) (bool, uint32) {
179 ctx, cancel := context.WithCancel(h.ctx)
180 defer cancel()
181
182 errCh := make(chan error, 1)
183 status <- svc.Status{State: svc.StartPending}
184 go func() {
185 errCh <- h.run(ctx)
186 }()
187 status <- svc.Status{State: svc.Running, Accepts: svc.AcceptStop | svc.AcceptShutdown}
188
189 for {
190 select {
191 case req := <-requests:
192 switch req.Cmd {
193 case svc.Interrogate:
194 status <- req.CurrentStatus
195 case svc.Stop, svc.Shutdown:
196 status <- svc.Status{State: svc.StopPending}
197 cancel()
198 }
199 case err := <-errCh:
200 if err != nil && !errors.Is(err, context.Canceled) {
201 return false, 1
202 }
203 return false, 0
204 }
205 }
206 }