11package server
22
33import (
4+ "bufio"
45 "bytes"
56 "context"
67 "crypto/tls"
@@ -24,7 +25,9 @@ import (
2425)
2526
2627const (
27- TCPBufferSize = 256
28+ TCPBufferSize = 256
29+ moduleRetryInterval = 30 * time .Second
30+ moduleDiscoveryBufSize = 4096
2831)
2932
3033var (
@@ -76,6 +79,13 @@ type Target struct {
7679 UseProxyProtocol bool
7780}
7881
82+ type upstreamConfig struct {
83+ Name string
84+ Target Target
85+ Modules []string
86+ DiscoverModules bool
87+ }
88+
7989type Server struct {
8090 // --- Options section
8191 // Listen Address
@@ -97,8 +107,14 @@ type Server struct {
97107 bufPool sync.Pool
98108 // name -> upstream targets
99109 modules map [string ][]Target
110+ upstreams []upstreamConfig
100111 tlsCertificate * tls.Certificate
101112 tlsConfig * tls.Config
113+ retryInterval time.Duration
114+
115+ discoveryMu sync.Mutex
116+ discoveryVersion uint64
117+ discoveryCancel context.CancelFunc
102118
103119 queue * queue.Queue
104120
@@ -135,10 +151,11 @@ func New() *Server {
135151 return & buf
136152 },
137153 },
138- dialer : net.Dialer {}, // customize keep alive interval?
139- accessLog : accessLog ,
140- errorLog : errorLog ,
141- queue : queue .New (0 , 0 ),
154+ dialer : net.Dialer {}, // customize keep alive interval?
155+ accessLog : accessLog ,
156+ errorLog : errorLog ,
157+ queue : queue .New (0 , 0 ),
158+ retryInterval : moduleRetryInterval ,
142159 }
143160 s .tlsConfig = & tls.Config {GetCertificate : s .getTLSCertificate }
144161 return s
@@ -176,25 +193,33 @@ func (s *Server) loadConfig(c *Config, openLog bool) error {
176193
177194 s .queue .SetMax (c .Proxy .MaxActiveConns , c .Proxy .MaxQueuedConns )
178195
179- modules := map [ string ][] Target {}
196+ upstreams := make ([] upstreamConfig , 0 , len ( c . Upstreams ))
180197 upstreamNames := make ([]string , 0 , len (c .Upstreams ))
181198 for upstreamName := range c .Upstreams {
182199 upstreamNames = append (upstreamNames , upstreamName )
183200 }
184201 sort .Strings (upstreamNames )
185202 for _ , upstreamName := range upstreamNames {
186203 v := c .Upstreams [upstreamName ]
204+ if len (v .Modules ) == 0 && ! v .DiscoverModules {
205+ return fmt .Errorf ("upstream=%s must set modules or discover_modules" , upstreamName )
206+ }
187207 addr := v .Address
188208 _ , err := net .ResolveTCPAddr ("tcp" , addr )
189209 if err != nil {
190210 return fmt .Errorf ("resolve address: %w, upstream=%s, address=%s" , err , upstreamName , addr )
191211 }
192- target := Target {Addr : addr , UseProxyProtocol : v .UseProxyProtocol }
193- for _ , moduleName := range v .Modules {
194- modules [moduleName ] = append (modules [moduleName ], target )
195- }
212+ upstreams = append (upstreams , upstreamConfig {
213+ Name : upstreamName ,
214+ Target : Target {Addr : addr , UseProxyProtocol : v .UseProxyProtocol },
215+ Modules : append ([]string (nil ), v .Modules ... ),
216+ DiscoverModules : v .DiscoverModules ,
217+ })
196218 }
197219
220+ discoveredModules , failed := s .discoverConfiguredModules (context .Background (), upstreams )
221+ modules := buildModuleTargets (upstreams , discoveredModules )
222+
198223 s .reloadLock .Lock ()
199224 defer s .reloadLock .Unlock ()
200225 if s .ListenAddr == "" {
@@ -216,10 +241,193 @@ func (s *Server) loadConfig(c *Config, openLog bool) error {
216241 }
217242 s .Motd = c .Proxy .Motd
218243 s .modules = modules
244+ s .upstreams = upstreams
219245 s .tlsCertificate = tlsCertificate
246+ s .restartModuleDiscovery (upstreams , discoveredModules , failed )
220247 return nil
221248}
222249
250+ func buildModuleTargets (upstreams []upstreamConfig , discovered map [string ][]string ) map [string ][]Target {
251+ modules := map [string ][]Target {}
252+ for _ , upstream := range upstreams {
253+ moduleNames := upstream .Modules
254+ if upstream .DiscoverModules {
255+ moduleNames = discovered [upstream .Name ]
256+ }
257+ for _ , moduleName := range moduleNames {
258+ modules [moduleName ] = append (modules [moduleName ], upstream .Target )
259+ }
260+ }
261+ return modules
262+ }
263+
264+ func (s * Server ) discoverConfiguredModules (ctx context.Context , upstreams []upstreamConfig ) (map [string ][]string , []upstreamConfig ) {
265+ discovered := map [string ][]string {}
266+ failed := make ([]upstreamConfig , 0 )
267+ for _ , upstream := range upstreams {
268+ if ! upstream .DiscoverModules {
269+ continue
270+ }
271+ modules , err := s .discoverModulesFromUpstream (ctx , upstream )
272+ if err != nil {
273+ s .logModuleDiscoveryFailure (upstream , err )
274+ failed = append (failed , upstream )
275+ continue
276+ }
277+ discovered [upstream .Name ] = modules
278+ s .logModuleDiscoverySuccess (upstream , modules )
279+ }
280+ return discovered , failed
281+ }
282+
283+ func (s * Server ) restartModuleDiscovery (upstreams []upstreamConfig , discovered map [string ][]string , pending []upstreamConfig ) {
284+ s .discoveryMu .Lock ()
285+ defer s .discoveryMu .Unlock ()
286+ s .discoveryVersion ++
287+ version := s .discoveryVersion
288+ if s .discoveryCancel != nil {
289+ s .discoveryCancel ()
290+ s .discoveryCancel = nil
291+ }
292+ if len (pending ) == 0 {
293+ return
294+ }
295+ ctx , cancel := context .WithCancel (context .Background ())
296+ s .discoveryCancel = cancel
297+ seed := cloneDiscoveredModules (discovered )
298+ allUpstreams := append ([]upstreamConfig (nil ), upstreams ... )
299+ failedUpstreams := append ([]upstreamConfig (nil ), pending ... )
300+ go s .retryDiscoverModules (ctx , version , allUpstreams , seed , failedUpstreams )
301+ }
302+
303+ func cloneDiscoveredModules (src map [string ][]string ) map [string ][]string {
304+ dup := make (map [string ][]string , len (src ))
305+ for name , modules := range src {
306+ dup [name ] = append ([]string (nil ), modules ... )
307+ }
308+ return dup
309+ }
310+
311+ func (s * Server ) retryDiscoverModules (ctx context.Context , version uint64 , upstreams []upstreamConfig , discovered map [string ][]string , pending []upstreamConfig ) {
312+ interval := s .retryInterval
313+ if interval <= 0 {
314+ interval = moduleRetryInterval
315+ }
316+ ticker := time .NewTicker (interval )
317+ defer ticker .Stop ()
318+
319+ for len (pending ) > 0 {
320+ select {
321+ case <- ctx .Done ():
322+ return
323+ case <- ticker .C :
324+ }
325+
326+ nextPending := pending [:0 ]
327+ updated := false
328+ for _ , upstream := range pending {
329+ modules , err := s .discoverModulesFromUpstream (ctx , upstream )
330+ if err != nil {
331+ s .logModuleDiscoveryFailure (upstream , err )
332+ nextPending = append (nextPending , upstream )
333+ continue
334+ }
335+ discovered [upstream .Name ] = modules
336+ updated = true
337+ s .logModuleDiscoverySuccess (upstream , modules )
338+ }
339+ pending = nextPending
340+ if updated {
341+ s .applyDiscoveredModules (version , upstreams , discovered )
342+ }
343+ }
344+
345+ s .discoveryMu .Lock ()
346+ if s .discoveryVersion == version {
347+ s .discoveryCancel = nil
348+ }
349+ s .discoveryMu .Unlock ()
350+ }
351+
352+ func (s * Server ) applyDiscoveredModules (version uint64 , upstreams []upstreamConfig , discovered map [string ][]string ) {
353+ s .discoveryMu .Lock ()
354+ if s .discoveryVersion != version {
355+ s .discoveryMu .Unlock ()
356+ return
357+ }
358+ s .discoveryMu .Unlock ()
359+
360+ modules := buildModuleTargets (upstreams , discovered )
361+
362+ s .reloadLock .Lock ()
363+ defer s .reloadLock .Unlock ()
364+ s .modules = modules
365+ s .upstreams = append ([]upstreamConfig (nil ), upstreams ... )
366+ }
367+
368+ func (s * Server ) discoverModulesFromUpstream (ctx context.Context , upstream upstreamConfig ) ([]string , error ) {
369+ conn , err := s .dialer .DialContext (ctx , "tcp" , upstream .Target .Addr )
370+ if err != nil {
371+ return nil , fmt .Errorf ("dial: %w" , err )
372+ }
373+ defer conn .Close ()
374+
375+ reader := bufio .NewReaderSize (conn , moduleDiscoveryBufSize )
376+ if _ , err := writeWithTimeout (conn , RsyncdServerVersion , s .WriteTimeout ); err != nil {
377+ return nil , fmt .Errorf ("send version: %w" , err )
378+ }
379+
380+ if s .ReadTimeout > 0 {
381+ _ = conn .SetReadDeadline (time .Now ().Add (s .ReadTimeout ))
382+ }
383+ line , err := reader .ReadString (lineFeed )
384+ if err != nil {
385+ return nil , fmt .Errorf ("read version: %w" , err )
386+ }
387+ if ! bytes .HasPrefix ([]byte (line ), RsyncdVersionPrefix ) {
388+ return nil , fmt .Errorf ("unexpected version response: %q" , line )
389+ }
390+
391+ if _ , err := writeWithTimeout (conn , []byte {'\n' }, s .WriteTimeout ); err != nil {
392+ return nil , fmt .Errorf ("request module list: %w" , err )
393+ }
394+
395+ modules := make ([]string , 0 )
396+ for {
397+ if s .ReadTimeout > 0 {
398+ _ = conn .SetReadDeadline (time .Now ().Add (s .ReadTimeout ))
399+ }
400+ line , err = reader .ReadString (lineFeed )
401+ if err != nil {
402+ return nil , fmt .Errorf ("read module list: %w" , err )
403+ }
404+ line = strings .TrimSuffix (line , string (lineFeed ))
405+ if line == strings .TrimSuffix (string (RsyncdExit ), string (lineFeed )) {
406+ break
407+ }
408+ if line == "" {
409+ continue
410+ }
411+ fields := strings .Fields (line )
412+ if len (fields ) == 0 {
413+ continue
414+ }
415+ modules = append (modules , fields [0 ])
416+ }
417+ sort .Strings (modules )
418+ return modules , nil
419+ }
420+
421+ func (s * Server ) logModuleDiscoveryFailure (upstream upstreamConfig , err error ) {
422+ log .Printf ("[WARN] discover modules from upstream %s (%s): %v" , upstream .Name , upstream .Target .Addr , err )
423+ s .errorLog .F ("[WARN] discover modules from upstream %s (%s): %v" , upstream .Name , upstream .Target .Addr , err )
424+ }
425+
426+ func (s * Server ) logModuleDiscoverySuccess (upstream upstreamConfig , modules []string ) {
427+ log .Printf ("[INFO] discovered modules from upstream %s (%s): %s" , upstream .Name , upstream .Target .Addr , strings .Join (modules , ", " ))
428+ s .errorLog .F ("[INFO] discovered modules from upstream %s (%s): %s" , upstream .Name , upstream .Target .Addr , strings .Join (modules , ", " ))
429+ }
430+
223431func chooseTargetByClientIP (ip net.IP , targetCount int ) int {
224432 if targetCount <= 1 {
225433 return 0
@@ -592,6 +800,12 @@ func (s *Server) Listen() error {
592800}
593801
594802func (s * Server ) Close () {
803+ s .discoveryMu .Lock ()
804+ if s .discoveryCancel != nil {
805+ s .discoveryCancel ()
806+ s .discoveryCancel = nil
807+ }
808+ s .discoveryMu .Unlock ()
595809 if s .TCPListener != nil {
596810 _ = s .TCPListener .Close ()
597811 }
0 commit comments