From 6878ad0d358327633ccf457c2ba0dced0724908b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Sat, 14 Mar 2026 21:06:25 +0800 Subject: [PATCH] ccm,ocm: fix naming and error-handling convention violations - Rename credential interface to Credential (exported), cred to credential - Rename mutex/saveMutex to access/saveAccess per go-syntax.md - Fix abbreviations: reverseHttpClient, allCreds, credOpt, extCred, credDialer, reverseCredDialer, portStr - Replace errors.Is(http.ErrServerClosed) with E.IsClosed - Add E.IsClosedOrCanceled guard before streaming write error logs --- service/ccm/credential.go | 8 +- service/ccm/credential_builder.go | 100 +++++++++++----------- service/ccm/credential_external.go | 52 ++++++------ service/ccm/credential_provider.go | 124 +++++++++++++-------------- service/ccm/reverse.go | 13 ++- service/ccm/service.go | 53 ++++++------ service/ccm/service_handler.go | 7 ++ service/ccm/service_status.go | 14 +-- service/ccm/service_usage.go | 32 +++---- service/ccm/service_user.go | 4 +- service/ocm/credential.go | 8 +- service/ocm/credential_builder.go | 110 ++++++++++++------------ service/ocm/credential_external.go | 94 ++++++++++---------- service/ocm/credential_provider.go | 132 ++++++++++++++--------------- service/ocm/reverse.go | 13 ++- service/ocm/service.go | 57 ++++++------- service/ocm/service_handler.go | 6 ++ service/ocm/service_status.go | 14 +-- service/ocm/service_usage.go | 32 +++---- service/ocm/service_user.go | 4 +- service/ocm/service_websocket.go | 10 +-- 21 files changed, 448 insertions(+), 439 deletions(-) diff --git a/service/ccm/credential.go b/service/ccm/credential.go index 8589676a8..d5cae9e1e 100644 --- a/service/ccm/credential.go +++ b/service/ccm/credential.go @@ -90,7 +90,7 @@ func (c *credentialRequestContext) cancelRequest() { c.cancelOnce.Do(c.cancelFunc) } -type credential interface { +type Credential interface { tagName() string isAvailable() bool isUsable() bool @@ -130,11 +130,11 @@ const ( type credentialSelection struct { scope credentialSelectionScope - filter func(credential) bool + filter func(Credential) bool } -func (s credentialSelection) allows(cred credential) bool { - return s.filter == nil || s.filter(cred) +func (s credentialSelection) allows(credential Credential) bool { + return s.filter == nil || s.filter(credential) } func (s credentialSelection) scopeOrDefault() credentialSelectionScope { diff --git a/service/ccm/credential_builder.go b/service/ccm/credential_builder.go index c49a20195..63bfd0395 100644 --- a/service/ccm/credential_builder.go +++ b/service/ccm/credential_builder.go @@ -14,55 +14,55 @@ func buildCredentialProviders( ctx context.Context, options option.CCMServiceOptions, logger log.ContextLogger, -) (map[string]credentialProvider, []credential, error) { - allCredentialMap := make(map[string]credential) - var allCreds []credential +) (map[string]credentialProvider, []Credential, error) { + allCredentialMap := make(map[string]Credential) + var allCredentials []Credential providers := make(map[string]credentialProvider) // Pass 1: create default and external credentials - for _, credOpt := range options.Credentials { - switch credOpt.Type { + for _, credentialOption := range options.Credentials { + switch credentialOption.Type { case "default": - cred, err := newDefaultCredential(ctx, credOpt.Tag, credOpt.DefaultOptions, logger) + credential, err := newDefaultCredential(ctx, credentialOption.Tag, credentialOption.DefaultOptions, logger) if err != nil { return nil, nil, err } - allCredentialMap[credOpt.Tag] = cred - allCreds = append(allCreds, cred) - providers[credOpt.Tag] = &singleCredentialProvider{cred: cred} + allCredentialMap[credentialOption.Tag] = credential + allCredentials = append(allCredentials, credential) + providers[credentialOption.Tag] = &singleCredentialProvider{credential: credential} case "external": - cred, err := newExternalCredential(ctx, credOpt.Tag, credOpt.ExternalOptions, logger) + credential, err := newExternalCredential(ctx, credentialOption.Tag, credentialOption.ExternalOptions, logger) if err != nil { return nil, nil, err } - allCredentialMap[credOpt.Tag] = cred - allCreds = append(allCreds, cred) - providers[credOpt.Tag] = &singleCredentialProvider{cred: cred} + allCredentialMap[credentialOption.Tag] = credential + allCredentials = append(allCredentials, credential) + providers[credentialOption.Tag] = &singleCredentialProvider{credential: credential} } } // Pass 2: create balancer providers - for _, credOpt := range options.Credentials { - if credOpt.Type == "balancer" { - subCredentials, err := resolveCredentialTags(credOpt.BalancerOptions.Credentials, allCredentialMap, credOpt.Tag) + for _, credentialOption := range options.Credentials { + if credentialOption.Type == "balancer" { + subCredentials, err := resolveCredentialTags(credentialOption.BalancerOptions.Credentials, allCredentialMap, credentialOption.Tag) if err != nil { return nil, nil, err } - providers[credOpt.Tag] = newBalancerProvider(subCredentials, credOpt.BalancerOptions.Strategy, time.Duration(credOpt.BalancerOptions.PollInterval), credOpt.BalancerOptions.RebalanceThreshold, logger) + providers[credentialOption.Tag] = newBalancerProvider(subCredentials, credentialOption.BalancerOptions.Strategy, time.Duration(credentialOption.BalancerOptions.PollInterval), credentialOption.BalancerOptions.RebalanceThreshold, logger) } } - return providers, allCreds, nil + return providers, allCredentials, nil } -func resolveCredentialTags(tags []string, allCredentials map[string]credential, parentTag string) ([]credential, error) { - credentials := make([]credential, 0, len(tags)) +func resolveCredentialTags(tags []string, allCredentials map[string]Credential, parentTag string) ([]Credential, error) { + credentials := make([]Credential, 0, len(tags)) for _, tag := range tags { - cred, exists := allCredentials[tag] + credential, exists := allCredentials[tag] if !exists { return nil, E.New("credential ", parentTag, " references unknown credential: ", tag) } - credentials = append(credentials, cred) + credentials = append(credentials, credential) } if len(credentials) == 0 { return nil, E.New("credential ", parentTag, " has no sub-credentials") @@ -89,48 +89,48 @@ func validateCCMOptions(options option.CCMServiceOptions) error { if hasCredentials { tags := make(map[string]bool) credentialTypes := make(map[string]string) - for _, cred := range options.Credentials { - if tags[cred.Tag] { - return E.New("duplicate credential tag: ", cred.Tag) + for _, credential := range options.Credentials { + if tags[credential.Tag] { + return E.New("duplicate credential tag: ", credential.Tag) } - tags[cred.Tag] = true - credentialTypes[cred.Tag] = cred.Type - if cred.Type == "default" || cred.Type == "" { - if cred.DefaultOptions.Reserve5h > 99 { - return E.New("credential ", cred.Tag, ": reserve_5h must be at most 99") + tags[credential.Tag] = true + credentialTypes[credential.Tag] = credential.Type + if credential.Type == "default" || credential.Type == "" { + if credential.DefaultOptions.Reserve5h > 99 { + return E.New("credential ", credential.Tag, ": reserve_5h must be at most 99") } - if cred.DefaultOptions.ReserveWeekly > 99 { - return E.New("credential ", cred.Tag, ": reserve_weekly must be at most 99") + if credential.DefaultOptions.ReserveWeekly > 99 { + return E.New("credential ", credential.Tag, ": reserve_weekly must be at most 99") } - if cred.DefaultOptions.Limit5h > 100 { - return E.New("credential ", cred.Tag, ": limit_5h must be at most 100") + if credential.DefaultOptions.Limit5h > 100 { + return E.New("credential ", credential.Tag, ": limit_5h must be at most 100") } - if cred.DefaultOptions.LimitWeekly > 100 { - return E.New("credential ", cred.Tag, ": limit_weekly must be at most 100") + if credential.DefaultOptions.LimitWeekly > 100 { + return E.New("credential ", credential.Tag, ": limit_weekly must be at most 100") } - if cred.DefaultOptions.Reserve5h > 0 && cred.DefaultOptions.Limit5h > 0 { - return E.New("credential ", cred.Tag, ": reserve_5h and limit_5h are mutually exclusive") + if credential.DefaultOptions.Reserve5h > 0 && credential.DefaultOptions.Limit5h > 0 { + return E.New("credential ", credential.Tag, ": reserve_5h and limit_5h are mutually exclusive") } - if cred.DefaultOptions.ReserveWeekly > 0 && cred.DefaultOptions.LimitWeekly > 0 { - return E.New("credential ", cred.Tag, ": reserve_weekly and limit_weekly are mutually exclusive") + if credential.DefaultOptions.ReserveWeekly > 0 && credential.DefaultOptions.LimitWeekly > 0 { + return E.New("credential ", credential.Tag, ": reserve_weekly and limit_weekly are mutually exclusive") } } - if cred.Type == "external" { - if cred.ExternalOptions.Token == "" { - return E.New("credential ", cred.Tag, ": external credential requires token") + if credential.Type == "external" { + if credential.ExternalOptions.Token == "" { + return E.New("credential ", credential.Tag, ": external credential requires token") } - if cred.ExternalOptions.Reverse && cred.ExternalOptions.URL == "" { - return E.New("credential ", cred.Tag, ": reverse external credential requires url") + if credential.ExternalOptions.Reverse && credential.ExternalOptions.URL == "" { + return E.New("credential ", credential.Tag, ": reverse external credential requires url") } } - if cred.Type == "balancer" { - switch cred.BalancerOptions.Strategy { + if credential.Type == "balancer" { + switch credential.BalancerOptions.Strategy { case "", C.BalancerStrategyLeastUsed, C.BalancerStrategyRoundRobin, C.BalancerStrategyRandom, C.BalancerStrategyFallback: default: - return E.New("credential ", cred.Tag, ": unknown balancer strategy: ", cred.BalancerOptions.Strategy) + return E.New("credential ", credential.Tag, ": unknown balancer strategy: ", credential.BalancerOptions.Strategy) } - if cred.BalancerOptions.RebalanceThreshold < 0 { - return E.New("credential ", cred.Tag, ": rebalance_threshold must not be negative") + if credential.BalancerOptions.RebalanceThreshold < 0 { + return E.New("credential ", credential.Tag, ": rebalance_threshold must not be negative") } } } diff --git a/service/ccm/credential_external.go b/service/ccm/credential_external.go index 24ddf6c4a..eb75c5b08 100644 --- a/service/ccm/credential_external.go +++ b/service/ccm/credential_external.go @@ -48,7 +48,7 @@ type externalCredential struct { // Reverse proxy fields reverse bool - reverseHttpClient *http.Client + reverseHTTPClient *http.Client reverseSession *yamux.Session reverseAccess sync.RWMutex closed bool @@ -63,9 +63,9 @@ type externalCredential struct { } func externalCredentialURLPort(parsedURL *url.URL) uint16 { - portStr := parsedURL.Port() - if portStr != "" { - port, err := strconv.ParseUint(portStr, 10, 16) + portString := parsedURL.Port() + if portString != "" { + port, err := strconv.ParseUint(portString, 10, 16) if err == nil { return uint16(port) } @@ -113,7 +113,7 @@ func newExternalCredential(ctx context.Context, tag string, options option.CCMEx requestContext, cancelRequests := context.WithCancel(context.Background()) reverseContext, reverseCancel := context.WithCancel(context.Background()) - cred := &externalCredential{ + credential := &externalCredential{ tag: tag, token: options.Token, pollInterval: pollInterval, @@ -127,12 +127,12 @@ func newExternalCredential(ctx context.Context, tag string, options option.CCMEx if options.URL == "" { // Receiver mode: no URL, wait for reverse connection - cred.baseURL = reverseProxyBaseURL - cred.forwardHTTPClient = &http.Client{ + credential.baseURL = reverseProxyBaseURL + credential.forwardHTTPClient = &http.Client{ Transport: &http.Transport{ ForceAttemptHTTP2: false, DialContext: func(ctx context.Context, network, address string) (net.Conn, error) { - return cred.openReverseConnection(ctx) + return credential.openReverseConnection(ctx) }, }, } @@ -173,34 +173,34 @@ func newExternalCredential(ctx context.Context, tag string, options option.CCMEx } } - cred.baseURL = externalCredentialBaseURL(parsedURL) + credential.baseURL = externalCredentialBaseURL(parsedURL) if options.Reverse { // Connector mode: we dial out to serve, not to proxy - cred.connectorDialer = credentialDialer + credential.connectorDialer = credentialDialer if options.Server != "" { - cred.connectorDestination = M.ParseSocksaddrHostPort(options.Server, externalCredentialServerPort(parsedURL, options.ServerPort)) + credential.connectorDestination = M.ParseSocksaddrHostPort(options.Server, externalCredentialServerPort(parsedURL, options.ServerPort)) } else { - cred.connectorDestination = M.ParseSocksaddrHostPort(parsedURL.Hostname(), externalCredentialURLPort(parsedURL)) + credential.connectorDestination = M.ParseSocksaddrHostPort(parsedURL.Hostname(), externalCredentialURLPort(parsedURL)) } - cred.connectorRequestPath = externalCredentialReversePath(parsedURL, "/ccm/v1/reverse") - cred.connectorURL = parsedURL + credential.connectorRequestPath = externalCredentialReversePath(parsedURL, "/ccm/v1/reverse") + credential.connectorURL = parsedURL if parsedURL.Scheme == "https" { - cred.connectorTLS = &stdTLS.Config{ + credential.connectorTLS = &stdTLS.Config{ ServerName: parsedURL.Hostname(), RootCAs: adapter.RootPoolFromContext(ctx), Time: ntp.TimeFuncFromContext(ctx), } } - cred.forwardHTTPClient = &http.Client{Transport: transport} + credential.forwardHTTPClient = &http.Client{Transport: transport} } else { // Normal mode: standard HTTP client for proxying - cred.forwardHTTPClient = &http.Client{Transport: transport} - cred.reverseHttpClient = &http.Client{ + credential.forwardHTTPClient = &http.Client{Transport: transport} + credential.reverseHTTPClient = &http.Client{ Transport: &http.Transport{ ForceAttemptHTTP2: false, DialContext: func(ctx context.Context, network, address string) (net.Conn, error) { - return cred.openReverseConnection(ctx) + return credential.openReverseConnection(ctx) }, }, } @@ -208,7 +208,7 @@ func newExternalCredential(ctx context.Context, tag string, options option.CCMEx } if options.UsagesPath != "" { - cred.usageTracker = &AggregatedUsage{ + credential.usageTracker = &AggregatedUsage{ LastUpdated: time.Now(), Combinations: make([]CostCombination, 0), filePath: options.UsagesPath, @@ -216,7 +216,7 @@ func newExternalCredential(ctx context.Context, tag string, options option.CCMEx } } - return cred, nil + return credential, nil } func (c *externalCredential) start() error { @@ -352,7 +352,7 @@ func (c *externalCredential) getAccessToken() (string, error) { func (c *externalCredential) buildProxyRequest(ctx context.Context, original *http.Request, bodyBytes []byte, _ http.Header) (*http.Request, error) { baseURL := c.baseURL - if c.reverseHttpClient != nil { + if c.reverseHTTPClient != nil { session := c.getReverseSession() if session != nil && !session.IsClosed() { baseURL = reverseProxyBaseURL @@ -511,7 +511,7 @@ func (c *externalCredential) doPollUsageRequest(ctx context.Context) (*http.Resp } } // Try reverse transport first (single attempt, no retry) - if c.reverseHttpClient != nil { + if c.reverseHTTPClient != nil { session := c.getReverseSession() if session != nil && !session.IsClosed() { request, err := buildRequest(reverseProxyBaseURL)() @@ -519,7 +519,7 @@ func (c *externalCredential) doPollUsageRequest(ctx context.Context) (*http.Resp return nil, err } reverseClient := &http.Client{ - Transport: c.reverseHttpClient.Transport, + Transport: c.reverseHTTPClient.Transport, Timeout: 5 * time.Second, } response, err := reverseClient.Do(request) @@ -660,10 +660,10 @@ func (c *externalCredential) usageTrackerOrNil() *AggregatedUsage { } func (c *externalCredential) httpClient() *http.Client { - if c.reverseHttpClient != nil { + if c.reverseHTTPClient != nil { session := c.getReverseSession() if session != nil && !session.IsClosed() { - return c.reverseHttpClient + return c.reverseHTTPClient } } return c.forwardHTTPClient diff --git a/service/ccm/credential_provider.go b/service/ccm/credential_provider.go index cd77bfcdc..5500df6a1 100644 --- a/service/ccm/credential_provider.go +++ b/service/ccm/credential_provider.go @@ -13,29 +13,29 @@ import ( ) type credentialProvider interface { - selectCredential(sessionID string, selection credentialSelection) (credential, bool, error) - onRateLimited(sessionID string, cred credential, resetAt time.Time, selection credentialSelection) credential - linkProviderInterrupt(cred credential, selection credentialSelection, onInterrupt func()) func() bool + selectCredential(sessionID string, selection credentialSelection) (Credential, bool, error) + onRateLimited(sessionID string, credential Credential, resetAt time.Time, selection credentialSelection) Credential + linkProviderInterrupt(credential Credential, selection credentialSelection, onInterrupt func()) func() bool pollIfStale(ctx context.Context) - allCredentials() []credential + allCredentials() []Credential close() } type singleCredentialProvider struct { - cred credential + credential Credential sessionAccess sync.RWMutex sessions map[string]time.Time } -func (p *singleCredentialProvider) selectCredential(sessionID string, selection credentialSelection) (credential, bool, error) { - if !selection.allows(p.cred) { - return nil, false, E.New("credential ", p.cred.tagName(), " is filtered out") +func (p *singleCredentialProvider) selectCredential(sessionID string, selection credentialSelection) (Credential, bool, error) { + if !selection.allows(p.credential) { + return nil, false, E.New("credential ", p.credential.tagName(), " is filtered out") } - if !p.cred.isAvailable() { - return nil, false, p.cred.unavailableError() + if !p.credential.isAvailable() { + return nil, false, p.credential.unavailableError() } - if !p.cred.isUsable() { - return nil, false, E.New("credential ", p.cred.tagName(), " is rate-limited") + if !p.credential.isUsable() { + return nil, false, E.New("credential ", p.credential.tagName(), " is rate-limited") } var isNew bool if sessionID != "" { @@ -50,11 +50,11 @@ func (p *singleCredentialProvider) selectCredential(sessionID string, selection } p.sessionAccess.Unlock() } - return p.cred, isNew, nil + return p.credential, isNew, nil } -func (p *singleCredentialProvider) onRateLimited(_ string, cred credential, resetAt time.Time, _ credentialSelection) credential { - cred.markRateLimited(resetAt) +func (p *singleCredentialProvider) onRateLimited(_ string, credential Credential, resetAt time.Time, _ credentialSelection) Credential { + credential.markRateLimited(resetAt) return nil } @@ -68,16 +68,16 @@ func (p *singleCredentialProvider) pollIfStale(ctx context.Context) { } p.sessionAccess.Unlock() - if time.Since(p.cred.lastUpdatedTime()) > p.cred.pollBackoff(defaultPollInterval) { - p.cred.pollUsage(ctx) + if time.Since(p.credential.lastUpdatedTime()) > p.credential.pollBackoff(defaultPollInterval) { + p.credential.pollUsage(ctx) } } -func (p *singleCredentialProvider) allCredentials() []credential { - return []credential{p.cred} +func (p *singleCredentialProvider) allCredentials() []Credential { + return []Credential{p.credential} } -func (p *singleCredentialProvider) linkProviderInterrupt(_ credential, _ credentialSelection, _ func()) func() bool { +func (p *singleCredentialProvider) linkProviderInterrupt(_ Credential, _ credentialSelection, _ func()) func() bool { return func() bool { return false } @@ -102,7 +102,7 @@ type credentialInterruptEntry struct { } type balancerProvider struct { - credentials []credential + credentials []Credential strategy string roundRobinIndex atomic.Uint64 pollInterval time.Duration @@ -114,7 +114,7 @@ type balancerProvider struct { logger log.ContextLogger } -func newBalancerProvider(credentials []credential, strategy string, pollInterval time.Duration, rebalanceThreshold float64, logger log.ContextLogger) *balancerProvider { +func newBalancerProvider(credentials []Credential, strategy string, pollInterval time.Duration, rebalanceThreshold float64, logger log.ContextLogger) *balancerProvider { if pollInterval <= 0 { pollInterval = defaultPollInterval } @@ -129,7 +129,7 @@ func newBalancerProvider(credentials []credential, strategy string, pollInterval } } -func (p *balancerProvider) selectCredential(sessionID string, selection credentialSelection) (credential, bool, error) { +func (p *balancerProvider) selectCredential(sessionID string, selection credentialSelection) (Credential, bool, error) { if p.strategy == C.BalancerStrategyFallback { best := p.pickCredential(selection.filter) if best == nil { @@ -145,23 +145,23 @@ func (p *balancerProvider) selectCredential(sessionID string, selection credenti p.sessionAccess.RUnlock() if exists { if entry.selectionScope == selectionScope { - for _, cred := range p.credentials { - if cred.tagName() == entry.tag && selection.allows(cred) && cred.isUsable() { + for _, credential := range p.credentials { + if credential.tagName() == entry.tag && selection.allows(credential) && credential.isUsable() { if p.rebalanceThreshold > 0 && (p.strategy == "" || p.strategy == C.BalancerStrategyLeastUsed) { better := p.pickLeastUsed(selection.filter) - if better != nil && better.tagName() != cred.tagName() { - effectiveThreshold := p.rebalanceThreshold / cred.planWeight() - delta := cred.weeklyUtilization() - better.weeklyUtilization() + if better != nil && better.tagName() != credential.tagName() { + effectiveThreshold := p.rebalanceThreshold / credential.planWeight() + delta := credential.weeklyUtilization() - better.weeklyUtilization() if delta > effectiveThreshold { - p.logger.Info("rebalancing away from ", cred.tagName(), + p.logger.Info("rebalancing away from ", credential.tagName(), ": utilization delta ", delta, "% exceeds effective threshold ", - effectiveThreshold, "% (weight ", cred.planWeight(), ")") - p.rebalanceCredential(cred.tagName(), selectionScope) + effectiveThreshold, "% (weight ", credential.planWeight(), ")") + p.rebalanceCredential(credential.tagName(), selectionScope) break } } } - return cred, false, nil + return credential, false, nil } } } @@ -208,12 +208,12 @@ func (p *balancerProvider) rebalanceCredential(tag string, selectionScope creden p.sessionAccess.Unlock() } -func (p *balancerProvider) linkProviderInterrupt(cred credential, selection credentialSelection, onInterrupt func()) func() bool { +func (p *balancerProvider) linkProviderInterrupt(credential Credential, selection credentialSelection, onInterrupt func()) func() bool { if p.strategy == C.BalancerStrategyFallback { return func() bool { return false } } key := credentialInterruptKey{ - tag: cred.tagName(), + tag: credential.tagName(), selectionScope: selection.scopeOrDefault(), } p.interruptAccess.Lock() @@ -227,8 +227,8 @@ func (p *balancerProvider) linkProviderInterrupt(cred credential, selection cred return context.AfterFunc(entry.context, onInterrupt) } -func (p *balancerProvider) onRateLimited(sessionID string, cred credential, resetAt time.Time, selection credentialSelection) credential { - cred.markRateLimited(resetAt) +func (p *balancerProvider) onRateLimited(sessionID string, credential Credential, resetAt time.Time, selection credentialSelection) Credential { + credential.markRateLimited(resetAt) if p.strategy == C.BalancerStrategyFallback { return p.pickCredential(selection.filter) } @@ -251,7 +251,7 @@ func (p *balancerProvider) onRateLimited(sessionID string, cred credential, rese return best } -func (p *balancerProvider) pickCredential(filter func(credential) bool) credential { +func (p *balancerProvider) pickCredential(filter func(Credential) bool) Credential { switch p.strategy { case C.BalancerStrategyRoundRobin: return p.pickRoundRobin(filter) @@ -264,13 +264,13 @@ func (p *balancerProvider) pickCredential(filter func(credential) bool) credenti } } -func (p *balancerProvider) pickFallback(filter func(credential) bool) credential { - for _, cred := range p.credentials { - if filter != nil && !filter(cred) { +func (p *balancerProvider) pickFallback(filter func(Credential) bool) Credential { + for _, credential := range p.credentials { + if filter != nil && !filter(credential) { continue } - if cred.isUsable() { - return cred + if credential.isUsable() { + return credential } } return nil @@ -278,20 +278,20 @@ func (p *balancerProvider) pickFallback(filter func(credential) bool) credential const weeklyWindowHours = 7 * 24 -func (p *balancerProvider) pickLeastUsed(filter func(credential) bool) credential { - var best credential +func (p *balancerProvider) pickLeastUsed(filter func(Credential) bool) Credential { + var best Credential bestScore := float64(-1) now := time.Now() - for _, cred := range p.credentials { - if filter != nil && !filter(cred) { + for _, credential := range p.credentials { + if filter != nil && !filter(credential) { continue } - if !cred.isUsable() { + if !credential.isUsable() { continue } - remaining := cred.weeklyCap() - cred.weeklyUtilization() - score := remaining * cred.planWeight() - resetTime := cred.weeklyResetTime() + remaining := credential.weeklyCap() - credential.weeklyUtilization() + score := remaining * credential.planWeight() + resetTime := credential.weeklyResetTime() if !resetTime.IsZero() { timeUntilReset := resetTime.Sub(now) if timeUntilReset < time.Hour { @@ -301,13 +301,13 @@ func (p *balancerProvider) pickLeastUsed(filter func(credential) bool) credentia } if score > bestScore { bestScore = score - best = cred + best = credential } } return best } -func (p *balancerProvider) pickRoundRobin(filter func(credential) bool) credential { +func (p *balancerProvider) pickRoundRobin(filter func(Credential) bool) Credential { start := int(p.roundRobinIndex.Add(1) - 1) count := len(p.credentials) for offset := range count { @@ -322,8 +322,8 @@ func (p *balancerProvider) pickRoundRobin(filter func(credential) bool) credenti return nil } -func (p *balancerProvider) pickRandom(filter func(credential) bool) credential { - var usable []credential +func (p *balancerProvider) pickRandom(filter func(Credential) bool) Credential { + var usable []Credential for _, candidate := range p.credentials { if filter != nil && !filter(candidate) { continue @@ -348,14 +348,14 @@ func (p *balancerProvider) pollIfStale(ctx context.Context) { } p.sessionAccess.Unlock() - for _, cred := range p.credentials { - if time.Since(cred.lastUpdatedTime()) > cred.pollBackoff(p.pollInterval) { - cred.pollUsage(ctx) + for _, credential := range p.credentials { + if time.Since(credential.lastUpdatedTime()) > credential.pollBackoff(p.pollInterval) { + credential.pollUsage(ctx) } } } -func (p *balancerProvider) allCredentials() []credential { +func (p *balancerProvider) allCredentials() []Credential { return p.credentials } @@ -382,15 +382,15 @@ func ccmPlanWeight(accountType string, rateLimitTier string) float64 { } } -func allCredentialsUnavailableError(credentials []credential) error { +func allCredentialsUnavailableError(credentials []Credential) error { var hasUnavailable bool var earliest time.Time - for _, cred := range credentials { - if cred.unavailableError() != nil { + for _, credential := range credentials { + if credential.unavailableError() != nil { hasUnavailable = true continue } - resetAt := cred.earliestReset() + resetAt := credential.earliestReset() if !resetAt.IsZero() && (earliest.IsZero() || resetAt.Before(earliest)) { earliest = resetAt } diff --git a/service/ccm/reverse.go b/service/ccm/reverse.go index 6ecc224f9..97ef1751c 100644 --- a/service/ccm/reverse.go +++ b/service/ccm/reverse.go @@ -4,7 +4,6 @@ import ( "bufio" "context" stdTLS "crypto/tls" - "errors" "io" "math/rand/v2" "net" @@ -124,13 +123,13 @@ func (s *Service) handleReverseConnect(ctx context.Context, w http.ResponseWrite } func (s *Service) findReceiverCredential(token string) *externalCredential { - for _, cred := range s.allCredentials { - extCred, ok := cred.(*externalCredential) - if !ok || extCred.connectorURL != nil { + for _, credential := range s.allCredentials { + external, ok := credential.(*externalCredential) + if !ok || external.connectorURL != nil { continue } - if extCred.token == token { - return extCred + if external.token == token { + return external } } return nil @@ -248,7 +247,7 @@ func (c *externalCredential) connectorConnect(ctx context.Context) (time.Duratio } err = httpServer.Serve(&yamuxNetListener{session: session}) sessionLifetime := time.Since(serveStart) - if err != nil && !errors.Is(err, http.ErrServerClosed) && ctx.Err() == nil { + if err != nil && !E.IsClosed(err) && ctx.Err() == nil { return sessionLifetime, E.Cause(err, "serve") } return sessionLifetime, E.New("connection closed") diff --git a/service/ccm/service.go b/service/ccm/service.go index 6dce1931b..69964c02c 100644 --- a/service/ccm/service.go +++ b/service/ccm/service.go @@ -3,7 +3,6 @@ package ccm import ( "context" "encoding/json" - "errors" "net/http" "strings" "sync" @@ -55,18 +54,18 @@ func writeJSONError(w http.ResponseWriter, r *http.Request, statusCode int, erro }) } -func hasAlternativeCredential(provider credentialProvider, currentCredential credential, selection credentialSelection) bool { +func hasAlternativeCredential(provider credentialProvider, currentCredential Credential, selection credentialSelection) bool { if provider == nil || currentCredential == nil { return false } - for _, cred := range provider.allCredentials() { - if cred == currentCredential { + for _, credential := range provider.allCredentials() { + if credential == currentCredential { continue } - if !selection.allows(cred) { + if !selection.allows(credential) { continue } - if cred.isUsable() { + if credential.isUsable() { return true } } @@ -96,7 +95,7 @@ func writeCredentialUnavailableError( w http.ResponseWriter, r *http.Request, provider credentialProvider, - currentCredential credential, + currentCredential Credential, selection credentialSelection, fallback string, ) { @@ -111,8 +110,8 @@ func credentialSelectionForUser(userConfig *option.CCMUser) credentialSelection selection := credentialSelection{scope: credentialSelectionScopeAll} if userConfig != nil && !userConfig.AllowExternalUsage { selection.scope = credentialSelectionScopeNonExternal - selection.filter = func(cred credential) bool { - return !cred.isExternal() + selection.filter = func(credential Credential) bool { + return !credential.isExternal() } } return selection @@ -159,7 +158,7 @@ type Service struct { // Multi-credential mode providers map[string]credentialProvider - allCredentials []credential + allCredentials []Credential userConfigMap map[string]*option.CCMUser } @@ -204,7 +203,7 @@ func NewService(ctx context.Context, logger log.ContextLogger, tag string, optio } service.userConfigMap = userConfigMap } else { - cred, err := newDefaultCredential(ctx, "default", option.CCMDefaultCredentialOptions{ + credential, err := newDefaultCredential(ctx, "default", option.CCMDefaultCredentialOptions{ CredentialPath: options.CredentialPath, UsagesPath: options.UsagesPath, Detour: options.Detour, @@ -212,9 +211,9 @@ func NewService(ctx context.Context, logger log.ContextLogger, tag string, optio if err != nil { return nil, err } - service.legacyCredential = cred - service.legacyProvider = &singleCredentialProvider{cred: cred} - service.allCredentials = []credential{cred} + service.legacyCredential = credential + service.legacyProvider = &singleCredentialProvider{credential: credential} + service.allCredentials = []Credential{credential} } if options.TLS != nil { @@ -235,11 +234,11 @@ func (s *Service) Start(stage adapter.StartStage) error { s.userManager.UpdateUsers(s.options.Users) - for _, cred := range s.allCredentials { - if extCred, ok := cred.(*externalCredential); ok && extCred.reverse && extCred.connectorURL != nil { - extCred.reverseService = s + for _, credential := range s.allCredentials { + if external, ok := credential.(*externalCredential); ok && external.reverse && external.connectorURL != nil { + external.reverseService = s } - err := cred.start() + err := credential.start() if err != nil { return err } @@ -271,7 +270,7 @@ func (s *Service) Start(stage adapter.StartStage) error { go func() { serveErr := s.httpServer.Serve(tcpListener) - if serveErr != nil && !errors.Is(serveErr, http.ErrServerClosed) { + if serveErr != nil && !E.IsClosed(serveErr) { s.logger.Error("serve error: ", serveErr) } }() @@ -280,15 +279,15 @@ func (s *Service) Start(stage adapter.StartStage) error { } func (s *Service) InterfaceUpdated() { - for _, cred := range s.allCredentials { - extCred, ok := cred.(*externalCredential) + for _, credential := range s.allCredentials { + external, ok := credential.(*externalCredential) if !ok { continue } - if extCred.reverse && extCred.connectorURL != nil { - extCred.reverseService = s - extCred.resetReverseContext() - go extCred.connectorLoop() + if external.reverse && external.connectorURL != nil { + external.reverseService = s + external.resetReverseContext() + go external.connectorLoop() } } } @@ -300,8 +299,8 @@ func (s *Service) Close() error { s.tlsConfig, ) - for _, cred := range s.allCredentials { - cred.close() + for _, credential := range s.allCredentials { + credential.close() } return err diff --git a/service/ccm/service_handler.go b/service/ccm/service_handler.go index 7dd0c6411..7a59cfe4a 100644 --- a/service/ccm/service_handler.go +++ b/service/ccm/service_handler.go @@ -14,6 +14,7 @@ import ( "github.com/sagernet/sing-box/log" "github.com/sagernet/sing-box/option" "github.com/sagernet/sing/common/buf" + E "github.com/sagernet/sing/common/exceptions" "github.com/anthropics/anthropic-sdk-go" ) @@ -336,6 +337,9 @@ func (s *Service) ServeHTTP(w http.ResponseWriter, r *http.Request) { if n > 0 { _, writeError := w.Write(buffer[:n]) if writeError != nil { + if E.IsClosedOrCanceled(writeError) { + return + } s.logger.ErrorContext(ctx, "write streaming response: ", writeError) return } @@ -462,6 +466,9 @@ func (s *Service) handleResponseWithTracking(ctx context.Context, writer http.Re _, writeError := writer.Write(buffer[:n]) if writeError != nil { + if E.IsClosedOrCanceled(writeError) { + return + } s.logger.ErrorContext(ctx, "write streaming response: ", writeError) return } diff --git a/service/ccm/service_status.go b/service/ccm/service_status.go index 3f91b4614..75929c59f 100644 --- a/service/ccm/service_status.go +++ b/service/ccm/service_status.go @@ -62,22 +62,22 @@ func (s *Service) handleStatusEndpoint(w http.ResponseWriter, r *http.Request) { func (s *Service) computeAggregatedUtilization(provider credentialProvider, userConfig *option.CCMUser) (float64, float64, float64) { var totalWeightedRemaining5h, totalWeightedRemainingWeekly, totalWeight float64 - for _, cred := range provider.allCredentials() { - if !cred.isAvailable() { + for _, credential := range provider.allCredentials() { + if !credential.isAvailable() { continue } - if userConfig.ExternalCredential != "" && cred.tagName() == userConfig.ExternalCredential { + if userConfig.ExternalCredential != "" && credential.tagName() == userConfig.ExternalCredential { continue } - if !userConfig.AllowExternalUsage && cred.isExternal() { + if !userConfig.AllowExternalUsage && credential.isExternal() { continue } - weight := cred.planWeight() - remaining5h := cred.fiveHourCap() - cred.fiveHourUtilization() + weight := credential.planWeight() + remaining5h := credential.fiveHourCap() - credential.fiveHourUtilization() if remaining5h < 0 { remaining5h = 0 } - remainingWeekly := cred.weeklyCap() - cred.weeklyUtilization() + remainingWeekly := credential.weeklyCap() - credential.weeklyUtilization() if remainingWeekly < 0 { remainingWeekly = 0 } diff --git a/service/ccm/service_usage.go b/service/ccm/service_usage.go index 36e9ee65d..e23db6654 100644 --- a/service/ccm/service_usage.go +++ b/service/ccm/service_usage.go @@ -35,13 +35,13 @@ type CostCombination struct { type AggregatedUsage struct { LastUpdated time.Time `json:"last_updated"` Combinations []CostCombination `json:"combinations"` - mutex sync.Mutex + access sync.Mutex filePath string logger log.ContextLogger lastSaveTime time.Time pendingSave bool saveTimer *time.Timer - saveMutex sync.Mutex + saveAccess sync.Mutex } type UsageStatsJSON struct { @@ -527,8 +527,8 @@ func deriveWeekStartUnix(cycleHint *WeeklyCycleHint) int64 { } func (u *AggregatedUsage) ToJSON() *AggregatedUsageJSON { - u.mutex.Lock() - defer u.mutex.Unlock() + u.access.Lock() + defer u.access.Unlock() result := &AggregatedUsageJSON{ LastUpdated: u.LastUpdated, @@ -561,8 +561,8 @@ func (u *AggregatedUsage) ToJSON() *AggregatedUsageJSON { } func (u *AggregatedUsage) Load() error { - u.mutex.Lock() - defer u.mutex.Unlock() + u.access.Lock() + defer u.access.Unlock() u.LastUpdated = time.Time{} u.Combinations = nil @@ -608,9 +608,9 @@ func (u *AggregatedUsage) Save() error { defer os.Remove(tmpFile) err = os.Rename(tmpFile, u.filePath) if err == nil { - u.saveMutex.Lock() + u.saveAccess.Lock() u.lastSaveTime = time.Now() - u.saveMutex.Unlock() + u.saveAccess.Unlock() } return err } @@ -644,8 +644,8 @@ func (u *AggregatedUsage) AddUsageWithCycleHint( observedAt = time.Now() } - u.mutex.Lock() - defer u.mutex.Unlock() + u.access.Lock() + defer u.access.Unlock() u.LastUpdated = observedAt weekStartUnix := deriveWeekStartUnix(cycleHint) @@ -660,8 +660,8 @@ func (u *AggregatedUsage) AddUsageWithCycleHint( func (u *AggregatedUsage) scheduleSave() { const saveInterval = time.Minute - u.saveMutex.Lock() - defer u.saveMutex.Unlock() + u.saveAccess.Lock() + defer u.saveAccess.Unlock() timeSinceLastSave := time.Since(u.lastSaveTime) @@ -678,9 +678,9 @@ func (u *AggregatedUsage) scheduleSave() { remainingTime := saveInterval - timeSinceLastSave u.saveTimer = time.AfterFunc(remainingTime, func() { - u.saveMutex.Lock() + u.saveAccess.Lock() u.pendingSave = false - u.saveMutex.Unlock() + u.saveAccess.Unlock() u.saveAsync() }) } @@ -695,8 +695,8 @@ func (u *AggregatedUsage) saveAsync() { } func (u *AggregatedUsage) cancelPendingSave() { - u.saveMutex.Lock() - defer u.saveMutex.Unlock() + u.saveAccess.Lock() + defer u.saveAccess.Unlock() if u.saveTimer != nil { u.saveTimer.Stop() diff --git a/service/ccm/service_user.go b/service/ccm/service_user.go index 149894c04..e3f52bdf0 100644 --- a/service/ccm/service_user.go +++ b/service/ccm/service_user.go @@ -7,8 +7,8 @@ import ( ) type UserManager struct { - access sync.RWMutex - tokenMap map[string]string + access sync.RWMutex + tokenMap map[string]string } func (m *UserManager) UpdateUsers(users []option.CCMUser) { diff --git a/service/ocm/credential.go b/service/ocm/credential.go index 27a889470..e0ad9f565 100644 --- a/service/ocm/credential.go +++ b/service/ocm/credential.go @@ -92,7 +92,7 @@ func (c *credentialRequestContext) cancelRequest() { c.cancelOnce.Do(c.cancelFunc) } -type credential interface { +type Credential interface { tagName() string isAvailable() bool isUsable() bool @@ -139,11 +139,11 @@ const ( type credentialSelection struct { scope credentialSelectionScope - filter func(credential) bool + filter func(Credential) bool } -func (s credentialSelection) allows(cred credential) bool { - return s.filter == nil || s.filter(cred) +func (s credentialSelection) allows(credential Credential) bool { + return s.filter == nil || s.filter(credential) } func (s credentialSelection) scopeOrDefault() credentialSelectionScope { diff --git a/service/ocm/credential_builder.go b/service/ocm/credential_builder.go index 5faaf67c6..e308d04d1 100644 --- a/service/ocm/credential_builder.go +++ b/service/ocm/credential_builder.go @@ -14,55 +14,55 @@ func buildOCMCredentialProviders( ctx context.Context, options option.OCMServiceOptions, logger log.ContextLogger, -) (map[string]credentialProvider, []credential, error) { - allCredentialMap := make(map[string]credential) - var allCreds []credential +) (map[string]credentialProvider, []Credential, error) { + allCredentialMap := make(map[string]Credential) + var allCredentials []Credential providers := make(map[string]credentialProvider) // Pass 1: create default and external credentials - for _, credOpt := range options.Credentials { - switch credOpt.Type { + for _, credentialOption := range options.Credentials { + switch credentialOption.Type { case "default": - cred, err := newDefaultCredential(ctx, credOpt.Tag, credOpt.DefaultOptions, logger) + credential, err := newDefaultCredential(ctx, credentialOption.Tag, credentialOption.DefaultOptions, logger) if err != nil { return nil, nil, err } - allCredentialMap[credOpt.Tag] = cred - allCreds = append(allCreds, cred) - providers[credOpt.Tag] = &singleCredentialProvider{cred: cred} + allCredentialMap[credentialOption.Tag] = credential + allCredentials = append(allCredentials, credential) + providers[credentialOption.Tag] = &singleCredentialProvider{credential: credential} case "external": - cred, err := newExternalCredential(ctx, credOpt.Tag, credOpt.ExternalOptions, logger) + credential, err := newExternalCredential(ctx, credentialOption.Tag, credentialOption.ExternalOptions, logger) if err != nil { return nil, nil, err } - allCredentialMap[credOpt.Tag] = cred - allCreds = append(allCreds, cred) - providers[credOpt.Tag] = &singleCredentialProvider{cred: cred} + allCredentialMap[credentialOption.Tag] = credential + allCredentials = append(allCredentials, credential) + providers[credentialOption.Tag] = &singleCredentialProvider{credential: credential} } } // Pass 2: create balancer providers - for _, credOpt := range options.Credentials { - if credOpt.Type == "balancer" { - subCredentials, err := resolveCredentialTags(credOpt.BalancerOptions.Credentials, allCredentialMap, credOpt.Tag) + for _, credentialOption := range options.Credentials { + if credentialOption.Type == "balancer" { + subCredentials, err := resolveCredentialTags(credentialOption.BalancerOptions.Credentials, allCredentialMap, credentialOption.Tag) if err != nil { return nil, nil, err } - providers[credOpt.Tag] = newBalancerProvider(subCredentials, credOpt.BalancerOptions.Strategy, time.Duration(credOpt.BalancerOptions.PollInterval), credOpt.BalancerOptions.RebalanceThreshold, logger) + providers[credentialOption.Tag] = newBalancerProvider(subCredentials, credentialOption.BalancerOptions.Strategy, time.Duration(credentialOption.BalancerOptions.PollInterval), credentialOption.BalancerOptions.RebalanceThreshold, logger) } } - return providers, allCreds, nil + return providers, allCredentials, nil } -func resolveCredentialTags(tags []string, allCredentials map[string]credential, parentTag string) ([]credential, error) { - credentials := make([]credential, 0, len(tags)) +func resolveCredentialTags(tags []string, allCredentials map[string]Credential, parentTag string) ([]Credential, error) { + credentials := make([]Credential, 0, len(tags)) for _, tag := range tags { - cred, exists := allCredentials[tag] + credential, exists := allCredentials[tag] if !exists { return nil, E.New("credential ", parentTag, " references unknown credential: ", tag) } - credentials = append(credentials, cred) + credentials = append(credentials, credential) } if len(credentials) == 0 { return nil, E.New("credential ", parentTag, " has no sub-credentials") @@ -89,48 +89,48 @@ func validateOCMOptions(options option.OCMServiceOptions) error { if hasCredentials { tags := make(map[string]bool) credentialTypes := make(map[string]string) - for _, cred := range options.Credentials { - if tags[cred.Tag] { - return E.New("duplicate credential tag: ", cred.Tag) + for _, credential := range options.Credentials { + if tags[credential.Tag] { + return E.New("duplicate credential tag: ", credential.Tag) } - tags[cred.Tag] = true - credentialTypes[cred.Tag] = cred.Type - if cred.Type == "default" || cred.Type == "" { - if cred.DefaultOptions.Reserve5h > 99 { - return E.New("credential ", cred.Tag, ": reserve_5h must be at most 99") + tags[credential.Tag] = true + credentialTypes[credential.Tag] = credential.Type + if credential.Type == "default" || credential.Type == "" { + if credential.DefaultOptions.Reserve5h > 99 { + return E.New("credential ", credential.Tag, ": reserve_5h must be at most 99") } - if cred.DefaultOptions.ReserveWeekly > 99 { - return E.New("credential ", cred.Tag, ": reserve_weekly must be at most 99") + if credential.DefaultOptions.ReserveWeekly > 99 { + return E.New("credential ", credential.Tag, ": reserve_weekly must be at most 99") } - if cred.DefaultOptions.Limit5h > 100 { - return E.New("credential ", cred.Tag, ": limit_5h must be at most 100") + if credential.DefaultOptions.Limit5h > 100 { + return E.New("credential ", credential.Tag, ": limit_5h must be at most 100") } - if cred.DefaultOptions.LimitWeekly > 100 { - return E.New("credential ", cred.Tag, ": limit_weekly must be at most 100") + if credential.DefaultOptions.LimitWeekly > 100 { + return E.New("credential ", credential.Tag, ": limit_weekly must be at most 100") } - if cred.DefaultOptions.Reserve5h > 0 && cred.DefaultOptions.Limit5h > 0 { - return E.New("credential ", cred.Tag, ": reserve_5h and limit_5h are mutually exclusive") + if credential.DefaultOptions.Reserve5h > 0 && credential.DefaultOptions.Limit5h > 0 { + return E.New("credential ", credential.Tag, ": reserve_5h and limit_5h are mutually exclusive") } - if cred.DefaultOptions.ReserveWeekly > 0 && cred.DefaultOptions.LimitWeekly > 0 { - return E.New("credential ", cred.Tag, ": reserve_weekly and limit_weekly are mutually exclusive") + if credential.DefaultOptions.ReserveWeekly > 0 && credential.DefaultOptions.LimitWeekly > 0 { + return E.New("credential ", credential.Tag, ": reserve_weekly and limit_weekly are mutually exclusive") } } - if cred.Type == "external" { - if cred.ExternalOptions.Token == "" { - return E.New("credential ", cred.Tag, ": external credential requires token") + if credential.Type == "external" { + if credential.ExternalOptions.Token == "" { + return E.New("credential ", credential.Tag, ": external credential requires token") } - if cred.ExternalOptions.Reverse && cred.ExternalOptions.URL == "" { - return E.New("credential ", cred.Tag, ": reverse external credential requires url") + if credential.ExternalOptions.Reverse && credential.ExternalOptions.URL == "" { + return E.New("credential ", credential.Tag, ": reverse external credential requires url") } } - if cred.Type == "balancer" { - switch cred.BalancerOptions.Strategy { + if credential.Type == "balancer" { + switch credential.BalancerOptions.Strategy { case "", C.BalancerStrategyLeastUsed, C.BalancerStrategyRoundRobin, C.BalancerStrategyRandom, C.BalancerStrategyFallback: default: - return E.New("credential ", cred.Tag, ": unknown balancer strategy: ", cred.BalancerOptions.Strategy) + return E.New("credential ", credential.Tag, ": unknown balancer strategy: ", credential.BalancerOptions.Strategy) } - if cred.BalancerOptions.RebalanceThreshold < 0 { - return E.New("credential ", cred.Tag, ": rebalance_threshold must not be negative") + if credential.BalancerOptions.RebalanceThreshold < 0 { + return E.New("credential ", credential.Tag, ": rebalance_threshold must not be negative") } } } @@ -160,14 +160,14 @@ func validateOCMCompositeCredentialModes( options option.OCMServiceOptions, providers map[string]credentialProvider, ) error { - for _, credOpt := range options.Credentials { - if credOpt.Type != "balancer" { + for _, credentialOption := range options.Credentials { + if credentialOption.Type != "balancer" { continue } - provider, exists := providers[credOpt.Tag] + provider, exists := providers[credentialOption.Tag] if !exists { - return E.New("unknown credential: ", credOpt.Tag) + return E.New("unknown credential: ", credentialOption.Tag) } for _, subCred := range provider.allCredentials() { @@ -176,7 +176,7 @@ func validateOCMCompositeCredentialModes( } if subCred.ocmIsAPIKeyMode() { return E.New( - "credential ", credOpt.Tag, + "credential ", credentialOption.Tag, " references API key default credential ", subCred.tagName(), "; balancer and fallback only support OAuth default credentials", ) diff --git a/service/ocm/credential_external.go b/service/ocm/credential_external.go index 968bf904d..02675780b 100644 --- a/service/ocm/credential_external.go +++ b/service/ocm/credential_external.go @@ -33,7 +33,7 @@ type externalCredential struct { tag string baseURL string token string - credDialer N.Dialer + credentialDialer N.Dialer forwardHTTPClient *http.Client state credentialState stateAccess sync.RWMutex @@ -49,20 +49,20 @@ type externalCredential struct { requestAccess sync.Mutex // Reverse proxy fields - reverse bool - reverseHttpClient *http.Client - reverseCredDialer N.Dialer - reverseSession *yamux.Session - reverseAccess sync.RWMutex - closed bool - reverseContext context.Context - reverseCancel context.CancelFunc - connectorDialer N.Dialer - connectorDestination M.Socksaddr - connectorRequestPath string - connectorURL *url.URL - connectorTLS *stdTLS.Config - reverseService http.Handler + reverse bool + reverseHTTPClient *http.Client + reverseCredentialDialer N.Dialer + reverseSession *yamux.Session + reverseAccess sync.RWMutex + closed bool + reverseContext context.Context + reverseCancel context.CancelFunc + connectorDialer N.Dialer + connectorDestination M.Socksaddr + connectorRequestPath string + connectorURL *url.URL + connectorTLS *stdTLS.Config + reverseService http.Handler } type reverseSessionDialer struct { @@ -81,9 +81,9 @@ func (d reverseSessionDialer) ListenPacket(ctx context.Context, destination M.So } func externalCredentialURLPort(parsedURL *url.URL) uint16 { - portStr := parsedURL.Port() - if portStr != "" { - port, err := strconv.ParseUint(portStr, 10, 16) + portString := parsedURL.Port() + if portString != "" { + port, err := strconv.ParseUint(portString, 10, 16) if err == nil { return uint16(port) } @@ -131,7 +131,7 @@ func newExternalCredential(ctx context.Context, tag string, options option.OCMEx requestContext, cancelRequests := context.WithCancel(context.Background()) reverseContext, reverseCancel := context.WithCancel(context.Background()) - cred := &externalCredential{ + credential := &externalCredential{ tag: tag, token: options.Token, pollInterval: pollInterval, @@ -145,13 +145,13 @@ func newExternalCredential(ctx context.Context, tag string, options option.OCMEx if options.URL == "" { // Receiver mode: no URL, wait for reverse connection - cred.baseURL = reverseProxyBaseURL - cred.credDialer = reverseSessionDialer{credential: cred} - cred.forwardHTTPClient = &http.Client{ + credential.baseURL = reverseProxyBaseURL + credential.credentialDialer = reverseSessionDialer{credential: credential} + credential.forwardHTTPClient = &http.Client{ Transport: &http.Transport{ ForceAttemptHTTP2: false, DialContext: func(ctx context.Context, network, address string) (net.Conn, error) { - return cred.openReverseConnection(ctx) + return credential.openReverseConnection(ctx) }, }, } @@ -192,36 +192,36 @@ func newExternalCredential(ctx context.Context, tag string, options option.OCMEx } } - cred.baseURL = externalCredentialBaseURL(parsedURL) + credential.baseURL = externalCredentialBaseURL(parsedURL) if options.Reverse { // Connector mode: we dial out to serve, not to proxy - cred.connectorDialer = credentialDialer + credential.connectorDialer = credentialDialer if options.Server != "" { - cred.connectorDestination = M.ParseSocksaddrHostPort(options.Server, externalCredentialServerPort(parsedURL, options.ServerPort)) + credential.connectorDestination = M.ParseSocksaddrHostPort(options.Server, externalCredentialServerPort(parsedURL, options.ServerPort)) } else { - cred.connectorDestination = M.ParseSocksaddrHostPort(parsedURL.Hostname(), externalCredentialURLPort(parsedURL)) + credential.connectorDestination = M.ParseSocksaddrHostPort(parsedURL.Hostname(), externalCredentialURLPort(parsedURL)) } - cred.connectorRequestPath = externalCredentialReversePath(parsedURL, "/ocm/v1/reverse") - cred.connectorURL = parsedURL + credential.connectorRequestPath = externalCredentialReversePath(parsedURL, "/ocm/v1/reverse") + credential.connectorURL = parsedURL if parsedURL.Scheme == "https" { - cred.connectorTLS = &stdTLS.Config{ + credential.connectorTLS = &stdTLS.Config{ ServerName: parsedURL.Hostname(), RootCAs: adapter.RootPoolFromContext(ctx), Time: ntp.TimeFuncFromContext(ctx), } } - cred.forwardHTTPClient = &http.Client{Transport: transport} + credential.forwardHTTPClient = &http.Client{Transport: transport} } else { // Normal mode: standard HTTP client for proxying - cred.credDialer = credentialDialer - cred.forwardHTTPClient = &http.Client{Transport: transport} - cred.reverseCredDialer = reverseSessionDialer{credential: cred} - cred.reverseHttpClient = &http.Client{ + credential.credentialDialer = credentialDialer + credential.forwardHTTPClient = &http.Client{Transport: transport} + credential.reverseCredentialDialer = reverseSessionDialer{credential: credential} + credential.reverseHTTPClient = &http.Client{ Transport: &http.Transport{ ForceAttemptHTTP2: false, DialContext: func(ctx context.Context, network, address string) (net.Conn, error) { - return cred.openReverseConnection(ctx) + return credential.openReverseConnection(ctx) }, }, } @@ -229,7 +229,7 @@ func newExternalCredential(ctx context.Context, tag string, options option.OCMEx } if options.UsagesPath != "" { - cred.usageTracker = &AggregatedUsage{ + credential.usageTracker = &AggregatedUsage{ LastUpdated: time.Now(), Combinations: make([]CostCombination, 0), filePath: options.UsagesPath, @@ -237,7 +237,7 @@ func newExternalCredential(ctx context.Context, tag string, options option.OCMEx } } - return cred, nil + return credential, nil } func (c *externalCredential) start() error { @@ -376,7 +376,7 @@ func (c *externalCredential) getAccessToken() (string, error) { func (c *externalCredential) buildProxyRequest(ctx context.Context, original *http.Request, bodyBytes []byte, _ http.Header) (*http.Request, error) { baseURL := c.baseURL - if c.reverseHttpClient != nil { + if c.reverseHTTPClient != nil { session := c.getReverseSession() if session != nil && !session.IsClosed() { baseURL = reverseProxyBaseURL @@ -550,7 +550,7 @@ func (c *externalCredential) doPollUsageRequest(ctx context.Context) (*http.Resp } } // Try reverse transport first (single attempt, no retry) - if c.reverseHttpClient != nil { + if c.reverseHTTPClient != nil { session := c.getReverseSession() if session != nil && !session.IsClosed() { request, err := buildRequest(reverseProxyBaseURL)() @@ -558,7 +558,7 @@ func (c *externalCredential) doPollUsageRequest(ctx context.Context) (*http.Resp return nil, err } reverseClient := &http.Client{ - Transport: c.reverseHttpClient.Transport, + Transport: c.reverseHTTPClient.Transport, Timeout: 5 * time.Second, } response, err := reverseClient.Do(request) @@ -699,23 +699,23 @@ func (c *externalCredential) usageTrackerOrNil() *AggregatedUsage { } func (c *externalCredential) httpClient() *http.Client { - if c.reverseHttpClient != nil { + if c.reverseHTTPClient != nil { session := c.getReverseSession() if session != nil && !session.IsClosed() { - return c.reverseHttpClient + return c.reverseHTTPClient } } return c.forwardHTTPClient } func (c *externalCredential) ocmDialer() N.Dialer { - if c.reverseCredDialer != nil { + if c.reverseCredentialDialer != nil { session := c.getReverseSession() if session != nil && !session.IsClosed() { - return c.reverseCredDialer + return c.reverseCredentialDialer } } - return c.credDialer + return c.credentialDialer } func (c *externalCredential) ocmIsAPIKeyMode() bool { @@ -727,7 +727,7 @@ func (c *externalCredential) ocmGetAccountID() string { } func (c *externalCredential) ocmGetBaseURL() string { - if c.reverseHttpClient != nil { + if c.reverseHTTPClient != nil { session := c.getReverseSession() if session != nil && !session.IsClosed() { return reverseProxyBaseURL diff --git a/service/ocm/credential_provider.go b/service/ocm/credential_provider.go index 53383e368..6f3da6b43 100644 --- a/service/ocm/credential_provider.go +++ b/service/ocm/credential_provider.go @@ -13,29 +13,29 @@ import ( ) type credentialProvider interface { - selectCredential(sessionID string, selection credentialSelection) (credential, bool, error) - onRateLimited(sessionID string, cred credential, resetAt time.Time, selection credentialSelection) credential - linkProviderInterrupt(cred credential, selection credentialSelection, onInterrupt func()) func() bool + selectCredential(sessionID string, selection credentialSelection) (Credential, bool, error) + onRateLimited(sessionID string, credential Credential, resetAt time.Time, selection credentialSelection) Credential + linkProviderInterrupt(credential Credential, selection credentialSelection, onInterrupt func()) func() bool pollIfStale(ctx context.Context) - allCredentials() []credential + allCredentials() []Credential close() } type singleCredentialProvider struct { - cred credential + credential Credential sessionAccess sync.RWMutex sessions map[string]time.Time } -func (p *singleCredentialProvider) selectCredential(sessionID string, selection credentialSelection) (credential, bool, error) { - if !selection.allows(p.cred) { - return nil, false, E.New("credential ", p.cred.tagName(), " is filtered out") +func (p *singleCredentialProvider) selectCredential(sessionID string, selection credentialSelection) (Credential, bool, error) { + if !selection.allows(p.credential) { + return nil, false, E.New("credential ", p.credential.tagName(), " is filtered out") } - if !p.cred.isAvailable() { - return nil, false, p.cred.unavailableError() + if !p.credential.isAvailable() { + return nil, false, p.credential.unavailableError() } - if !p.cred.isUsable() { - return nil, false, E.New("credential ", p.cred.tagName(), " is rate-limited") + if !p.credential.isUsable() { + return nil, false, E.New("credential ", p.credential.tagName(), " is rate-limited") } var isNew bool if sessionID != "" { @@ -50,11 +50,11 @@ func (p *singleCredentialProvider) selectCredential(sessionID string, selection } p.sessionAccess.Unlock() } - return p.cred, isNew, nil + return p.credential, isNew, nil } -func (p *singleCredentialProvider) onRateLimited(_ string, cred credential, resetAt time.Time, _ credentialSelection) credential { - cred.markRateLimited(resetAt) +func (p *singleCredentialProvider) onRateLimited(_ string, credential Credential, resetAt time.Time, _ credentialSelection) Credential { + credential.markRateLimited(resetAt) return nil } @@ -68,16 +68,16 @@ func (p *singleCredentialProvider) pollIfStale(ctx context.Context) { } p.sessionAccess.Unlock() - if time.Since(p.cred.lastUpdatedTime()) > p.cred.pollBackoff(defaultPollInterval) { - p.cred.pollUsage(ctx) + if time.Since(p.credential.lastUpdatedTime()) > p.credential.pollBackoff(defaultPollInterval) { + p.credential.pollUsage(ctx) } } -func (p *singleCredentialProvider) allCredentials() []credential { - return []credential{p.cred} +func (p *singleCredentialProvider) allCredentials() []Credential { + return []Credential{p.credential} } -func (p *singleCredentialProvider) linkProviderInterrupt(_ credential, _ credentialSelection, _ func()) func() bool { +func (p *singleCredentialProvider) linkProviderInterrupt(_ Credential, _ credentialSelection, _ func()) func() bool { return func() bool { return false } @@ -102,7 +102,7 @@ type credentialInterruptEntry struct { } type balancerProvider struct { - credentials []credential + credentials []Credential strategy string roundRobinIndex atomic.Uint64 pollInterval time.Duration @@ -114,11 +114,11 @@ type balancerProvider struct { logger log.ContextLogger } -func compositeCredentialSelectable(cred credential) bool { - return !cred.ocmIsAPIKeyMode() +func compositeCredentialSelectable(credential Credential) bool { + return !credential.ocmIsAPIKeyMode() } -func newBalancerProvider(credentials []credential, strategy string, pollInterval time.Duration, rebalanceThreshold float64, logger log.ContextLogger) *balancerProvider { +func newBalancerProvider(credentials []Credential, strategy string, pollInterval time.Duration, rebalanceThreshold float64, logger log.ContextLogger) *balancerProvider { if pollInterval <= 0 { pollInterval = defaultPollInterval } @@ -133,7 +133,7 @@ func newBalancerProvider(credentials []credential, strategy string, pollInterval } } -func (p *balancerProvider) selectCredential(sessionID string, selection credentialSelection) (credential, bool, error) { +func (p *balancerProvider) selectCredential(sessionID string, selection credentialSelection) (Credential, bool, error) { if p.strategy == C.BalancerStrategyFallback { best := p.pickCredential(selection.filter) if best == nil { @@ -149,23 +149,23 @@ func (p *balancerProvider) selectCredential(sessionID string, selection credenti p.sessionAccess.RUnlock() if exists { if entry.selectionScope == selectionScope { - for _, cred := range p.credentials { - if cred.tagName() == entry.tag && compositeCredentialSelectable(cred) && selection.allows(cred) && cred.isUsable() { + for _, credential := range p.credentials { + if credential.tagName() == entry.tag && compositeCredentialSelectable(credential) && selection.allows(credential) && credential.isUsable() { if p.rebalanceThreshold > 0 && (p.strategy == "" || p.strategy == C.BalancerStrategyLeastUsed) { better := p.pickLeastUsed(selection.filter) - if better != nil && better.tagName() != cred.tagName() { - effectiveThreshold := p.rebalanceThreshold / cred.planWeight() - delta := cred.weeklyUtilization() - better.weeklyUtilization() + if better != nil && better.tagName() != credential.tagName() { + effectiveThreshold := p.rebalanceThreshold / credential.planWeight() + delta := credential.weeklyUtilization() - better.weeklyUtilization() if delta > effectiveThreshold { - p.logger.Info("rebalancing away from ", cred.tagName(), + p.logger.Info("rebalancing away from ", credential.tagName(), ": utilization delta ", delta, "% exceeds effective threshold ", - effectiveThreshold, "% (weight ", cred.planWeight(), ")") - p.rebalanceCredential(cred.tagName(), selectionScope) + effectiveThreshold, "% (weight ", credential.planWeight(), ")") + p.rebalanceCredential(credential.tagName(), selectionScope) break } } } - return cred, false, nil + return credential, false, nil } } } @@ -212,12 +212,12 @@ func (p *balancerProvider) rebalanceCredential(tag string, selectionScope creden p.sessionAccess.Unlock() } -func (p *balancerProvider) linkProviderInterrupt(cred credential, selection credentialSelection, onInterrupt func()) func() bool { +func (p *balancerProvider) linkProviderInterrupt(credential Credential, selection credentialSelection, onInterrupt func()) func() bool { if p.strategy == C.BalancerStrategyFallback { return func() bool { return false } } key := credentialInterruptKey{ - tag: cred.tagName(), + tag: credential.tagName(), selectionScope: selection.scopeOrDefault(), } p.interruptAccess.Lock() @@ -231,8 +231,8 @@ func (p *balancerProvider) linkProviderInterrupt(cred credential, selection cred return context.AfterFunc(entry.context, onInterrupt) } -func (p *balancerProvider) onRateLimited(sessionID string, cred credential, resetAt time.Time, selection credentialSelection) credential { - cred.markRateLimited(resetAt) +func (p *balancerProvider) onRateLimited(sessionID string, credential Credential, resetAt time.Time, selection credentialSelection) Credential { + credential.markRateLimited(resetAt) if p.strategy == C.BalancerStrategyFallback { return p.pickCredential(selection.filter) } @@ -255,7 +255,7 @@ func (p *balancerProvider) onRateLimited(sessionID string, cred credential, rese return best } -func (p *balancerProvider) pickCredential(filter func(credential) bool) credential { +func (p *balancerProvider) pickCredential(filter func(Credential) bool) Credential { switch p.strategy { case C.BalancerStrategyRoundRobin: return p.pickRoundRobin(filter) @@ -268,16 +268,16 @@ func (p *balancerProvider) pickCredential(filter func(credential) bool) credenti } } -func (p *balancerProvider) pickFallback(filter func(credential) bool) credential { - for _, cred := range p.credentials { - if filter != nil && !filter(cred) { +func (p *balancerProvider) pickFallback(filter func(Credential) bool) Credential { + for _, credential := range p.credentials { + if filter != nil && !filter(credential) { continue } - if !compositeCredentialSelectable(cred) { + if !compositeCredentialSelectable(credential) { continue } - if cred.isUsable() { - return cred + if credential.isUsable() { + return credential } } return nil @@ -285,23 +285,23 @@ func (p *balancerProvider) pickFallback(filter func(credential) bool) credential const weeklyWindowHours = 7 * 24 -func (p *balancerProvider) pickLeastUsed(filter func(credential) bool) credential { - var best credential +func (p *balancerProvider) pickLeastUsed(filter func(Credential) bool) Credential { + var best Credential bestScore := float64(-1) now := time.Now() - for _, cred := range p.credentials { - if filter != nil && !filter(cred) { + for _, credential := range p.credentials { + if filter != nil && !filter(credential) { continue } - if !compositeCredentialSelectable(cred) { + if !compositeCredentialSelectable(credential) { continue } - if !cred.isUsable() { + if !credential.isUsable() { continue } - remaining := cred.weeklyCap() - cred.weeklyUtilization() - score := remaining * cred.planWeight() - resetTime := cred.weeklyResetTime() + remaining := credential.weeklyCap() - credential.weeklyUtilization() + score := remaining * credential.planWeight() + resetTime := credential.weeklyResetTime() if !resetTime.IsZero() { timeUntilReset := resetTime.Sub(now) if timeUntilReset < time.Hour { @@ -311,7 +311,7 @@ func (p *balancerProvider) pickLeastUsed(filter func(credential) bool) credentia } if score > bestScore { bestScore = score - best = cred + best = credential } } return best @@ -328,7 +328,7 @@ func ocmPlanWeight(accountType string) float64 { } } -func (p *balancerProvider) pickRoundRobin(filter func(credential) bool) credential { +func (p *balancerProvider) pickRoundRobin(filter func(Credential) bool) Credential { start := int(p.roundRobinIndex.Add(1) - 1) count := len(p.credentials) for offset := range count { @@ -346,8 +346,8 @@ func (p *balancerProvider) pickRoundRobin(filter func(credential) bool) credenti return nil } -func (p *balancerProvider) pickRandom(filter func(credential) bool) credential { - var usable []credential +func (p *balancerProvider) pickRandom(filter func(Credential) bool) Credential { + var usable []Credential for _, candidate := range p.credentials { if filter != nil && !filter(candidate) { continue @@ -375,28 +375,28 @@ func (p *balancerProvider) pollIfStale(ctx context.Context) { } p.sessionAccess.Unlock() - for _, cred := range p.credentials { - if time.Since(cred.lastUpdatedTime()) > cred.pollBackoff(p.pollInterval) { - cred.pollUsage(ctx) + for _, credential := range p.credentials { + if time.Since(credential.lastUpdatedTime()) > credential.pollBackoff(p.pollInterval) { + credential.pollUsage(ctx) } } } -func (p *balancerProvider) allCredentials() []credential { +func (p *balancerProvider) allCredentials() []Credential { return p.credentials } func (p *balancerProvider) close() {} -func allRateLimitedError(credentials []credential) error { +func allRateLimitedError(credentials []Credential) error { var hasUnavailable bool var earliest time.Time - for _, cred := range credentials { - if cred.unavailableError() != nil { + for _, credential := range credentials { + if credential.unavailableError() != nil { hasUnavailable = true continue } - resetAt := cred.earliestReset() + resetAt := credential.earliestReset() if !resetAt.IsZero() && (earliest.IsZero() || resetAt.Before(earliest)) { earliest = resetAt } diff --git a/service/ocm/reverse.go b/service/ocm/reverse.go index ab99c77a6..494cb4716 100644 --- a/service/ocm/reverse.go +++ b/service/ocm/reverse.go @@ -4,7 +4,6 @@ import ( "bufio" "context" stdTLS "crypto/tls" - "errors" "io" "math/rand/v2" "net" @@ -124,13 +123,13 @@ func (s *Service) handleReverseConnect(ctx context.Context, w http.ResponseWrite } func (s *Service) findReceiverCredential(token string) *externalCredential { - for _, cred := range s.allCredentials { - extCred, ok := cred.(*externalCredential) - if !ok || extCred.connectorURL != nil { + for _, credential := range s.allCredentials { + external, ok := credential.(*externalCredential) + if !ok || external.connectorURL != nil { continue } - if extCred.token == token { - return extCred + if external.token == token { + return external } } return nil @@ -248,7 +247,7 @@ func (c *externalCredential) connectorConnect(ctx context.Context) (time.Duratio } err = httpServer.Serve(&yamuxNetListener{session: session}) sessionLifetime := time.Since(serveStart) - if err != nil && !errors.Is(err, http.ErrServerClosed) && ctx.Err() == nil { + if err != nil && !E.IsClosed(err) && ctx.Err() == nil { return sessionLifetime, E.Cause(err, "serve") } return sessionLifetime, E.New("connection closed") diff --git a/service/ocm/service.go b/service/ocm/service.go index 101f90492..272bbb3a5 100644 --- a/service/ocm/service.go +++ b/service/ocm/service.go @@ -3,7 +3,6 @@ package ocm import ( "context" "encoding/json" - "errors" "io" "net/http" "strings" @@ -68,18 +67,18 @@ const ( retryableUsageCode = "credential_usage_exhausted" ) -func hasAlternativeCredential(provider credentialProvider, currentCredential credential, selection credentialSelection) bool { +func hasAlternativeCredential(provider credentialProvider, currentCredential Credential, selection credentialSelection) bool { if provider == nil || currentCredential == nil { return false } - for _, cred := range provider.allCredentials() { - if cred == currentCredential { + for _, credential := range provider.allCredentials() { + if credential == currentCredential { continue } - if !selection.allows(cred) { + if !selection.allows(credential) { continue } - if cred.isUsable() { + if credential.isUsable() { return true } } @@ -109,7 +108,7 @@ func writeCredentialUnavailableError( w http.ResponseWriter, r *http.Request, provider credentialProvider, - currentCredential credential, + currentCredential Credential, selection credentialSelection, fallback string, ) { @@ -124,8 +123,8 @@ func credentialSelectionForUser(userConfig *option.OCMUser) credentialSelection selection := credentialSelection{scope: credentialSelectionScopeAll} if userConfig != nil && !userConfig.AllowExternalUsage { selection.scope = credentialSelectionScopeNonExternal - selection.filter = func(cred credential) bool { - return !cred.isExternal() + selection.filter = func(credential Credential) bool { + return !credential.isExternal() } } return selection @@ -174,7 +173,7 @@ type Service struct { // Multi-credential mode providers map[string]credentialProvider - allCredentials []credential + allCredentials []Credential userConfigMap map[string]*option.OCMUser } @@ -218,7 +217,7 @@ func NewService(ctx context.Context, logger log.ContextLogger, tag string, optio } service.userConfigMap = userConfigMap } else { - cred, err := newDefaultCredential(ctx, "default", option.OCMDefaultCredentialOptions{ + credential, err := newDefaultCredential(ctx, "default", option.OCMDefaultCredentialOptions{ CredentialPath: options.CredentialPath, UsagesPath: options.UsagesPath, Detour: options.Detour, @@ -226,9 +225,9 @@ func NewService(ctx context.Context, logger log.ContextLogger, tag string, optio if err != nil { return nil, err } - service.legacyCredential = cred - service.legacyProvider = &singleCredentialProvider{cred: cred} - service.allCredentials = []credential{cred} + service.legacyCredential = credential + service.legacyProvider = &singleCredentialProvider{credential: credential} + service.allCredentials = []Credential{credential} } if options.TLS != nil { @@ -249,16 +248,16 @@ func (s *Service) Start(stage adapter.StartStage) error { s.userManager.UpdateUsers(s.options.Users) - for _, cred := range s.allCredentials { - if extCred, ok := cred.(*externalCredential); ok && extCred.reverse && extCred.connectorURL != nil { - extCred.reverseService = s + for _, credential := range s.allCredentials { + if external, ok := credential.(*externalCredential); ok && external.reverse && external.connectorURL != nil { + external.reverseService = s } - err := cred.start() + err := credential.start() if err != nil { return err } - tag := cred.tagName() - cred.setOnBecameUnusable(func() { + tag := credential.tagName() + credential.setOnBecameUnusable(func() { s.interruptWebSocketSessionsForCredential(tag) }) } @@ -295,7 +294,7 @@ func (s *Service) Start(stage adapter.StartStage) error { go func() { serveErr := s.httpServer.Serve(tcpListener) - if serveErr != nil && !errors.Is(serveErr, http.ErrServerClosed) { + if serveErr != nil && !E.IsClosed(serveErr) { s.logger.Error("serve error: ", serveErr) } }() @@ -304,15 +303,15 @@ func (s *Service) Start(stage adapter.StartStage) error { } func (s *Service) InterfaceUpdated() { - for _, cred := range s.allCredentials { - extCred, ok := cred.(*externalCredential) + for _, credential := range s.allCredentials { + external, ok := credential.(*externalCredential) if !ok { continue } - if extCred.reverse && extCred.connectorURL != nil { - extCred.reverseService = s - extCred.resetReverseContext() - go extCred.connectorLoop() + if external.reverse && external.connectorURL != nil { + external.reverseService = s + external.resetReverseContext() + go external.connectorLoop() } } } @@ -330,8 +329,8 @@ func (s *Service) Close() error { } s.webSocketGroup.Wait() - for _, cred := range s.allCredentials { - cred.close() + for _, credential := range s.allCredentials { + credential.close() } return err diff --git a/service/ocm/service_handler.go b/service/ocm/service_handler.go index 9fb9c96d7..1a247d6cc 100644 --- a/service/ocm/service_handler.go +++ b/service/ocm/service_handler.go @@ -318,6 +318,9 @@ func (s *Service) ServeHTTP(w http.ResponseWriter, r *http.Request) { if n > 0 { _, writeError := w.Write(buffer[:n]) if writeError != nil { + if E.IsClosedOrCanceled(writeError) { + return + } s.logger.ErrorContext(ctx, "write streaming response: ", writeError) return } @@ -471,6 +474,9 @@ func (s *Service) handleResponseWithTracking(ctx context.Context, writer http.Re _, writeError := writer.Write(buffer[:n]) if writeError != nil { + if E.IsClosedOrCanceled(writeError) { + return + } s.logger.ErrorContext(ctx, "write streaming response: ", writeError) return } diff --git a/service/ocm/service_status.go b/service/ocm/service_status.go index 29b95d063..915fb837d 100644 --- a/service/ocm/service_status.go +++ b/service/ocm/service_status.go @@ -62,22 +62,22 @@ func (s *Service) handleStatusEndpoint(w http.ResponseWriter, r *http.Request) { func (s *Service) computeAggregatedUtilization(provider credentialProvider, userConfig *option.OCMUser) (float64, float64, float64) { var totalWeightedRemaining5h, totalWeightedRemainingWeekly, totalWeight float64 - for _, cred := range provider.allCredentials() { - if !cred.isAvailable() { + for _, credential := range provider.allCredentials() { + if !credential.isAvailable() { continue } - if userConfig.ExternalCredential != "" && cred.tagName() == userConfig.ExternalCredential { + if userConfig.ExternalCredential != "" && credential.tagName() == userConfig.ExternalCredential { continue } - if !userConfig.AllowExternalUsage && cred.isExternal() { + if !userConfig.AllowExternalUsage && credential.isExternal() { continue } - weight := cred.planWeight() - remaining5h := cred.fiveHourCap() - cred.fiveHourUtilization() + weight := credential.planWeight() + remaining5h := credential.fiveHourCap() - credential.fiveHourUtilization() if remaining5h < 0 { remaining5h = 0 } - remainingWeekly := cred.weeklyCap() - cred.weeklyUtilization() + remainingWeekly := credential.weeklyCap() - credential.weeklyUtilization() if remainingWeekly < 0 { remainingWeekly = 0 } diff --git a/service/ocm/service_usage.go b/service/ocm/service_usage.go index 589fd093a..19a853a7c 100644 --- a/service/ocm/service_usage.go +++ b/service/ocm/service_usage.go @@ -55,13 +55,13 @@ type CostCombination struct { type AggregatedUsage struct { LastUpdated time.Time `json:"last_updated"` Combinations []CostCombination `json:"combinations"` - mutex sync.Mutex + access sync.Mutex filePath string logger log.ContextLogger lastSaveTime time.Time pendingSave bool saveTimer *time.Timer - saveMutex sync.Mutex + saveAccess sync.Mutex } type UsageStatsJSON struct { @@ -1035,8 +1035,8 @@ func deriveWeekStartUnix(cycleHint *WeeklyCycleHint) int64 { } func (u *AggregatedUsage) ToJSON() *AggregatedUsageJSON { - u.mutex.Lock() - defer u.mutex.Unlock() + u.access.Lock() + defer u.access.Unlock() result := &AggregatedUsageJSON{ LastUpdated: u.LastUpdated, @@ -1069,8 +1069,8 @@ func (u *AggregatedUsage) ToJSON() *AggregatedUsageJSON { } func (u *AggregatedUsage) Load() error { - u.mutex.Lock() - defer u.mutex.Unlock() + u.access.Lock() + defer u.access.Unlock() u.LastUpdated = time.Time{} u.Combinations = nil @@ -1116,9 +1116,9 @@ func (u *AggregatedUsage) Save() error { defer os.Remove(tmpFile) err = os.Rename(tmpFile, u.filePath) if err == nil { - u.saveMutex.Lock() + u.saveAccess.Lock() u.lastSaveTime = time.Now() - u.saveMutex.Unlock() + u.saveAccess.Unlock() } return err } @@ -1140,8 +1140,8 @@ func (u *AggregatedUsage) AddUsageWithCycleHint(model string, contextWindow int, observedAt = time.Now() } - u.mutex.Lock() - defer u.mutex.Unlock() + u.access.Lock() + defer u.access.Unlock() u.LastUpdated = observedAt weekStartUnix := deriveWeekStartUnix(cycleHint) @@ -1156,8 +1156,8 @@ func (u *AggregatedUsage) AddUsageWithCycleHint(model string, contextWindow int, func (u *AggregatedUsage) scheduleSave() { const saveInterval = time.Minute - u.saveMutex.Lock() - defer u.saveMutex.Unlock() + u.saveAccess.Lock() + defer u.saveAccess.Unlock() timeSinceLastSave := time.Since(u.lastSaveTime) @@ -1174,9 +1174,9 @@ func (u *AggregatedUsage) scheduleSave() { remainingTime := saveInterval - timeSinceLastSave u.saveTimer = time.AfterFunc(remainingTime, func() { - u.saveMutex.Lock() + u.saveAccess.Lock() u.pendingSave = false - u.saveMutex.Unlock() + u.saveAccess.Unlock() u.saveAsync() }) } @@ -1191,8 +1191,8 @@ func (u *AggregatedUsage) saveAsync() { } func (u *AggregatedUsage) cancelPendingSave() { - u.saveMutex.Lock() - defer u.saveMutex.Unlock() + u.saveAccess.Lock() + defer u.saveAccess.Unlock() if u.saveTimer != nil { u.saveTimer.Stop() diff --git a/service/ocm/service_user.go b/service/ocm/service_user.go index b69655e9a..5f7680837 100644 --- a/service/ocm/service_user.go +++ b/service/ocm/service_user.go @@ -7,8 +7,8 @@ import ( ) type UserManager struct { - access sync.RWMutex - tokenMap map[string]string + access sync.RWMutex + tokenMap map[string]string } func (m *UserManager) UpdateUsers(users []option.OCMUser) { diff --git a/service/ocm/service_websocket.go b/service/ocm/service_websocket.go index 4b640d9c5..1b7b5baac 100644 --- a/service/ocm/service_websocket.go +++ b/service/ocm/service_websocket.go @@ -98,7 +98,7 @@ func (s *Service) handleWebSocket( sessionID string, userConfig *option.OCMUser, provider credentialProvider, - selectedCredential credential, + selectedCredential Credential, selection credentialSelection, isNew bool, ) { @@ -307,7 +307,7 @@ func (s *Service) handleWebSocket( waitGroup.Wait() } -func (s *Service) proxyWebSocketClientToUpstream(ctx context.Context, clientConn net.Conn, upstreamConn net.Conn, selectedCredential credential, modelChannel chan<- string, isNew bool, username string, sessionID string) { +func (s *Service) proxyWebSocketClientToUpstream(ctx context.Context, clientConn net.Conn, upstreamConn net.Conn, selectedCredential Credential, modelChannel chan<- string, isNew bool, username string, sessionID string) { logged := false for { data, opCode, err := wsutil.ReadClientData(clientConn) @@ -359,7 +359,7 @@ func (s *Service) proxyWebSocketClientToUpstream(ctx context.Context, clientConn } } -func (s *Service) proxyWebSocketUpstreamToClient(ctx context.Context, upstreamReadWriter io.ReadWriter, clientConn net.Conn, selectedCredential credential, userConfig *option.OCMUser, provider credentialProvider, modelChannel <-chan string, username string, weeklyCycleHint *WeeklyCycleHint) { +func (s *Service) proxyWebSocketUpstreamToClient(ctx context.Context, upstreamReadWriter io.ReadWriter, clientConn net.Conn, selectedCredential Credential, userConfig *option.OCMUser, provider credentialProvider, modelChannel <-chan string, username string, weeklyCycleHint *WeeklyCycleHint) { usageTracker := selectedCredential.usageTrackerOrNil() var requestModel string for { @@ -413,7 +413,7 @@ func (s *Service) proxyWebSocketUpstreamToClient(ctx context.Context, upstreamRe } } -func (s *Service) handleWebSocketRateLimitsEvent(data []byte, selectedCredential credential) { +func (s *Service) handleWebSocketRateLimitsEvent(data []byte, selectedCredential Credential) { var rateLimitsEvent struct { RateLimits struct { Primary *struct { @@ -462,7 +462,7 @@ func (s *Service) handleWebSocketRateLimitsEvent(data []byte, selectedCredential selectedCredential.updateStateFromHeaders(headers) } -func (s *Service) handleWebSocketErrorRateLimited(data []byte, selectedCredential credential) { +func (s *Service) handleWebSocketErrorRateLimited(data []byte, selectedCredential Credential) { var errorEvent struct { Headers map[string]string `json:"headers"` }