Merge pull request #2881 from xiaoyiluck666/fix/openai-oauth-refresh-enrichment

修复 OpenAI OAuth 刷新未补全账号信息
This commit is contained in:
Wesley Liddick 2026-06-01 10:04:16 +08:00 committed by GitHub
commit 418d09be56
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
6 changed files with 164 additions and 26 deletions

View File

@ -137,7 +137,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
httpUpstream := repository.NewHTTPUpstream(configConfig) httpUpstream := repository.NewHTTPUpstream(configConfig)
deferredService := service.ProvideDeferredService(accountRepository, timingWheelService) deferredService := service.ProvideDeferredService(accountRepository, timingWheelService)
openAIOAuthClient := repository.NewOpenAIOAuthClient() openAIOAuthClient := repository.NewOpenAIOAuthClient()
openAIOAuthService := service.NewOpenAIOAuthService(proxyRepository, openAIOAuthClient) openAIOAuthService := service.ProvideOpenAIOAuthService(proxyRepository, openAIOAuthClient, privacyClientFactory)
oAuthRefreshAPI := service.ProvideOAuthRefreshAPI(accountRepository, geminiTokenCache) oAuthRefreshAPI := service.ProvideOAuthRefreshAPI(accountRepository, geminiTokenCache)
openAITokenProvider := service.ProvideOpenAITokenProvider(accountRepository, geminiTokenCache, openAIOAuthService, oAuthRefreshAPI) openAITokenProvider := service.ProvideOpenAITokenProvider(accountRepository, geminiTokenCache, openAIOAuthService, oAuthRefreshAPI)
channelRepository := repository.NewChannelRepository(db) channelRepository := repository.NewChannelRepository(db)

View File

