HH-635: let SSClash own Mihomo lifecycle (#6)
This commit was merged in pull request #6.
This commit is contained in:
+195
-34
@@ -1,6 +1,7 @@
|
||||
package bootstrap
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
@@ -18,8 +19,14 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
maxSubscriptionSize = 16 << 20
|
||||
minAdminPasswordLength = 12
|
||||
maxSubscriptionSize = 16 << 20
|
||||
minAdminPasswordLength = 12
|
||||
managedConfigVersion = "1"
|
||||
managedConfigVersionFile = ".mohomo-docker-config-version"
|
||||
managedChinaIPProvider = " ChinaIp: {type: file, behavior: ipcidr, format: yaml, path: /usr/local/share/ssclash/rules/ChinaIp.yaml}\n"
|
||||
managedChinaIPRule = " - RULE-SET,ChinaIp,🎯 全球直连"
|
||||
legacyChinaIPRule = " - GEOIP,CN,🎯 全球直连"
|
||||
defaultControllerURL = "http://127.0.0.1:9090"
|
||||
)
|
||||
|
||||
var errMihomoStateUncertain = errors.New("Mihomo subscription state could not be restored")
|
||||
@@ -55,15 +62,18 @@ type Config struct {
|
||||
type Result struct {
|
||||
CoreInitialized bool
|
||||
ConfigInitialized bool
|
||||
ConfigMigrated bool
|
||||
ServerSettingsChanged bool
|
||||
}
|
||||
|
||||
type RuntimeConfig struct {
|
||||
Root string
|
||||
CoreBinary string
|
||||
SSClashBinary string
|
||||
ConfigSource string
|
||||
RuntimeDir string
|
||||
SubscriptionURL string
|
||||
ControllerURL string
|
||||
UpdateInterval time.Duration
|
||||
}
|
||||
|
||||
@@ -103,9 +113,13 @@ func Prepare(config Config) (Result, error) {
|
||||
if err != nil {
|
||||
return result, fmt.Errorf("initialize Mihomo core: %w", err)
|
||||
}
|
||||
result.ConfigInitialized, err = copyIfAbsent(config.ConfigSource, filepath.Join(root, "config.yaml"), 0o644)
|
||||
result.ConfigInitialized, result.ConfigMigrated, err = prepareManagedConfig(
|
||||
config.ConfigSource,
|
||||
filepath.Join(root, "config.yaml"),
|
||||
filepath.Join(root, managedConfigVersionFile),
|
||||
)
|
||||
if err != nil {
|
||||
return result, fmt.Errorf("initialize config: %w", err)
|
||||
return result, fmt.Errorf("prepare config: %w", err)
|
||||
}
|
||||
result.ServerSettingsChanged, err = enforceServerSettings(filepath.Join(root, ".ssclash", "settings"))
|
||||
if err != nil {
|
||||
@@ -266,6 +280,10 @@ func Run(ctx context.Context, config RuntimeConfig) error {
|
||||
if config.UpdateInterval <= 0 {
|
||||
return errors.New("subscription update interval must be positive")
|
||||
}
|
||||
root := filepath.Clean(config.Root)
|
||||
if !filepath.IsAbs(root) || root == string(filepath.Separator) {
|
||||
return fmt.Errorf("unsafe root %q", config.Root)
|
||||
}
|
||||
runtimeDir := filepath.Clean(config.RuntimeDir)
|
||||
if !filepath.IsAbs(runtimeDir) || runtimeDir == string(filepath.Separator) {
|
||||
return fmt.Errorf("unsafe runtime directory %q", config.RuntimeDir)
|
||||
@@ -289,11 +307,15 @@ func Run(ctx context.Context, config RuntimeConfig) error {
|
||||
}
|
||||
activeSubscription := filepath.Join(runtimeDir, "subscription.yaml")
|
||||
client := &http.Client{Timeout: 30 * time.Second}
|
||||
controllerURL := strings.TrimRight(config.ControllerURL, "/")
|
||||
if controllerURL == "" {
|
||||
controllerURL = defaultControllerURL
|
||||
}
|
||||
validate := func(candidate string) error {
|
||||
return validateSubscription(ctx, config, runtimeConfig, candidate)
|
||||
}
|
||||
reload := func(ctx context.Context) error {
|
||||
return reloadSubscription(ctx, client)
|
||||
return reloadSubscription(ctx, client, controllerURL)
|
||||
}
|
||||
if err := updateSubscription(ctx, client, config.SubscriptionURL, activeSubscription, validate); err != nil {
|
||||
return fmt.Errorf("initial subscription update failed: %w", err)
|
||||
@@ -301,28 +323,20 @@ func Run(ctx context.Context, config RuntimeConfig) error {
|
||||
if err := validateMihomoConfig(ctx, config.CoreBinary, runtimeDir, runtimeConfig); err != nil {
|
||||
return errors.New("generated Mihomo configuration failed validation")
|
||||
}
|
||||
if err := ensureSubscriptionLink(filepath.Join(root, "subscription.yaml"), activeSubscription); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
serviceCtx, cancel := context.WithCancel(ctx)
|
||||
defer cancel()
|
||||
ssclash := serviceCommand(serviceCtx, config.SSClashBinary, "serve")
|
||||
mihomo := serviceCommand(serviceCtx, config.CoreBinary, "-d", runtimeDir, "-f", runtimeConfig)
|
||||
if err := ssclash.Start(); err != nil {
|
||||
return fmt.Errorf("start SSClash: %w", err)
|
||||
}
|
||||
if err := mihomo.Start(); err != nil {
|
||||
cancel()
|
||||
_ = ssclash.Wait()
|
||||
return fmt.Errorf("start Mihomo: %w", err)
|
||||
}
|
||||
log.Printf("bootstrap: services started mode=server subscription_update_interval=%s", config.UpdateInterval)
|
||||
log.Printf("bootstrap: SSClash started mode=server core_owner=ssclash subscription_update_interval=%s", config.UpdateInterval)
|
||||
|
||||
type processResult struct {
|
||||
name string
|
||||
err error
|
||||
}
|
||||
exits := make(chan processResult, 2)
|
||||
go func() { exits <- processResult{name: "SSClash", err: ssclash.Wait()} }()
|
||||
go func() { exits <- processResult{name: "Mihomo", err: mihomo.Wait()} }()
|
||||
exit := make(chan error, 1)
|
||||
go func() { exit <- ssclash.Wait() }()
|
||||
ticker := time.NewTicker(config.UpdateInterval)
|
||||
defer ticker.Stop()
|
||||
|
||||
@@ -330,32 +344,63 @@ func Run(ctx context.Context, config RuntimeConfig) error {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
cancel()
|
||||
<-exits
|
||||
<-exits
|
||||
<-exit
|
||||
return ctx.Err()
|
||||
case result := <-exits:
|
||||
cancel()
|
||||
<-exits
|
||||
if result.err == nil {
|
||||
return fmt.Errorf("%s exited", result.name)
|
||||
case err := <-exit:
|
||||
if err == nil {
|
||||
return errors.New("SSClash exited")
|
||||
}
|
||||
return fmt.Errorf("%s exited: %w", result.name, result.err)
|
||||
return fmt.Errorf("SSClash exited: %w", err)
|
||||
case <-ticker.C:
|
||||
if err := updateAndReload(ctx, client, config, activeSubscription, validate, reload); errors.Is(err, errMihomoStateUncertain) {
|
||||
log.Print("bootstrap: subscription rollback failed; stopping services")
|
||||
running := mihomoRunning(ctx, client, controllerURL)
|
||||
var err error
|
||||
if running {
|
||||
err = updateAndReload(ctx, client, config, activeSubscription, validate, reload)
|
||||
} else {
|
||||
err = updateSubscription(ctx, client, config.SubscriptionURL, activeSubscription, validate)
|
||||
}
|
||||
if errors.Is(err, errMihomoStateUncertain) {
|
||||
log.Print("bootstrap: subscription rollback failed; stopping SSClash")
|
||||
cancel()
|
||||
<-exits
|
||||
<-exits
|
||||
<-exit
|
||||
return err
|
||||
} else if err != nil {
|
||||
log.Print("bootstrap: subscription update rejected; keeping previous valid configuration")
|
||||
continue
|
||||
}
|
||||
log.Print("bootstrap: subscription updated and reloaded")
|
||||
if running {
|
||||
log.Print("bootstrap: subscription updated and reloaded")
|
||||
} else {
|
||||
log.Print("bootstrap: subscription updated; Mihomo is stopped")
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func ensureSubscriptionLink(path, target string) error {
|
||||
info, err := os.Lstat(path)
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
if err := os.Symlink(target, path); err != nil {
|
||||
return fmt.Errorf("create runtime subscription link: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
return fmt.Errorf("inspect runtime subscription link: %w", err)
|
||||
}
|
||||
if info.Mode()&os.ModeSymlink == 0 {
|
||||
return errors.New("runtime subscription path is not a managed symlink")
|
||||
}
|
||||
existingTarget, err := os.Readlink(path)
|
||||
if err != nil {
|
||||
return fmt.Errorf("read runtime subscription link: %w", err)
|
||||
}
|
||||
if existingTarget != target {
|
||||
return fmt.Errorf("runtime subscription link has unexpected target %q", existingTarget)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateSubscriptionURL(raw string) error {
|
||||
parsed, err := url.ParseRequestURI(raw)
|
||||
if err != nil || parsed.Host == "" || (parsed.Scheme != "https" && parsed.Scheme != "http") {
|
||||
@@ -461,8 +506,8 @@ func updateAndReload(ctx context.Context, client *http.Client, config RuntimeCon
|
||||
return errors.New("new subscription reload failed; previous subscription restored")
|
||||
}
|
||||
|
||||
func reloadSubscription(ctx context.Context, client *http.Client) error {
|
||||
request, err := http.NewRequestWithContext(ctx, http.MethodPut, "http://127.0.0.1:9090/providers/proxies/subscription", nil)
|
||||
func reloadSubscription(ctx context.Context, client *http.Client, controllerURL string) error {
|
||||
request, err := http.NewRequestWithContext(ctx, http.MethodPut, controllerURL+"/providers/proxies/subscription", nil)
|
||||
if err != nil {
|
||||
return errors.New("create Mihomo reload request")
|
||||
}
|
||||
@@ -478,6 +523,20 @@ func reloadSubscription(ctx context.Context, client *http.Client) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func mihomoRunning(ctx context.Context, client *http.Client, controllerURL string) bool {
|
||||
request, err := http.NewRequestWithContext(ctx, http.MethodGet, controllerURL+"/version", nil)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
response, err := client.Do(request)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
defer response.Body.Close()
|
||||
_, _ = io.Copy(io.Discard, response.Body)
|
||||
return true
|
||||
}
|
||||
|
||||
func serviceCommand(ctx context.Context, binary string, arguments ...string) *exec.Cmd {
|
||||
command := exec.CommandContext(ctx, binary, arguments...)
|
||||
command.Env = childEnvironment()
|
||||
@@ -558,6 +617,108 @@ func copyIfAbsent(source, target string, mode os.FileMode) (bool, error) {
|
||||
return err == nil, err
|
||||
}
|
||||
|
||||
func prepareManagedConfig(source, target, versionPath string) (bool, bool, error) {
|
||||
if err := validateManagedConfigVersion(versionPath); err != nil {
|
||||
return false, false, err
|
||||
}
|
||||
current, err := os.ReadFile(source)
|
||||
if err != nil {
|
||||
return false, false, fmt.Errorf("read managed config source: %w", err)
|
||||
}
|
||||
|
||||
info, err := os.Lstat(target)
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
if err := writeManagedConfig(target, current); err != nil {
|
||||
return false, false, err
|
||||
}
|
||||
if err := writeManagedConfigVersion(versionPath); err != nil {
|
||||
return false, false, err
|
||||
}
|
||||
return true, false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return false, false, fmt.Errorf("inspect config %q: %w", target, err)
|
||||
}
|
||||
if !info.Mode().IsRegular() {
|
||||
return false, false, fmt.Errorf("config %q is not a regular file", target)
|
||||
}
|
||||
if info.Size() == 0 {
|
||||
return false, false, fmt.Errorf("config %q is empty", target)
|
||||
}
|
||||
existing, err := os.ReadFile(target)
|
||||
if err != nil {
|
||||
return false, false, fmt.Errorf("read config %q: %w", target, err)
|
||||
}
|
||||
if bytes.Equal(existing, current) {
|
||||
return false, false, writeManagedConfigVersion(versionPath)
|
||||
}
|
||||
|
||||
legacy, legacyErr := legacyManagedConfig(current)
|
||||
if legacyErr == nil && bytes.Equal(existing, legacy) {
|
||||
if err := writeManagedConfig(target, current); err != nil {
|
||||
return false, false, err
|
||||
}
|
||||
if err := writeManagedConfigVersion(versionPath); err != nil {
|
||||
return false, true, err
|
||||
}
|
||||
return false, true, nil
|
||||
}
|
||||
if bytes.Contains(existing, []byte("GEOIP,CN")) {
|
||||
return false, false, fmt.Errorf("custom config uses GEOIP,CN and was preserved; replace it with the packaged local ChinaIp rule before retrying")
|
||||
}
|
||||
return false, false, nil
|
||||
}
|
||||
|
||||
func legacyManagedConfig(current []byte) ([]byte, error) {
|
||||
text := string(current)
|
||||
if strings.Count(text, managedChinaIPProvider) != 1 || strings.Count(text, managedChinaIPRule) != 1 {
|
||||
return nil, errors.New("packaged config is missing the managed ChinaIp rule")
|
||||
}
|
||||
text = strings.Replace(text, managedChinaIPProvider, "", 1)
|
||||
text = strings.Replace(text, managedChinaIPRule, legacyChinaIPRule, 1)
|
||||
return []byte(text), nil
|
||||
}
|
||||
|
||||
func validateManagedConfigVersion(path string) error {
|
||||
info, err := os.Lstat(path)
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
return fmt.Errorf("inspect managed config version: %w", err)
|
||||
}
|
||||
if !info.Mode().IsRegular() {
|
||||
return errors.New("managed config version marker is not a regular file")
|
||||
}
|
||||
content, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return fmt.Errorf("read managed config version: %w", err)
|
||||
}
|
||||
if string(content) != managedConfigVersion+"\n" {
|
||||
return fmt.Errorf("unsupported managed config version %q", strings.TrimSpace(string(content)))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func writeManagedConfig(path string, content []byte) error {
|
||||
return atomicWrite(path, 0o644, func(output *os.File) error {
|
||||
_, err := output.Write(content)
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
func writeManagedConfigVersion(path string) error {
|
||||
if _, err := os.Lstat(path); err == nil {
|
||||
return nil
|
||||
} else if !errors.Is(err, os.ErrNotExist) {
|
||||
return fmt.Errorf("inspect managed config version: %w", err)
|
||||
}
|
||||
return atomicWrite(path, 0o644, func(output *os.File) error {
|
||||
_, err := output.WriteString(managedConfigVersion + "\n")
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
func enforceServerSettings(path string) (bool, error) {
|
||||
content, err := os.ReadFile(path)
|
||||
if err != nil && !errors.Is(err, os.ErrNotExist) {
|
||||
|
||||
@@ -92,6 +92,61 @@ func TestPreparePreservesUserDataAndForcesServerMode(t *testing.T) {
|
||||
assertFileContent(t, filepath.Join(root, ".ssclash", "settings"), "LOG_LEVEL=debug\nOPERATING_MODE=server\nPROXY_MODE=none\n")
|
||||
}
|
||||
|
||||
func TestPrepareMigratesExactLegacyManagedConfig(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tempDir := t.TempDir()
|
||||
root := filepath.Join(tempDir, "data")
|
||||
if err := os.MkdirAll(root, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
current := "rule-providers:\n" + managedChinaIPProvider + "rules:\n" + managedChinaIPRule + "\n"
|
||||
legacy := "rule-providers:\nrules:\n" + legacyChinaIPRule + "\n"
|
||||
writeFixture(t, root, "config.yaml", legacy)
|
||||
|
||||
result, err := Prepare(Config{
|
||||
Root: root,
|
||||
SSClashTemp: filepath.Join(tempDir, "tmp"),
|
||||
CoreSource: writeFixture(t, tempDir, "mihomo", "core"),
|
||||
ConfigSource: writeFixture(t, tempDir, "current.yaml", current),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Prepare() error = %v", err)
|
||||
}
|
||||
if result.ConfigInitialized || !result.ConfigMigrated {
|
||||
t.Fatalf("Prepare() result = %+v, want migrated existing config", result)
|
||||
}
|
||||
assertFileContent(t, filepath.Join(root, "config.yaml"), current)
|
||||
assertFileContent(t, filepath.Join(root, managedConfigVersionFile), managedConfigVersion+"\n")
|
||||
}
|
||||
|
||||
func TestPreparePreservesAndRejectsCustomLegacyGeoIPConfig(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tempDir := t.TempDir()
|
||||
root := filepath.Join(tempDir, "data")
|
||||
if err := os.MkdirAll(root, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
current := "rule-providers:\n" + managedChinaIPProvider + "rules:\n" + managedChinaIPRule + "\n"
|
||||
custom := "rule-providers:\nrules:\n" + legacyChinaIPRule + "\n# user managed\n"
|
||||
target := writeFixture(t, root, "config.yaml", custom)
|
||||
|
||||
_, err := Prepare(Config{
|
||||
Root: root,
|
||||
SSClashTemp: filepath.Join(tempDir, "tmp"),
|
||||
CoreSource: writeFixture(t, tempDir, "mihomo", "core"),
|
||||
ConfigSource: writeFixture(t, tempDir, "current.yaml", current),
|
||||
})
|
||||
if err == nil || !strings.Contains(err.Error(), "custom config uses GEOIP,CN") {
|
||||
t.Fatalf("Prepare() error = %v, want explicit custom config migration error", err)
|
||||
}
|
||||
assertFileContent(t, target, custom)
|
||||
if _, statErr := os.Stat(filepath.Join(root, managedConfigVersionFile)); !errors.Is(statErr, os.ErrNotExist) {
|
||||
t.Fatalf("managed config version marker unexpectedly created: %v", statErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureAdminPasswordFailsClosedOnFreshVolume(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -479,47 +534,121 @@ func TestUpdateAndReloadRestoresRuntimeAfterAmbiguousFailure(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunStopsServicesWhenRollbackReloadFails(t *testing.T) {
|
||||
func TestRunLeavesMihomoLifecycleToSSClash(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
binary := writeFixture(t, tempDir, "fake-service", "#!/bin/sh\nif [ \"$1\" = -t ]; then exit 0; fi\nexec sleep 3600\n")
|
||||
if err := os.Chmod(binary, 0o755); err != nil {
|
||||
root := filepath.Join(tempDir, "root")
|
||||
if err := os.MkdirAll(root, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
config := writeFixture(t, tempDir, "config.yaml", "proxy-providers:\n subscription:\n type: file\n path: ./subscription.yaml\n")
|
||||
var requests atomic.Int32
|
||||
coreMarker := filepath.Join(tempDir, "core-started")
|
||||
ssclashMarker := filepath.Join(tempDir, "ssclash-started")
|
||||
t.Setenv("MIHOMO_TEST_MARKER", coreMarker)
|
||||
t.Setenv("SSCLASH_TEST_MARKER", ssclashMarker)
|
||||
core := writeFixture(t, tempDir, "fake-core", "#!/bin/sh\nif [ \"$1\" = -t ]; then exit 0; fi\ntouch \"$MIHOMO_TEST_MARKER\"\nexit 1\n")
|
||||
ssclash := writeFixture(t, tempDir, "fake-ssclash", "#!/bin/sh\n[ \"$1\" = serve ]\ntouch \"$SSCLASH_TEST_MARKER\"\nsleep 1\n")
|
||||
for _, binary := range []string{core, ssclash} {
|
||||
if err := os.Chmod(binary, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
config := writeFixture(t, root, "config.yaml", "proxy-providers:\n subscription:\n type: file\n path: ./subscription.yaml\n")
|
||||
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
|
||||
_, _ = writer.Write([]byte("proxies:\n - name: initial\n"))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
runtimeDir := filepath.Join(tempDir, "runtime")
|
||||
err := Run(context.Background(), RuntimeConfig{
|
||||
Root: root,
|
||||
CoreBinary: core,
|
||||
SSClashBinary: ssclash,
|
||||
ConfigSource: config,
|
||||
RuntimeDir: runtimeDir,
|
||||
SubscriptionURL: server.URL,
|
||||
UpdateInterval: time.Hour,
|
||||
})
|
||||
if err == nil || !strings.Contains(err.Error(), "SSClash exited") {
|
||||
t.Fatalf("Run() error = %v, want SSClash exit", err)
|
||||
}
|
||||
if _, err := os.Stat(ssclashMarker); err != nil {
|
||||
t.Fatalf("SSClash was not started: %v", err)
|
||||
}
|
||||
if _, err := os.Stat(coreMarker); !errors.Is(err, os.ErrNotExist) {
|
||||
t.Fatalf("bootstrap started Mihomo outside SSClash: %v", err)
|
||||
}
|
||||
linkTarget, err := os.Readlink(filepath.Join(root, "subscription.yaml"))
|
||||
if err != nil {
|
||||
t.Fatalf("read subscription link: %v", err)
|
||||
}
|
||||
if want := filepath.Join(runtimeDir, "subscription.yaml"); linkTarget != want {
|
||||
t.Fatalf("subscription link = %q, want %q", linkTarget, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunStopsSSClashWhenControllerCannotConfirmRollback(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
root := filepath.Join(tempDir, "root")
|
||||
if err := os.MkdirAll(root, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
core := writeFixture(t, tempDir, "fake-core", "#!/bin/sh\n[ \"$1\" = -t ]\n")
|
||||
ssclash := writeFixture(t, tempDir, "fake-ssclash", "#!/bin/sh\n[ \"$1\" = serve ]\nexec sleep 3600\n")
|
||||
for _, binary := range []string{core, ssclash} {
|
||||
if err := os.Chmod(binary, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
config := writeFixture(t, root, "config.yaml", "proxy-providers:\n subscription:\n type: file\n path: ./subscription.yaml\n")
|
||||
|
||||
var subscriptionRequests atomic.Int32
|
||||
subscription := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
|
||||
name := "updated"
|
||||
if requests.Add(1) == 1 {
|
||||
if subscriptionRequests.Add(1) == 1 {
|
||||
name = "initial"
|
||||
}
|
||||
_, _ = writer.Write([]byte("proxies:\n - name: " + name + "\n"))
|
||||
}))
|
||||
defer server.Close()
|
||||
defer subscription.Close()
|
||||
var reloadRequests atomic.Int32
|
||||
controller := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
|
||||
switch {
|
||||
case request.Method == http.MethodGet && request.URL.Path == "/version":
|
||||
writer.WriteHeader(http.StatusOK)
|
||||
case request.Method == http.MethodPut && request.URL.Path == "/providers/proxies/subscription":
|
||||
reloadRequests.Add(1)
|
||||
http.Error(writer, "reload failed", http.StatusInternalServerError)
|
||||
default:
|
||||
http.NotFound(writer, request)
|
||||
}
|
||||
}))
|
||||
defer controller.Close()
|
||||
|
||||
runResult := make(chan error, 1)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
go func() {
|
||||
runResult <- Run(ctx, RuntimeConfig{
|
||||
CoreBinary: binary,
|
||||
SSClashBinary: binary,
|
||||
Root: root,
|
||||
CoreBinary: core,
|
||||
SSClashBinary: ssclash,
|
||||
ConfigSource: config,
|
||||
RuntimeDir: filepath.Join(tempDir, "runtime"),
|
||||
SubscriptionURL: server.URL,
|
||||
SubscriptionURL: subscription.URL,
|
||||
ControllerURL: controller.URL,
|
||||
UpdateInterval: 20 * time.Millisecond,
|
||||
})
|
||||
}()
|
||||
var err error
|
||||
|
||||
select {
|
||||
case err = <-runResult:
|
||||
case err := <-runResult:
|
||||
if !errors.Is(err, errMihomoStateUncertain) {
|
||||
t.Fatalf("Run() error = %v, want uncertain Mihomo state", err)
|
||||
}
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("Run() did not reach rollback failure")
|
||||
t.Fatal("Run() did not stop SSClash after rollback reload failure")
|
||||
}
|
||||
if !errors.Is(err, errMihomoStateUncertain) {
|
||||
t.Fatalf("Run() error = %v, want uncertain Mihomo state", err)
|
||||
}
|
||||
if requests.Load() < 2 {
|
||||
t.Fatalf("subscription requests = %d, want initial fetch and timed update", requests.Load())
|
||||
if subscriptionRequests.Load() < 2 || reloadRequests.Load() != 2 {
|
||||
t.Fatalf("requests = subscription:%d reload:%d, want at least 2 and exactly 2", subscriptionRequests.Load(), reloadRequests.Load())
|
||||
}
|
||||
assertFileContent(t, filepath.Join(tempDir, "runtime", "subscription.yaml"), "proxies:\n - name: initial\n")
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user