mirror of
https://github.com/router-for-me/CLIProxyAPI.git
synced 2026-10-08 08:40:44 +08:00
fix(meta): validate catalog section, singleflight management mints, and synchronize singleflight test
This commit is contained in:
@@ -16,6 +16,7 @@ import (
|
||||
coreauth "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth"
|
||||
"github.com/router-for-me/CLIProxyAPI/v7/sdk/proxyutil"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.org/x/sync/singleflight"
|
||||
)
|
||||
|
||||
const defaultAPICallTimeout = 60 * time.Second
|
||||
@@ -359,6 +360,31 @@ func (h *Handler) refreshAntigravityOAuthAccessToken(ctx context.Context, auth *
|
||||
return strings.TrimSpace(tokenResp.AccessToken), nil
|
||||
}
|
||||
|
||||
var metaManagementMintGroup singleflight.Group
|
||||
|
||||
func metaTokenFromAuth(auth *coreauth.Auth) string {
|
||||
if auth == nil {
|
||||
return ""
|
||||
}
|
||||
if auth.Metadata != nil {
|
||||
if k, ok := auth.Metadata["api_key"].(string); ok && strings.TrimSpace(k) != "" && !strings.HasPrefix(strings.TrimSpace(k), "dca:") {
|
||||
return strings.TrimSpace(k)
|
||||
}
|
||||
if t, ok := auth.Metadata["access_token"].(string); ok && strings.TrimSpace(t) != "" && !strings.HasPrefix(strings.TrimSpace(t), "dca:") {
|
||||
return strings.TrimSpace(t)
|
||||
}
|
||||
}
|
||||
if auth.Attributes != nil {
|
||||
if k := strings.TrimSpace(auth.Attributes["api_key"]); k != "" && !strings.HasPrefix(k, "dca:") {
|
||||
return k
|
||||
}
|
||||
if t := strings.TrimSpace(auth.Attributes["access_token"]); t != "" && !strings.HasPrefix(t, "dca:") {
|
||||
return t
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (h *Handler) resolveMetaToken(ctx context.Context, auth *coreauth.Auth, requestProxyURL string) (string, error) {
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
@@ -366,21 +392,8 @@ func (h *Handler) resolveMetaToken(ctx context.Context, auth *coreauth.Auth, req
|
||||
if auth == nil {
|
||||
return "", nil
|
||||
}
|
||||
if auth.Metadata != nil {
|
||||
if k, ok := auth.Metadata["api_key"].(string); ok && strings.TrimSpace(k) != "" && !strings.HasPrefix(strings.TrimSpace(k), "dca:") {
|
||||
return strings.TrimSpace(k), nil
|
||||
}
|
||||
if t, ok := auth.Metadata["access_token"].(string); ok && strings.TrimSpace(t) != "" && !strings.HasPrefix(strings.TrimSpace(t), "dca:") {
|
||||
return strings.TrimSpace(t), nil
|
||||
}
|
||||
}
|
||||
if auth.Attributes != nil {
|
||||
if k := strings.TrimSpace(auth.Attributes["api_key"]); k != "" && !strings.HasPrefix(k, "dca:") {
|
||||
return k, nil
|
||||
}
|
||||
if t := strings.TrimSpace(auth.Attributes["access_token"]); t != "" && !strings.HasPrefix(t, "dca:") {
|
||||
return t, nil
|
||||
}
|
||||
if tok := metaTokenFromAuth(auth); tok != "" {
|
||||
return tok, nil
|
||||
}
|
||||
|
||||
var dcaToken string
|
||||
@@ -405,101 +418,157 @@ func (h *Handler) resolveMetaToken(ctx context.Context, auth *coreauth.Auth, req
|
||||
return "", nil
|
||||
}
|
||||
|
||||
proxyURL := firstNonEmptyString(&requestProxyURL, &auth.ProxyURL)
|
||||
var cfg *config.Config
|
||||
if h != nil {
|
||||
cfg = h.cfg
|
||||
}
|
||||
authSvc := metaauth.NewMetaAuthWithProxyURL(cfg, proxyURL)
|
||||
minted, err := authSvc.MintAPIKey(ctx, dcaToken)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("meta token mint failed: %w", err)
|
||||
}
|
||||
if minted == nil || minted.APIKey == "" {
|
||||
return "", fmt.Errorf("meta token mint returned empty key")
|
||||
}
|
||||
flightKey := firstNonEmptyString(&auth.ID, &dcaToken)
|
||||
|
||||
base := auth.Clone()
|
||||
baseURL := ""
|
||||
if auth.Attributes != nil {
|
||||
baseURL = strings.TrimSpace(auth.Attributes["base_url"])
|
||||
}
|
||||
if baseURL == "" {
|
||||
baseURL = stringValue(auth.Metadata, "base_url")
|
||||
}
|
||||
if baseURL == "" {
|
||||
baseURL = stringValue(auth.Metadata, "api_base_url")
|
||||
}
|
||||
if baseURL == "" {
|
||||
baseURL = metaauth.DefaultAPIBaseURL
|
||||
}
|
||||
if mintedURL := strings.TrimSpace(minted.BaseURL); mintedURL != "" {
|
||||
baseURL = mintedURL
|
||||
}
|
||||
if auth.Metadata == nil {
|
||||
auth.Metadata = make(map[string]any)
|
||||
}
|
||||
auth.Metadata["base_url"] = baseURL
|
||||
auth.Metadata["api_key"] = minted.APIKey
|
||||
auth.Metadata["access_token"] = minted.APIKey
|
||||
auth.Metadata["dca_token"] = dcaToken
|
||||
delete(auth.Metadata, "expired")
|
||||
if minted.UserEmail != "" {
|
||||
auth.Metadata["email"] = minted.UserEmail
|
||||
}
|
||||
if minted.UserFullName != "" {
|
||||
auth.Metadata["name"] = minted.UserFullName
|
||||
}
|
||||
now := time.Now()
|
||||
nowStr := now.Format(time.RFC3339)
|
||||
auth.LastRefreshedAt = now
|
||||
auth.UpdatedAt = now
|
||||
auth.Metadata["last_refresh"] = nowStr
|
||||
|
||||
if auth.Attributes == nil {
|
||||
auth.Attributes = make(map[string]string)
|
||||
}
|
||||
auth.Attributes["base_url"] = baseURL
|
||||
auth.Attributes["api_key"] = minted.APIKey
|
||||
auth.Attributes["access_token"] = minted.APIKey
|
||||
|
||||
storage := &metaauth.MetaTokenStorage{Type: "meta", AuthKind: "oauth"}
|
||||
if existing, ok := auth.Storage.(*metaauth.MetaTokenStorage); ok && existing != nil {
|
||||
copyStorage := *existing
|
||||
storage = ©Storage
|
||||
}
|
||||
storage.APIKey = minted.APIKey
|
||||
storage.AccessToken = minted.APIKey
|
||||
storage.DCAToken = dcaToken
|
||||
storage.Expired = ""
|
||||
storage.BaseURL = baseURL
|
||||
storage.LastRefresh = nowStr
|
||||
storage.Metadata = auth.Metadata
|
||||
if minted.UserEmail != "" {
|
||||
storage.Email = minted.UserEmail
|
||||
}
|
||||
if minted.UserFullName != "" {
|
||||
storage.Name = minted.UserFullName
|
||||
}
|
||||
auth.Storage = storage
|
||||
|
||||
// Use the configured backend; direct file writes bypass remote token stores.
|
||||
if !coreauth.IsConfigAPIKeyAuth(auth) {
|
||||
store := h.tokenStoreWithBaseDir()
|
||||
if store == nil {
|
||||
return "", fmt.Errorf("meta token store unavailable")
|
||||
mintRes, errMint, _ := metaManagementMintGroup.Do(flightKey, func() (any, error) {
|
||||
if h != nil && h.authManager != nil {
|
||||
var latest *coreauth.Auth
|
||||
if auth.ID != "" {
|
||||
if a, ok := h.authManager.GetByID(auth.ID); ok && a != nil {
|
||||
latest = a
|
||||
}
|
||||
}
|
||||
if latest == nil && auth.Index != "" {
|
||||
latest = h.authByIndex(auth.Index)
|
||||
}
|
||||
if latest != nil {
|
||||
if k := metaTokenFromAuth(latest); k != "" {
|
||||
return k, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
if _, errSave := store.Save(ctx, auth); errSave != nil {
|
||||
return "", fmt.Errorf("persist meta token: %w", errSave)
|
||||
|
||||
proxyURL := firstNonEmptyString(&requestProxyURL, &auth.ProxyURL)
|
||||
var cfg *config.Config
|
||||
if h != nil {
|
||||
cfg = h.cfg
|
||||
}
|
||||
authSvc := metaauth.NewMetaAuthWithProxyURL(cfg, proxyURL)
|
||||
minted, err := authSvc.MintAPIKey(ctx, dcaToken)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("meta token mint failed: %w", err)
|
||||
}
|
||||
if minted == nil || minted.APIKey == "" {
|
||||
return "", fmt.Errorf("meta token mint returned empty key")
|
||||
}
|
||||
|
||||
base := auth.Clone()
|
||||
baseURL := ""
|
||||
if auth.Attributes != nil {
|
||||
baseURL = strings.TrimSpace(auth.Attributes["base_url"])
|
||||
}
|
||||
if baseURL == "" {
|
||||
baseURL = stringValue(auth.Metadata, "base_url")
|
||||
}
|
||||
if baseURL == "" {
|
||||
baseURL = stringValue(auth.Metadata, "api_base_url")
|
||||
}
|
||||
if baseURL == "" {
|
||||
baseURL = metaauth.DefaultAPIBaseURL
|
||||
}
|
||||
if mintedURL := strings.TrimSpace(minted.BaseURL); mintedURL != "" {
|
||||
baseURL = mintedURL
|
||||
}
|
||||
if auth.Metadata == nil {
|
||||
auth.Metadata = make(map[string]any)
|
||||
}
|
||||
auth.Metadata["base_url"] = baseURL
|
||||
auth.Metadata["api_key"] = minted.APIKey
|
||||
auth.Metadata["access_token"] = minted.APIKey
|
||||
auth.Metadata["dca_token"] = dcaToken
|
||||
delete(auth.Metadata, "expired")
|
||||
if minted.UserEmail != "" {
|
||||
auth.Metadata["email"] = minted.UserEmail
|
||||
}
|
||||
if minted.UserFullName != "" {
|
||||
auth.Metadata["name"] = minted.UserFullName
|
||||
}
|
||||
now := time.Now()
|
||||
nowStr := now.Format(time.RFC3339)
|
||||
auth.LastRefreshedAt = now
|
||||
auth.UpdatedAt = now
|
||||
auth.Metadata["last_refresh"] = nowStr
|
||||
|
||||
if auth.Attributes == nil {
|
||||
auth.Attributes = make(map[string]string)
|
||||
}
|
||||
auth.Attributes["base_url"] = baseURL
|
||||
auth.Attributes["api_key"] = minted.APIKey
|
||||
auth.Attributes["access_token"] = minted.APIKey
|
||||
|
||||
storage := &metaauth.MetaTokenStorage{Type: "meta", AuthKind: "oauth"}
|
||||
if existing, ok := auth.Storage.(*metaauth.MetaTokenStorage); ok && existing != nil {
|
||||
copyStorage := *existing
|
||||
storage = ©Storage
|
||||
}
|
||||
storage.APIKey = minted.APIKey
|
||||
storage.AccessToken = minted.APIKey
|
||||
storage.DCAToken = dcaToken
|
||||
storage.Expired = ""
|
||||
storage.BaseURL = baseURL
|
||||
storage.LastRefresh = nowStr
|
||||
storage.Metadata = auth.Metadata
|
||||
if minted.UserEmail != "" {
|
||||
storage.Email = minted.UserEmail
|
||||
}
|
||||
if minted.UserFullName != "" {
|
||||
storage.Name = minted.UserFullName
|
||||
}
|
||||
auth.Storage = storage
|
||||
|
||||
// Use the configured backend; direct file writes bypass remote token stores.
|
||||
if !coreauth.IsConfigAPIKeyAuth(auth) {
|
||||
store := h.tokenStoreWithBaseDir()
|
||||
if store == nil {
|
||||
return "", fmt.Errorf("meta token store unavailable")
|
||||
}
|
||||
if _, errSave := store.Save(ctx, auth); errSave != nil {
|
||||
return "", fmt.Errorf("persist meta token: %w", errSave)
|
||||
}
|
||||
}
|
||||
if h != nil && h.authManager != nil {
|
||||
if _, errUpdate := h.authManager.UpdateRefreshedAuth(coreauth.WithSkipPersist(ctx), base, auth); errUpdate != nil {
|
||||
return "", fmt.Errorf("update meta auth: %w", errUpdate)
|
||||
}
|
||||
}
|
||||
|
||||
return minted.APIKey, nil
|
||||
})
|
||||
if errMint != nil {
|
||||
return "", errMint
|
||||
}
|
||||
|
||||
key, _ := mintRes.(string)
|
||||
if h != nil && h.authManager != nil {
|
||||
if _, errUpdate := h.authManager.UpdateRefreshedAuth(coreauth.WithSkipPersist(ctx), base, auth); errUpdate != nil {
|
||||
return "", fmt.Errorf("update meta auth: %w", errUpdate)
|
||||
var latest *coreauth.Auth
|
||||
if auth.ID != "" {
|
||||
if a, ok := h.authManager.GetByID(auth.ID); ok && a != nil {
|
||||
latest = a
|
||||
}
|
||||
}
|
||||
if latest == nil && auth.Index != "" {
|
||||
latest = h.authByIndex(auth.Index)
|
||||
}
|
||||
if latest != nil {
|
||||
if auth.Metadata == nil {
|
||||
auth.Metadata = make(map[string]any)
|
||||
}
|
||||
for k, v := range latest.Metadata {
|
||||
auth.Metadata[k] = v
|
||||
}
|
||||
if auth.Attributes == nil {
|
||||
auth.Attributes = make(map[string]string)
|
||||
}
|
||||
for k, v := range latest.Attributes {
|
||||
auth.Attributes[k] = v
|
||||
}
|
||||
auth.Storage = latest.Storage
|
||||
auth.LastRefreshedAt = latest.LastRefreshedAt
|
||||
auth.UpdatedAt = latest.UpdatedAt
|
||||
}
|
||||
}
|
||||
|
||||
return minted.APIKey, nil
|
||||
return key, nil
|
||||
}
|
||||
|
||||
func antigravityTokenNeedsRefresh(metadata map[string]any) bool {
|
||||
|
||||
@@ -6,7 +6,10 @@ import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"runtime"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -474,3 +477,82 @@ func TestResolveMetaTokenPropagatesStoreFailure(t *testing.T) {
|
||||
t.Error("failed save installed a token in the manager")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveMetaToken_ConcurrentSingleflight(t *testing.T) {
|
||||
var mints int64
|
||||
serverStarted := make(chan struct{})
|
||||
releaseServer := make(chan struct{})
|
||||
var startOnce sync.Once
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
atomic.AddInt64(&mints, 1)
|
||||
startOnce.Do(func() { close(serverStarted) })
|
||||
<-releaseServer
|
||||
_ = json.NewEncoder(w).Encode(map[string]string{
|
||||
"api_key": "LLM|minted-concurrent",
|
||||
"base_url": "https://regional.meta.example/v1",
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
t.Setenv("META_MINT_URL", server.URL)
|
||||
store := &memoryAuthStore{}
|
||||
manager := coreauth.NewManager(store, nil, nil)
|
||||
auth, err := manager.Register(coreauth.WithSkipPersist(context.Background()), &coreauth.Auth{
|
||||
ID: "meta-mgmt-concurrent",
|
||||
Provider: "meta",
|
||||
Metadata: map[string]any{"dca_token": "dca:concurrent-test"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
h := &Handler{cfg: &config.Config{}, authManager: manager, tokenStore: store}
|
||||
|
||||
const callers = 5
|
||||
var entryWg sync.WaitGroup
|
||||
entryWg.Add(callers)
|
||||
ready := make(chan struct{})
|
||||
var inFlight sync.WaitGroup
|
||||
inFlight.Add(callers)
|
||||
var doneWg sync.WaitGroup
|
||||
doneWg.Add(callers)
|
||||
|
||||
for i := 0; i < callers; i++ {
|
||||
go func() {
|
||||
defer doneWg.Done()
|
||||
entryWg.Done()
|
||||
<-ready
|
||||
inFlight.Done()
|
||||
token, errToken := h.resolveTokenForAuth(context.Background(), h.authByIndex(auth.Index), "")
|
||||
if errToken != nil {
|
||||
t.Errorf("resolveTokenForAuth error: %v", errToken)
|
||||
}
|
||||
if token != "LLM|minted-concurrent" {
|
||||
t.Errorf("expected LLM|minted-concurrent, got %q", token)
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
entryWg.Wait()
|
||||
close(ready)
|
||||
|
||||
<-serverStarted
|
||||
inFlight.Wait()
|
||||
for i := 0; i < 50; i++ {
|
||||
runtime.Gosched()
|
||||
}
|
||||
close(releaseServer)
|
||||
doneWg.Wait()
|
||||
|
||||
if totalMints := atomic.LoadInt64(&mints); totalMints != 1 {
|
||||
t.Errorf("expected 1 mint request, got %d", totalMints)
|
||||
}
|
||||
live := h.authByIndex(auth.Index)
|
||||
if live.Metadata["api_key"] != "LLM|minted-concurrent" {
|
||||
t.Error("live manager did not retain minted key")
|
||||
}
|
||||
if live.Attributes["api_key"] != "LLM|minted-concurrent" {
|
||||
t.Error("live manager did not retain minted key in attributes")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -111,3 +111,38 @@ func TestAntigravityWebSearchModelForRequiresRequestedModelCapability(t *testing
|
||||
t.Fatalf("unknown model should not get Antigravity web search model, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateModelsCatalog_Meta(t *testing.T) {
|
||||
valid := &staticModelsJSON{
|
||||
Meta: []*ModelInfo{
|
||||
{ID: "muse-spark-1.3"},
|
||||
},
|
||||
}
|
||||
if err := validateModelsCatalog(valid); err != nil {
|
||||
t.Fatalf("expected valid Meta catalog to pass, got: %v", err)
|
||||
}
|
||||
|
||||
withNull := &staticModelsJSON{
|
||||
Meta: []*ModelInfo{nil},
|
||||
}
|
||||
if err := validateModelsCatalog(withNull); err == nil {
|
||||
t.Fatal("expected error for Meta section with null model, got nil")
|
||||
}
|
||||
|
||||
withEmptyID := &staticModelsJSON{
|
||||
Meta: []*ModelInfo{{ID: " "}},
|
||||
}
|
||||
if err := validateModelsCatalog(withEmptyID); err == nil {
|
||||
t.Fatal("expected error for Meta section with empty model id, got nil")
|
||||
}
|
||||
|
||||
withDuplicate := &staticModelsJSON{
|
||||
Meta: []*ModelInfo{
|
||||
{ID: "muse-spark-1.3"},
|
||||
{ID: "muse-spark-1.3"},
|
||||
},
|
||||
}
|
||||
if err := validateModelsCatalog(withDuplicate); err == nil {
|
||||
t.Fatal("expected error for Meta section with duplicate model id, got nil")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -341,6 +341,7 @@ func validateModelsCatalog(data *staticModelsJSON) error {
|
||||
{name: "kimi", models: data.Kimi},
|
||||
{name: "antigravity", models: data.Antigravity},
|
||||
{name: "xai", models: data.XAI},
|
||||
{name: "meta", models: data.Meta},
|
||||
}
|
||||
|
||||
for _, section := range requiredSections {
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
@@ -272,6 +273,10 @@ func TestMetaExecutor_Refresh_DCA_MintAndPersist(t *testing.T) {
|
||||
func TestMetaExecutor_Refresh_SingleflightAndMultiAccount(t *testing.T) {
|
||||
var count1, count2 int64
|
||||
|
||||
serverStarted := make(chan struct{})
|
||||
releaseServer := make(chan struct{})
|
||||
var startOnce sync.Once
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == "/key" {
|
||||
var req map[string]string
|
||||
@@ -280,7 +285,8 @@ func TestMetaExecutor_Refresh_SingleflightAndMultiAccount(t *testing.T) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
if dca == "dca:acct1" {
|
||||
atomic.AddInt64(&count1, 1)
|
||||
time.Sleep(50 * time.Millisecond) // artificial delay to allow concurrent calls to coalesce
|
||||
startOnce.Do(func() { close(serverStarted) })
|
||||
<-releaseServer
|
||||
_ = json.NewEncoder(w).Encode(map[string]string{
|
||||
"api_key": "LLM|key-acct1",
|
||||
"user_email": "acct1@meta.com",
|
||||
@@ -310,10 +316,20 @@ func TestMetaExecutor_Refresh_SingleflightAndMultiAccount(t *testing.T) {
|
||||
Metadata: map[string]any{"dca_token": "dca:acct1"},
|
||||
}
|
||||
|
||||
for i := 0; i < 10; i++ {
|
||||
const goroutines = 10
|
||||
var entryWg sync.WaitGroup
|
||||
entryWg.Add(goroutines)
|
||||
ready := make(chan struct{})
|
||||
var inFlight sync.WaitGroup
|
||||
inFlight.Add(goroutines)
|
||||
|
||||
for i := 0; i < goroutines; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
entryWg.Done()
|
||||
<-ready
|
||||
inFlight.Done()
|
||||
res, err := exec.Refresh(context.Background(), auth1.Clone())
|
||||
if err != nil {
|
||||
t.Errorf("Refresh acct1 error: %v", err)
|
||||
@@ -323,6 +339,22 @@ func TestMetaExecutor_Refresh_SingleflightAndMultiAccount(t *testing.T) {
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// Release all goroutines simultaneously.
|
||||
entryWg.Wait()
|
||||
close(ready)
|
||||
|
||||
// Wait until the singleflight mint request has arrived at the server.
|
||||
<-serverStarted
|
||||
inFlight.Wait()
|
||||
|
||||
// Allow pending goroutines to enter singleflight.Do while the server holds the in-flight request.
|
||||
for i := 0; i < 50; i++ {
|
||||
runtime.Gosched()
|
||||
}
|
||||
|
||||
// Release the HTTP server handler to complete the single in-flight mint.
|
||||
close(releaseServer)
|
||||
wg.Wait()
|
||||
|
||||
if totalMint1 := atomic.LoadInt64(&count1); totalMint1 != 1 {
|
||||
|
||||
Reference in New Issue
Block a user