Skip to content

Commit 86db6a8

Browse files
committed
feat: Discover modules
1 parent 5b6e2ea commit 86db6a8

8 files changed

Lines changed: 406 additions & 10 deletions

File tree

‎assets/config.example.toml‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,10 @@ motd = "Served by rsync-proxy (https://github.com/ustclug/rsync-proxy)"
1717
address = "127.0.0.1:1234"
1818
modules = ["foo"]
1919

20+
[upstreams.u1_auto]
21+
address = "127.0.0.1:1234"
22+
discover_modules = true
23+
2024
[upstreams.u2]
2125
address = "192.168.0.10:1235"
2226
# Modules that multiple upstreams provide would be load-balanced by client IP.

‎pkg/server/config.go‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@ import (
1111
type Upstream struct {
1212
Address string `toml:"address"`
1313
Modules []string `toml:"modules"`
14+
DiscoverModules bool `toml:"discover_modules"`
1415
UseProxyProtocol bool `toml:"use_proxy_protocol"`
1516
}
1617

‎pkg/server/config_test.go‎

Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,9 +3,12 @@ package server
33
import (
44
"strings"
55
"testing"
6+
"time"
67

78
"github.com/stretchr/testify/assert"
89
"github.com/stretchr/testify/require"
10+
11+
"github.com/ustclug/rsync-proxy/test/fake/rsync"
912
)
1013

1114
func TestReadConfig(t *testing.T) {
@@ -103,3 +106,35 @@ modules = ["foo1"]
103106
require.Error(t, err, "load config")
104107
assert.Contains(t, err.Error(), "listen_tls requires tls_cert_file and tls_key_file", "unexpected error message")
105108
}
109+
110+
func TestReadConfigRequiresModulesOrDiscovery(t *testing.T) {
111+
s := New()
112+
configContent := `
113+
[upstreams.u1]
114+
address = "127.0.0.1:1234"
115+
`
116+
err := s.ReadConfig(strings.NewReader(configContent), true)
117+
require.Error(t, err, "load config")
118+
assert.Contains(t, err.Error(), "must set modules or discover_modules")
119+
}
120+
121+
func TestReadConfigDiscoversModules(t *testing.T) {
122+
upstream := rsync.NewModuleListServer([]string{"bar", "foo"})
123+
upstream.Start()
124+
defer upstream.Close()
125+
126+
s := New()
127+
s.ReadTimeout = time.Second
128+
s.WriteTimeout = time.Second
129+
configContent := `
130+
[upstreams.u1]
131+
address = "` + upstream.Listener.Addr().String() + `"
132+
discover_modules = true
133+
`
134+
err := s.ReadConfig(strings.NewReader(configContent), true)
135+
require.NoError(t, err, "load config")
136+
assert.Equal(t, map[string][]Target{
137+
"bar": {{Addr: upstream.Listener.Addr().String(), UseProxyProtocol: false}},
138+
"foo": {{Addr: upstream.Listener.Addr().String(), UseProxyProtocol: false}},
139+
}, s.modules)
140+
}

‎pkg/server/server.go‎

Lines changed: 224 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
package server
22

33
import (
4+
"bufio"
45
"bytes"
56
"context"
67
"crypto/tls"
@@ -24,7 +25,9 @@ import (
2425
)
2526

2627
const (
27-
TCPBufferSize = 256
28+
TCPBufferSize = 256
29+
moduleRetryInterval = 30 * time.Second
30+
moduleDiscoveryBufSize = 4096
2831
)
2932

3033
var (
@@ -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+
7989
type 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+
223431
func 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

594802
func (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

Comments
 (0)