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
This commit is contained in:
世界
2026-03-14 21:06:25 +08:00
parent 04bd63b455
commit 6878ad0d35
21 changed files with 448 additions and 439 deletions
+4 -4
View File
@@ -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 {
+50 -50
View File
@@ -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")
}
}
}
+26 -26
View File
@@ -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
+62 -62
View File
@@ -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
}
+6 -7
View File
@@ -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")
+26 -27
View File
@@ -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
+7
View File
@@ -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
}
+7 -7
View File
@@ -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
}
+16 -16
View File
@@ -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()
+2 -2
View File
@@ -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) {
+4 -4
View File
@@ -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 {
+55 -55
View File
@@ -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",
)
+47 -47
View File
@@ -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
+66 -66
View File
@@ -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
}
+6 -7
View File
@@ -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")
+28 -29
View File
@@ -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
+6
View File
@@ -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
}
+7 -7
View File
@@ -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
}
+16 -16
View File
@@ -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()
+2 -2
View File
@@ -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) {
+5 -5
View File
@@ -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"`
}