@ -278,11 +278,29 @@ func (s *OpenAIOAuthService) enrichTokenInfo(ctx context.Context, tokenInfo *Ope
tokenInfo.Email = info.Email tokenInfo.Email = info.Email
} }
} }
if strings.TrimSpace(tokenInfo.SubscriptionExpiresAt) == "" {
if expiresAt := fetchChatGPTSubscriptionExpiresAt(ctx, s.privacyClientFactory, tokenInfo.AccessToken, proxyURL, resolveChatGPTSubscriptionAccountID(tokenInfo, orgID)); expiresAt != "" {
tokenInfo.SubscriptionExpiresAt = expiresAt
}
}
// 尝试设置隐私关闭训练数据共享best-effort // 尝试设置隐私关闭训练数据共享best-effort
tokenInfo.PrivacyMode = disableOpenAITraining(ctx, s.privacyClientFactory, tokenInfo.AccessToken, proxyURL) tokenInfo.PrivacyMode = disableOpenAITraining(ctx, s.privacyClientFactory, tokenInfo.AccessToken, proxyURL)
} }
func resolveChatGPTSubscriptionAccountID(tokenInfo *OpenAITokenInfo, orgID string) string {
for _, candidate := range []string{
tokenInfo.ChatGPTAccountID,
tokenInfo.OrganizationID,
orgID,
} {
if trimmed := strings.TrimSpace(candidate); trimmed != "" {
return trimmed
}
}
return ""
}
// RefreshAccountToken refreshes token for an OpenAI OAuth account // RefreshAccountToken refreshes token for an OpenAI OAuth account
func (s *OpenAIOAuthService) RefreshAccountToken(ctx context.Context, account *Account) (*OpenAITokenInfo, error) { func (s *OpenAIOAuthService) RefreshAccountToken(ctx context.Context, account *Account) (*OpenAITokenInfo, error) {
if account.Platform != PlatformOpenAI { if account.Platform != PlatformOpenAI {
@ -292,30 +310,6 @@ func (s *OpenAIOAuthService) RefreshAccountToken(ctx context.Context, account *A
return nil, infraerrors.New(http.StatusBadRequest, "OPENAI_OAUTH_INVALID_ACCOUNT_TYPE", "account is not an OAuth account") return nil, infraerrors.New(http.StatusBadRequest, "OPENAI_OAUTH_INVALID_ACCOUNT_TYPE", "account is not an OAuth account")
} }
refreshToken := account.GetCredential("refresh_token")
if refreshToken == "" {
accessToken := account.GetCredential("access_token")
if accessToken != "" {
tokenInfo := &OpenAITokenInfo{
AccessToken: accessToken,
RefreshToken: "",
IDToken: account.GetCredential("id_token"),
ClientID: account.GetCredential("client_id"),
Email: account.GetCredential("email"),
ChatGPTAccountID: account.GetCredential("chatgpt_account_id"),
ChatGPTUserID: account.GetCredential("chatgpt_user_id"),
OrganizationID: account.GetCredential("organization_id"),
PlanType: account.GetCredential("plan_type"),
}
if expiresAt := account.GetCredentialAsTime("expires_at"); expiresAt != nil {
tokenInfo.ExpiresAt = expiresAt.Unix()
tokenInfo.ExpiresIn = int64(time.Until(*expiresAt).Seconds())
}
return tokenInfo, nil
}
return nil, infraerrors.New(http.StatusBadRequest, "OPENAI_OAUTH_NO_REFRESH_TOKEN", "no refresh token available")
}
var proxyURL string var proxyURL string
if account.ProxyID != nil { if account.ProxyID != nil {
proxy, err := s.proxyRepo.GetByID(ctx, *account.ProxyID) proxy, err := s.proxyRepo.GetByID(ctx, *account.ProxyID)
@ -324,6 +318,32 @@ func (s *OpenAIOAuthService) RefreshAccountToken(ctx context.Context, account *A
} }
} }
refreshToken := account.GetCredential("refresh_token")
if refreshToken == "" {
accessToken := account.GetCredential("access_token")
if accessToken != "" {
tokenInfo := &OpenAITokenInfo{
AccessToken: accessToken,
RefreshToken: "",
IDToken: account.GetCredential("id_token"),
ClientID: account.GetCredential("client_id"),
Email: account.GetCredential("email"),
ChatGPTAccountID: account.GetCredential("chatgpt_account_id"),
ChatGPTUserID: account.GetCredential("chatgpt_user_id"),
OrganizationID: account.GetCredential("organization_id"),
PlanType: account.GetCredential("plan_type"),
SubscriptionExpiresAt: account.GetCredential("subscription_expires_at"),
}
if expiresAt := account.GetCredentialAsTime("expires_at"); expiresAt != nil {
tokenInfo.ExpiresAt = expiresAt.Unix()
tokenInfo.ExpiresIn = int64(time.Until(*expiresAt).Seconds())
}
s.enrichTokenInfo(ctx, tokenInfo, proxyURL)
return tokenInfo, nil
}
return nil, infraerrors.New(http.StatusBadRequest, "OPENAI_OAUTH_NO_REFRESH_TOKEN", "no refresh token available")
}
clientID := account.GetCredential("client_id") clientID := account.GetCredential("client_id")
return s.RefreshTokenWithClientID(ctx, refreshToken, proxyURL, clientID) return s.RefreshTokenWithClientID(ctx, refreshToken, proxyURL, clientID)
} }

View File

