diff --git a/cli/main.go b/cli/main.go index 0c1e645..be622e6 100644 --- a/cli/main.go +++ b/cli/main.go @@ -225,52 +225,68 @@ func runHook(c *cli.Command) { } } -func keepAliveLoop(c *cli.Command, campusOnly bool) (ret error) { - logger.Infof("Accessing websites periodically to keep you online") - - accessTarget := func(url string, ipv6 bool) (ret error) { - network := "tcp4" - if ipv6 { - network = "tcp6" - } - netClient := &http.Client{ - Timeout: time.Second * 10, - Transport: &http.Transport{ - DialContext: func(ctx context.Context, _network, addr string) (net.Conn, error) { - logger.Debugf("DialContext %s (%s)\n", addr, network) - myDial := &net.Dialer{ - Timeout: 6 * time.Second, - KeepAlive: 0, - FallbackDelay: -1, // disable RFC 6555 Fast Fallback - } - return myDial.DialContext(ctx, network, addr) - }, +func newKeepAliveClient(network string) *http.Client { + dialer := &net.Dialer{ + Timeout: 6 * time.Second, + KeepAlive: 0, + FallbackDelay: -1, // disable RFC 6555 Fast Fallback + } + return &http.Client{ + Timeout: 10 * time.Second, + Transport: &http.Transport{ + DialContext: func(ctx context.Context, _ string, addr string) (net.Conn, error) { + logger.Debugf("DialContext %s (%s)\n", addr, network) + return dialer.DialContext(ctx, network, addr) }, - } - resp, ret := netClient.Head(url) - if ret != nil { - return - } - defer resp.Body.Close() - logger.Debugf("HTTP status code %d\n", resp.StatusCode) - return + }, } - targetInside := "https://www.tsinghua.edu.cn/" - targetOutside := "https://www.baidu.com/" +} + +func accessKeepAliveTarget(client *http.Client, url string) error { + resp, err := client.Head(url) + if err != nil { + return err + } + defer resp.Body.Close() + logger.Debugf("HTTP status code %d\n", resp.StatusCode) + return nil +} - stop := make(chan int, 1) - defer func() { stop <- 1 }() +func startIPv6KeepAlive(stop <-chan struct{}, client *http.Client, url string, interval time.Duration) <-chan struct{} { + done := make(chan struct{}) go func() { - // Keep IPv6 online, ignore any errors + defer close(done) + timer := time.NewTimer(interval) + defer timer.Stop() for { select { case <-stop: - break - case <-time.After(13 * time.Minute): - _ = accessTarget(targetInside, true) + return + case <-timer.C: + // Keep IPv6 online, ignore any errors. + _ = accessKeepAliveTarget(client, url) + timer.Reset(interval) } } }() + return done +} + +func keepAliveLoop(c *cli.Command, campusOnly bool) (ret error) { + logger.Infof("Accessing websites periodically to keep you online") + client4 := newKeepAliveClient("tcp4") + client6 := newKeepAliveClient("tcp6") + targetInside := "https://www.tsinghua.edu.cn/" + targetOutside := "https://www.baidu.com/" + + stop := make(chan struct{}) + done := startIPv6KeepAlive(stop, client6, targetInside, 13*time.Minute) + defer func() { + close(stop) + <-done + client4.CloseIdleConnections() + client6.CloseIdleConnections() + }() errorCount := 0 for { @@ -278,7 +294,11 @@ func keepAliveLoop(c *cli.Command, campusOnly bool) (ret error) { if campusOnly || settings.V6 { target = targetInside } - if ret = accessTarget(target, settings.V6); ret != nil { + client := client4 + if settings.V6 { + client = client6 + } + if ret = accessKeepAliveTarget(client, target); ret != nil { errorCount++ if errorCount >= settings.OnRetry { ret = fmt.Errorf("keepAlive request error (re-login might be required): %w\n", ret)