From 99e19e70330585ebfcfaf3e72f994cc9f39fc55f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Tue, 17 Mar 2026 20:47:42 +0800 Subject: [PATCH] service: stop retrying fatal watch status errors --- service/ccm/credential_external.go | 9 + service/ccm/credential_status_test.go | 265 -------------------------- service/ocm/credential_external.go | 9 + service/ocm/credential_status_test.go | 246 ------------------------ 4 files changed, 18 insertions(+), 511 deletions(-) delete mode 100644 service/ccm/credential_status_test.go delete mode 100644 service/ocm/credential_status_test.go diff --git a/service/ccm/credential_external.go b/service/ccm/credential_external.go index 2445a8509..ba42ad64e 100644 --- a/service/ccm/credential_external.go +++ b/service/ccm/credential_external.go @@ -5,6 +5,7 @@ import ( "context" stdTLS "crypto/tls" "encoding/json" + "errors" "io" "net" "net/http" @@ -677,6 +678,10 @@ func (c *externalCredential) statusStreamLoop() { if ctx.Err() != nil { return } + if !shouldRetryStatusStreamError(err) { + c.logger.Warn("status stream for ", c.tag, " disconnected: ", err, ", not retrying") + return + } var backoff time.Duration consecutiveFailures, backoff = c.nextStatusStreamBackoff(result, consecutiveFailures) c.logger.Debug("status stream for ", c.tag, " disconnected: ", err, ", reconnecting in ", backoff) @@ -760,6 +765,10 @@ func (c *externalCredential) connectStatusStream(ctx context.Context) (statusStr } } +func shouldRetryStatusStreamError(err error) bool { + return errors.Is(err, io.ErrUnexpectedEOF) || E.IsClosedOrCanceled(err) +} + func (c *externalCredential) nextStatusStreamBackoff(result statusStreamResult, consecutiveFailures int) (int, time.Duration) { if result.duration >= connectorBackoffResetThreshold { consecutiveFailures = 0 diff --git a/service/ccm/credential_status_test.go b/service/ccm/credential_status_test.go deleted file mode 100644 index 9353f1d83..000000000 --- a/service/ccm/credential_status_test.go +++ /dev/null @@ -1,265 +0,0 @@ -package ccm - -import ( - "context" - "errors" - "io" - "net" - "net/http" - "os" - "path/filepath" - "strings" - "testing" - "time" - - "github.com/sagernet/sing-box/log" - "github.com/sagernet/sing/common/observable" - - "github.com/hashicorp/yamux" -) - -type roundTripperFunc func(*http.Request) (*http.Response, error) - -func (f roundTripperFunc) RoundTrip(request *http.Request) (*http.Response, error) { - return f(request) -} - -func drainStatusEvents(subscription observable.Subscription[struct{}]) int { - var count int - for { - select { - case <-subscription: - count++ - default: - return count - } - } -} - -func newTestLogger() log.ContextLogger { - return log.NewNOPFactory().Logger() -} - -func newTestCCMExternalCredential(t *testing.T, body string, headers http.Header) (*externalCredential, observable.Subscription[struct{}]) { - t.Helper() - subscriber := observable.NewSubscriber[struct{}](8) - subscription, _ := subscriber.Subscription() - credential := &externalCredential{ - tag: "test", - baseURL: "http://example.com", - token: "token", - pollInterval: 25 * time.Millisecond, - forwardHTTPClient: &http.Client{Transport: roundTripperFunc(func(request *http.Request) (*http.Response, error) { - if request.URL.String() != "http://example.com/ccm/v1/status?watch=true" { - t.Fatalf("unexpected request URL: %s", request.URL.String()) - } - return &http.Response{ - StatusCode: http.StatusOK, - Header: headers.Clone(), - Body: io.NopCloser(strings.NewReader(body)), - }, nil - })}, - logger: newTestLogger(), - statusSubscriber: subscriber, - } - return credential, subscription -} - -func newTestYamuxSessionPair(t *testing.T) (*yamux.Session, *yamux.Session) { - t.Helper() - clientConn, serverConn := net.Pipe() - clientSession, err := yamux.Client(clientConn, defaultYamuxConfig) - if err != nil { - t.Fatalf("create yamux client: %v", err) - } - serverSession, err := yamux.Server(serverConn, defaultYamuxConfig) - if err != nil { - clientSession.Close() - t.Fatalf("create yamux server: %v", err) - } - t.Cleanup(func() { - clientSession.Close() - serverSession.Close() - }) - return clientSession, serverSession -} - -func TestExternalCredentialConnectStatusStreamSingleFrameStreamReconnects(t *testing.T) { - credential, subscription := newTestCCMExternalCredential(t, "{\"five_hour_utilization\":12,\"weekly_utilization\":34,\"plan_weight\":2}\n", nil) - oldTime := time.Unix(123, 0) - credential.stateAccess.Lock() - credential.state.lastUpdated = oldTime - credential.stateAccess.Unlock() - - result, err := credential.connectStatusStream(context.Background()) - if err != io.EOF { - t.Fatalf("expected EOF, got %v", err) - } - if result.frames != 1 { - t.Fatalf("expected 1 frame, got %d", result.frames) - } - if credential.lastUpdatedTime().Equal(oldTime) { - t.Fatal("expected lastUpdated to remain refreshed") - } - if credential.fiveHourUtilization() != 12 || credential.weeklyUtilization() != 34 { - t.Fatalf("unexpected utilizations: 5h=%v weekly=%v", credential.fiveHourUtilization(), credential.weeklyUtilization()) - } - if count := drainStatusEvents(subscription); count != 1 { - t.Fatalf("expected 1 status event, got %d", count) - } - - failures, backoff := credential.nextStatusStreamBackoff(result, 3) - if failures != 4 { - t.Fatalf("expected failures incremented to 4, got %d", failures) - } - if backoff < 16*time.Second || backoff >= 24*time.Second { - t.Fatalf("expected connector backoff in [16s, 24s), got %v", backoff) - } -} - -func TestExternalCredentialConnectStatusStreamMultiFrameKeepsLastUpdated(t *testing.T) { - credential, subscription := newTestCCMExternalCredential(t, strings.Join([]string{ - "{\"five_hour_utilization\":12,\"weekly_utilization\":34,\"plan_weight\":2}", - "{\"five_hour_utilization\":13,\"weekly_utilization\":35,\"plan_weight\":3}", - }, "\n"), nil) - oldTime := time.Unix(123, 0) - credential.stateAccess.Lock() - credential.state.lastUpdated = oldTime - credential.stateAccess.Unlock() - - result, err := credential.connectStatusStream(context.Background()) - if err != io.EOF { - t.Fatalf("expected EOF, got %v", err) - } - if result.frames != 2 { - t.Fatalf("expected 2 frames, got %d", result.frames) - } - if credential.lastUpdatedTime().Equal(oldTime) { - t.Fatal("expected lastUpdated to remain refreshed") - } - if credential.fiveHourUtilization() != 13 || credential.weeklyUtilization() != 35 { - t.Fatalf("unexpected utilizations: 5h=%v weekly=%v", credential.fiveHourUtilization(), credential.weeklyUtilization()) - } - if count := drainStatusEvents(subscription); count != 2 { - t.Fatalf("expected 2 status events, got %d", count) - } -} - -func TestExternalCredentialPlanWeightOnlyHeaderEmitsStatus(t *testing.T) { - subscriber := observable.NewSubscriber[struct{}](8) - subscription, _ := subscriber.Subscription() - credential := &externalCredential{ - tag: "test", - logger: newTestLogger(), - statusSubscriber: subscriber, - } - credential.stateAccess.Lock() - credential.state.remotePlanWeight = 2 - oldTime := time.Unix(123, 0) - credential.state.lastUpdated = oldTime - credential.stateAccess.Unlock() - - headers := make(http.Header) - headers.Set("X-CCM-Plan-Weight", "3") - credential.updateStateFromHeaders(headers) - - if weight := credential.planWeight(); weight != 3 { - t.Fatalf("expected plan weight 3, got %v", weight) - } - if count := drainStatusEvents(subscription); count != 1 { - t.Fatalf("expected 1 status event, got %d", count) - } - if !credential.lastUpdatedTime().Equal(oldTime) { - t.Fatalf("expected lastUpdated to stay %v, got %v", oldTime, credential.lastUpdatedTime()) - } - - credential.updateStateFromHeaders(headers) - - if count := drainStatusEvents(subscription); count != 0 { - t.Fatalf("expected no status event for unchanged plan weight, got %d", count) - } -} - -func TestDefaultCredentialStatusChangesEmitStatus(t *testing.T) { - credentialPath := filepath.Join(t.TempDir(), "credentials.json") - err := os.WriteFile(credentialPath, []byte("{\"claudeAiOauth\":{\"accessToken\":\"token\",\"refreshToken\":\"\",\"expiresAt\":0,\"subscriptionType\":\"max\"}}\n"), 0o600) - if err != nil { - t.Fatalf("write credential file: %v", err) - } - - subscriber := observable.NewSubscriber[struct{}](8) - subscription, _ := subscriber.Subscription() - credential := &defaultCredential{ - tag: "test", - credentialPath: credentialPath, - logger: newTestLogger(), - statusSubscriber: subscriber, - } - - err = credential.markCredentialsUnavailable(errors.New("boom")) - if err == nil { - t.Fatal("expected error from markCredentialsUnavailable") - } - if count := drainStatusEvents(subscription); count != 1 { - t.Fatalf("expected 1 status event after unavailable transition, got %d", count) - } - - err = credential.reloadCredentials(true) - if err != nil { - t.Fatalf("reload credentials: %v", err) - } - if count := drainStatusEvents(subscription); count != 1 { - t.Fatalf("expected 1 status event after recovery, got %d", count) - } - if weight := credential.planWeight(); weight != 5 { - t.Fatalf("expected initial max weight 5, got %v", weight) - } - - profileClient := &http.Client{Transport: roundTripperFunc(func(request *http.Request) (*http.Response, error) { - return &http.Response{ - StatusCode: http.StatusOK, - Header: make(http.Header), - Body: io.NopCloser(strings.NewReader( - "{\"organization\":{\"organization_type\":\"claude_max\",\"rate_limit_tier\":\"default_claude_max_20x\"}}", - )), - }, nil - })} - credential.fetchProfile(context.Background(), profileClient, "token") - if count := drainStatusEvents(subscription); count != 1 { - t.Fatalf("expected 1 status event after weight change, got %d", count) - } - if weight := credential.planWeight(); weight != 10 { - t.Fatalf("expected upgraded max weight 10, got %v", weight) - } -} - -func TestExternalCredentialReverseSessionChangesEmitStatus(t *testing.T) { - subscriber := observable.NewSubscriber[struct{}](8) - subscription, _ := subscriber.Subscription() - credential := &externalCredential{ - tag: "receiver", - baseURL: reverseProxyBaseURL, - pollInterval: time.Minute, - logger: newTestLogger(), - statusSubscriber: subscriber, - } - - clientSession, _ := newTestYamuxSessionPair(t) - if !credential.setReverseSession(clientSession) { - t.Fatal("expected reverse session to be accepted") - } - if count := drainStatusEvents(subscription); count != 1 { - t.Fatalf("expected 1 status event after reverse session up, got %d", count) - } - if !credential.isAvailable() { - t.Fatal("expected receiver credential to become available") - } - - credential.clearReverseSession(clientSession) - if count := drainStatusEvents(subscription); count != 1 { - t.Fatalf("expected 1 status event after reverse session down, got %d", count) - } - if credential.isAvailable() { - t.Fatal("expected receiver credential to become unavailable") - } -} diff --git a/service/ocm/credential_external.go b/service/ocm/credential_external.go index 27342aa72..b06171d03 100644 --- a/service/ocm/credential_external.go +++ b/service/ocm/credential_external.go @@ -5,6 +5,7 @@ import ( "context" stdTLS "crypto/tls" "encoding/json" + "errors" "io" "net" "net/http" @@ -719,6 +720,10 @@ func (c *externalCredential) statusStreamLoop() { if ctx.Err() != nil { return } + if !shouldRetryStatusStreamError(err) { + c.logger.Warn("status stream for ", c.tag, " disconnected: ", err, ", not retrying") + return + } var backoff time.Duration consecutiveFailures, backoff = c.nextStatusStreamBackoff(result, consecutiveFailures) c.logger.Debug("status stream for ", c.tag, " disconnected: ", err, ", reconnecting in ", backoff) @@ -802,6 +807,10 @@ func (c *externalCredential) connectStatusStream(ctx context.Context) (statusStr } } +func shouldRetryStatusStreamError(err error) bool { + return errors.Is(err, io.ErrUnexpectedEOF) || E.IsClosedOrCanceled(err) +} + func (c *externalCredential) nextStatusStreamBackoff(result statusStreamResult, consecutiveFailures int) (int, time.Duration) { if result.duration >= connectorBackoffResetThreshold { consecutiveFailures = 0 diff --git a/service/ocm/credential_status_test.go b/service/ocm/credential_status_test.go deleted file mode 100644 index 2865a2380..000000000 --- a/service/ocm/credential_status_test.go +++ /dev/null @@ -1,246 +0,0 @@ -package ocm - -import ( - "context" - "errors" - "io" - "net" - "net/http" - "os" - "path/filepath" - "strings" - "testing" - "time" - - "github.com/sagernet/sing-box/log" - "github.com/sagernet/sing/common/observable" - - "github.com/hashicorp/yamux" -) - -type roundTripperFunc func(*http.Request) (*http.Response, error) - -func (f roundTripperFunc) RoundTrip(request *http.Request) (*http.Response, error) { - return f(request) -} - -func drainStatusEvents(subscription observable.Subscription[struct{}]) int { - var count int - for { - select { - case <-subscription: - count++ - default: - return count - } - } -} - -func newTestLogger() log.ContextLogger { - return log.NewNOPFactory().Logger() -} - -func newTestOCMExternalCredential(t *testing.T, body string, headers http.Header) (*externalCredential, observable.Subscription[struct{}]) { - t.Helper() - subscriber := observable.NewSubscriber[struct{}](8) - subscription, _ := subscriber.Subscription() - credential := &externalCredential{ - tag: "test", - baseURL: "http://example.com", - token: "token", - pollInterval: 25 * time.Millisecond, - forwardHTTPClient: &http.Client{Transport: roundTripperFunc(func(request *http.Request) (*http.Response, error) { - if request.URL.String() != "http://example.com/ocm/v1/status?watch=true" { - t.Fatalf("unexpected request URL: %s", request.URL.String()) - } - return &http.Response{ - StatusCode: http.StatusOK, - Header: headers.Clone(), - Body: io.NopCloser(strings.NewReader(body)), - }, nil - })}, - logger: newTestLogger(), - statusSubscriber: subscriber, - } - return credential, subscription -} - -func newTestYamuxSessionPair(t *testing.T) (*yamux.Session, *yamux.Session) { - t.Helper() - clientConn, serverConn := net.Pipe() - clientSession, err := yamux.Client(clientConn, defaultYamuxConfig) - if err != nil { - t.Fatalf("create yamux client: %v", err) - } - serverSession, err := yamux.Server(serverConn, defaultYamuxConfig) - if err != nil { - clientSession.Close() - t.Fatalf("create yamux server: %v", err) - } - t.Cleanup(func() { - clientSession.Close() - serverSession.Close() - }) - return clientSession, serverSession -} - -func TestExternalCredentialConnectStatusStreamSingleFrameStreamReconnects(t *testing.T) { - credential, subscription := newTestOCMExternalCredential(t, "{\"five_hour_utilization\":12,\"weekly_utilization\":34,\"plan_weight\":2}\n", nil) - oldTime := time.Unix(123, 0) - credential.stateAccess.Lock() - credential.state.lastUpdated = oldTime - credential.stateAccess.Unlock() - - result, err := credential.connectStatusStream(context.Background()) - if err != io.EOF { - t.Fatalf("expected EOF, got %v", err) - } - if result.frames != 1 { - t.Fatalf("expected 1 frame, got %d", result.frames) - } - if credential.lastUpdatedTime().Equal(oldTime) { - t.Fatal("expected lastUpdated to remain refreshed") - } - if credential.fiveHourUtilization() != 12 || credential.weeklyUtilization() != 34 { - t.Fatalf("unexpected utilizations: 5h=%v weekly=%v", credential.fiveHourUtilization(), credential.weeklyUtilization()) - } - if count := drainStatusEvents(subscription); count != 1 { - t.Fatalf("expected 1 status event, got %d", count) - } - - failures, backoff := credential.nextStatusStreamBackoff(result, 3) - if failures != 4 { - t.Fatalf("expected failures incremented to 4, got %d", failures) - } - if backoff < 16*time.Second || backoff >= 24*time.Second { - t.Fatalf("expected connector backoff in [16s, 24s), got %v", backoff) - } -} - -func TestExternalCredentialConnectStatusStreamMultiFrameKeepsLastUpdated(t *testing.T) { - credential, subscription := newTestOCMExternalCredential(t, strings.Join([]string{ - "{\"five_hour_utilization\":12,\"weekly_utilization\":34,\"plan_weight\":2}", - "{\"five_hour_utilization\":13,\"weekly_utilization\":35,\"plan_weight\":3}", - }, "\n"), nil) - oldTime := time.Unix(123, 0) - credential.stateAccess.Lock() - credential.state.lastUpdated = oldTime - credential.stateAccess.Unlock() - - result, err := credential.connectStatusStream(context.Background()) - if err != io.EOF { - t.Fatalf("expected EOF, got %v", err) - } - if result.frames != 2 { - t.Fatalf("expected 2 frames, got %d", result.frames) - } - if credential.lastUpdatedTime().Equal(oldTime) { - t.Fatal("expected lastUpdated to remain refreshed") - } - if credential.fiveHourUtilization() != 13 || credential.weeklyUtilization() != 35 { - t.Fatalf("unexpected utilizations: 5h=%v weekly=%v", credential.fiveHourUtilization(), credential.weeklyUtilization()) - } - if count := drainStatusEvents(subscription); count != 2 { - t.Fatalf("expected 2 status events, got %d", count) - } -} - -func TestExternalCredentialPlanWeightOnlyRateLimitsEventEmitsStatus(t *testing.T) { - subscriber := observable.NewSubscriber[struct{}](8) - subscription, _ := subscriber.Subscription() - credential := &externalCredential{ - tag: "test", - logger: newTestLogger(), - statusSubscriber: subscriber, - } - credential.stateAccess.Lock() - credential.state.remotePlanWeight = 2 - oldTime := time.Unix(123, 0) - credential.state.lastUpdated = oldTime - credential.stateAccess.Unlock() - - (&Service{}).handleWebSocketRateLimitsEvent([]byte(`{"plan_weight":3}`), credential) - - if weight := credential.planWeight(); weight != 3 { - t.Fatalf("expected plan weight 3, got %v", weight) - } - if count := drainStatusEvents(subscription); count != 1 { - t.Fatalf("expected 1 status event, got %d", count) - } - if !credential.lastUpdatedTime().Equal(oldTime) { - t.Fatalf("expected lastUpdated to stay %v, got %v", oldTime, credential.lastUpdatedTime()) - } - - (&Service{}).handleWebSocketRateLimitsEvent([]byte(`{"plan_weight":3}`), credential) - - if count := drainStatusEvents(subscription); count != 0 { - t.Fatalf("expected no status event for unchanged plan weight, got %d", count) - } -} - -func TestDefaultCredentialAvailabilityChangesEmitStatus(t *testing.T) { - credentialPath := filepath.Join(t.TempDir(), "auth.json") - err := os.WriteFile(credentialPath, []byte("{\"OPENAI_API_KEY\":\"sk-test\"}\n"), 0o600) - if err != nil { - t.Fatalf("write credential file: %v", err) - } - - subscriber := observable.NewSubscriber[struct{}](8) - subscription, _ := subscriber.Subscription() - credential := &defaultCredential{ - tag: "test", - credentialPath: credentialPath, - logger: newTestLogger(), - statusSubscriber: subscriber, - } - - err = credential.markCredentialsUnavailable(errors.New("boom")) - if err == nil { - t.Fatal("expected error from markCredentialsUnavailable") - } - if count := drainStatusEvents(subscription); count != 1 { - t.Fatalf("expected 1 status event after unavailable transition, got %d", count) - } - - err = credential.reloadCredentials(true) - if err != nil { - t.Fatalf("reload credentials: %v", err) - } - if count := drainStatusEvents(subscription); count != 1 { - t.Fatalf("expected 1 status event after recovery, got %d", count) - } - if !credential.isAvailable() { - t.Fatal("expected credential to become available") - } -} - -func TestExternalCredentialReverseSessionChangesEmitStatus(t *testing.T) { - subscriber := observable.NewSubscriber[struct{}](8) - subscription, _ := subscriber.Subscription() - credential := &externalCredential{ - tag: "receiver", - baseURL: reverseProxyBaseURL, - pollInterval: time.Minute, - logger: newTestLogger(), - statusSubscriber: subscriber, - } - - clientSession, _ := newTestYamuxSessionPair(t) - if !credential.setReverseSession(clientSession) { - t.Fatal("expected reverse session to be accepted") - } - if count := drainStatusEvents(subscription); count != 1 { - t.Fatalf("expected 1 status event after reverse session up, got %d", count) - } - if !credential.isAvailable() { - t.Fatal("expected receiver credential to become available") - } - - credential.clearReverseSession(clientSession) - if count := drainStatusEvents(subscription); count != 1 { - t.Fatalf("expected 1 status event after reverse session down, got %d", count) - } - if credential.isAvailable() { - t.Fatal("expected receiver credential to become unavailable") - } -}