@ -8,6 +8,7 @@ import (
"time" "time"
"github.com/Wei-Shaw/sub2api/internal/pkg/openai" "github.com/Wei-Shaw/sub2api/internal/pkg/openai"
"github.com/imroc/req/v3"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
@ -32,6 +33,11 @@ func (s *openaiOAuthClientRefreshStub) RefreshTokenWithClientID(ctx context.Cont
func TestOpenAIOAuthService_RefreshAccountToken_NoRefreshTokenUsesExistingAccessToken(t *testing.T) { func TestOpenAIOAuthService_RefreshAccountToken_NoRefreshTokenUsesExistingAccessToken(t *testing.T) {
client := &openaiOAuthClientRefreshStub{} client := &openaiOAuthClientRefreshStub{}
svc := NewOpenAIOAuthService(nil, client) svc := NewOpenAIOAuthService(nil, client)
var privacyClientCalls int32
svc.SetPrivacyClientFactory(func(proxyURL string) (*req.Client, error) {
atomic.AddInt32(&privacyClientCalls, 1)
return nil, errors.New("stop before request")
})
expiresAt := time.Now().Add(30 * time.Minute).UTC().Format(time.RFC3339) expiresAt := time.Now().Add(30 * time.Minute).UTC().Format(time.RFC3339)
account := &Account{ account := &Account{
@ -51,6 +57,7 @@ func TestOpenAIOAuthService_RefreshAccountToken_NoRefreshTokenUsesExistingAccess
require.Equal(t, "existing-access-token", info.AccessToken) require.Equal(t, "existing-access-token", info.AccessToken)
require.Equal(t, "client-id-1", info.ClientID) require.Equal(t, "client-id-1", info.ClientID)
require.Zero(t, atomic.LoadInt32(&client.refreshCalls), "existing access token should be reused without calling refresh") require.Zero(t, atomic.LoadInt32(&client.refreshCalls), "existing access token should be reused without calling refresh")
require.Positive(t, atomic.LoadInt32(&privacyClientCalls), "existing access token should still run enrichment")
} }
func TestOpenAITokenRefresher_NeedsRefresh_SkipsAccountWithoutRefreshToken(t *testing.T) { func TestOpenAITokenRefresher_NeedsRefresh_SkipsAccountWithoutRefreshToken(t *testing.T) {

View File

@ -95,6 +95,8 @@ type ChatGPTAccountInfo struct {
const chatGPTAccountsCheckURL = "https://chatgpt.com/backend-api/accounts/check/v4-2023-04-27" const chatGPTAccountsCheckURL = "https://chatgpt.com/backend-api/accounts/check/v4-2023-04-27"
var chatGPTSubscriptionsURL = "https://chatgpt.com/backend-api/subscriptions"
// fetchChatGPTAccountInfo calls ChatGPT backend-api to get account info (plan_type, etc.). // fetchChatGPTAccountInfo calls ChatGPT backend-api to get account info (plan_type, etc.).
// Used as fallback when id_token doesn't contain these fields (e.g., Mobile RT). // Used as fallback when id_token doesn't contain these fields (e.g., Mobile RT).
// orgID is used to match the correct account when multiple accounts exist (e.g., personal + team). // orgID is used to match the correct account when multiple accounts exist (e.g., personal + team).
@ -199,6 +201,62 @@ func fetchChatGPTAccountInfo(ctx context.Context, clientFactory PrivacyClientFac
return info return info
} }
// fetchChatGPTSubscriptionExpiresAt reads the lightweight subscription endpoint used by
// ChatGPT/Codex clients. Some Plus accounts no longer expose entitlement.expires_at in
// accounts/check, but this endpoint still returns active_until.
func fetchChatGPTSubscriptionExpiresAt(ctx context.Context, clientFactory PrivacyClientFactory, accessToken, proxyURL, accountID string) string {
accountID = strings.TrimSpace(accountID)
if accessToken == "" || accountID == "" || clientFactory == nil {
return ""
}
ctx, cancel := context.WithTimeout(ctx, 15*time.Second)
defer cancel()
client, err := clientFactory(proxyURL)
if err != nil {
slog.Debug("chatgpt_subscription_client_error", "error", err.Error())
return ""
}
var result struct {
PlanType string `json:"plan_type"`
ActiveUntil string `json:"active_until"`
WillRenew bool `json:"will_renew"`
ID string `json:"id"`
}
resp, err := client.R().
SetContext(ctx).
SetHeader("Authorization", "Bearer "+accessToken).
SetHeader("Origin", "https://chatgpt.com").
SetHeader("Referer", "https://chatgpt.com/").
SetHeader("Accept", "application/json").
SetSuccessResult(&result).
SetQueryParam("account_id", accountID).
Get(chatGPTSubscriptionsURL)
if err != nil {
slog.Debug("chatgpt_subscription_request_error", "error", err.Error())
return ""
}
if !resp.IsSuccessState() {
slog.Debug("chatgpt_subscription_failed", "status", resp.StatusCode, "body", truncate(resp.String(), 200))
return ""
}
activeUntil := strings.TrimSpace(result.ActiveUntil)
if activeUntil == "" {
slog.Debug("chatgpt_subscription_no_active_until", "plan_type", result.PlanType, "has_subscription_id", strings.TrimSpace(result.ID) != "", "will_renew", result.WillRenew)
return ""
}
if _, err := time.Parse(time.RFC3339, activeUntil); err != nil {
slog.Debug("chatgpt_subscription_bad_active_until", "active_until", activeUntil, "error", err.Error())
return ""
}
slog.Info("chatgpt_subscription_success", "plan_type", result.PlanType, "subscription_expires_at", activeUntil, "account_id", accountID)
return activeUntil
}
// fillAccountInfo 从单个 account 对象中提取 plan_type 和 subscription_expires_at // fillAccountInfo 从单个 account 对象中提取 plan_type 和 subscription_expires_at
func fillAccountInfo(info *ChatGPTAccountInfo, acct map[string]any) { func fillAccountInfo(info *ChatGPTAccountInfo, acct map[string]any) {
info.PlanType = extractPlanType(acct) info.PlanType = extractPlanType(acct)

View File

@ -0,0 +1,42 @@
package service
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/imroc/req/v3"
"github.com/stretchr/testify/require"
)
func TestFetchChatGPTSubscriptionExpiresAt(t *testing.T) {
const wantExpiresAt = "2026-06-10T02:52:15Z"
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
require.Equal(t, "/backend-api/subscriptions", r.URL.Path)
require.Equal(t, "acc_123", r.URL.Query().Get("account_id"))
require.Equal(t, "Bearer access-token", r.Header.Get("Authorization"))
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(map[string]any{
"plan_type": "plus",
"active_until": wantExpiresAt,
"will_renew": true,
"id": "sub_123",
})
}))
defer server.Close()
oldURL := chatGPTSubscriptionsURL
chatGPTSubscriptionsURL = server.URL + "/backend-api/subscriptions"
t.Cleanup(func() { chatGPTSubscriptionsURL = oldURL })
got := fetchChatGPTSubscriptionExpiresAt(context.Background(), func(proxyURL string) (*req.Client, error) {
return req.C().SetTimeout(5 * time.Second), nil
}, "access-token", "", "acc_123")
require.Equal(t, wantExpiresAt, got)
}

View File

@ -45,6 +45,17 @@ func ProvideOAuthRefreshAPI(accountRepo AccountRepository, tokenCache GeminiToke
return NewOAuthRefreshAPI(accountRepo, tokenCache) return NewOAuthRefreshAPI(accountRepo, tokenCache)
} }
// ProvideOpenAIOAuthService creates OpenAIOAuthService with privacy/account enrichment support.
func ProvideOpenAIOAuthService(
proxyRepo ProxyRepository,
oauthClient OpenAIOAuthClient,
privacyClientFactory PrivacyClientFactory,
) *OpenAIOAuthService {
svc := NewOpenAIOAuthService(proxyRepo, oauthClient)
svc.SetPrivacyClientFactory(privacyClientFactory)
return svc
}
// ProvideTokenRefreshService creates and starts TokenRefreshService // ProvideTokenRefreshService creates and starts TokenRefreshService
func ProvideTokenRefreshService( func ProvideTokenRefreshService(
accountRepo AccountRepository, accountRepo AccountRepository,
@ -501,7 +512,7 @@ var ProviderSet = wire.NewSet(
NewOpenAIGatewayService, NewOpenAIGatewayService,
wire.Bind(new(AccountRuntimeBlocker), new(*OpenAIGatewayService)), wire.Bind(new(AccountRuntimeBlocker), new(*OpenAIGatewayService)),
NewOAuthService, NewOAuthService,
NewOpenAIOAuthService, ProvideOpenAIOAuthService,
NewGeminiOAuthService, NewGeminiOAuthService,
NewGeminiQuotaService, NewGeminiQuotaService,
NewCompositeTokenCacheInvalidator, NewCompositeTokenCacheInvalidator,