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:
@@ -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 {
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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"`
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user