Implement global OAuth service and enhance cache management in Signin API

- Introduced a global OAuth service variable to streamline access across the application.
- Added new methods in the OAuth Service for retrieving cache and store, improving data management.
- Refactored session state handling in the Signin API to utilize the global OAuth service for better cache interactions.
- Updated key management functions for user info and redirect URIs to enhance clarity and maintainability.
This commit is contained in:
Max 2025-08-02 08:34:18 +08:00
parent 2f0a31894c
commit 7111d9fab2
2 changed files with 54 additions and 47 deletions

View file

@ -11,6 +11,9 @@ import (
"github.com/yaoapp/yao/share"
)
// OAuth is the global OAuth service
var OAuth *Service = nil
// Service OAuth service
type Service struct {
config *Config
@ -147,6 +150,8 @@ func NewService(config *Config) (*Service, error) {
signingCerts: signingCerts,
}
// Set the global OAuth service
OAuth = service
return service, nil
}
@ -165,6 +170,16 @@ func (s *Service) GetClientProvider() types.ClientProvider {
return s.clientProvider
}
// GetCache returns the cache for the service
func (s *Service) GetCache() store.Store {
return s.cache
}
// GetStore returns the store for the service
func (s *Service) GetStore() store.Store {
return s.store
}
// setConfigDefaults sets default values for configuration
func setConfigDefaults(config *Config) error {
// Certificate defaults

View file

@ -11,8 +11,8 @@ import (
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"github.com/yaoapp/gou/session"
"github.com/yaoapp/gou/store"
"github.com/yaoapp/kun/log"
"github.com/yaoapp/yao/openapi/oauth"
"github.com/yaoapp/yao/openapi/oauth/types"
"github.com/yaoapp/yao/openapi/response"
"github.com/yaoapp/yao/openapi/utils"
@ -193,6 +193,8 @@ func authback(c *gin.Context) {
userInfo, err = provider.GetUserInfo(tokenResponse.AccessToken, tokenResponse.TokenType)
}
// Create / Update / User then login (Generate access_token and id_token)
if err != nil {
errorResp := &response.ErrorResponse{
Code: response.ErrInvalidRequest.Code,
@ -293,8 +295,6 @@ func getOAuthAuthorizationURL(c *gin.Context) {
// Add scopes
if len(provider.Scopes) > 0 {
params.Add("scope", strings.Join(provider.Scopes, " "))
} else {
params.Add("scope", "openid profile email")
}
// Add response_mode if specified (required for Apple with name/email scopes)
@ -398,30 +398,37 @@ func generateSessionID() string {
return session.ID()
}
// userInfoKey returns the key for the user info
func userInfoKey(providerID, state string) string {
return fmt.Sprintf("signin:user_info:%s:%s", providerID, state)
}
// stateKey returns the key for the state
func stateKey(providerID string) string {
return fmt.Sprintf("signin:state:%s", providerID)
}
// redirectURIKey returns the key for the redirect URI
func redirectURIKey(providerID, state string) string {
return fmt.Sprintf("signin:redirect_uri:%s:%s", providerID, state)
}
// saveState saves the state to the session
func saveState(providerID, sid, state string) error {
return session.Global().ID(sid).SetWithEx(fmt.Sprintf("oauth_state_%s", providerID), state, 20*time.Minute)
return session.Global().ID(sid).SetWithEx(stateKey(providerID), state, 20*time.Minute)
}
// saveRedirectURI saves the redirect URI to the session
func saveRedirectURI(providerID, state, redirectURI string) error {
key := fmt.Sprintf("oauth_redirect_uri_%s_%s", providerID, state)
store, err := store.Get("__yao.oauth.cache")
if err != nil {
return err
}
store.Set(key, redirectURI, 20*time.Minute)
return nil
key := redirectURIKey(providerID, state)
store := oauth.OAuth.GetCache()
return store.Set(key, redirectURI, 20*time.Minute)
}
// getRedirectURI gets the redirect URI from the session
func getRedirectURI(providerID, state string) (string, error) {
key := fmt.Sprintf("oauth_redirect_uri_%s_%s", providerID, state)
store, err := store.Get("__yao.oauth.cache")
if err != nil {
return "", err
}
key := redirectURIKey(providerID, state)
store := oauth.OAuth.GetCache()
value, ok := store.Get(key)
if !ok || value == nil {
return "", fmt.Errorf("redirect URI not found")
@ -430,34 +437,22 @@ func getRedirectURI(providerID, state string) (string, error) {
}
func removeRedirectURI(providerID, state string) error {
key := fmt.Sprintf("oauth_redirect_uri_%s_%s", providerID, state)
store, err := store.Get("__yao.oauth.cache")
if err != nil {
return err
}
store.Del(key)
return nil
key := redirectURIKey(providerID, state)
store := oauth.OAuth.GetCache()
return store.Del(key)
}
// saveUserInfo saves the user info to cache (for form_post mode)
func saveUserInfo(providerID, state, userInfo string) error {
key := fmt.Sprintf("oauth_user_info_%s_%s", providerID, state)
store, err := store.Get("__yao.oauth.cache")
if err != nil {
return err
}
store.Set(key, userInfo, 20*time.Minute)
return nil
key := userInfoKey(providerID, state)
store := oauth.OAuth.GetCache()
return store.Set(key, userInfo, 20*time.Minute)
}
// getUserInfo gets the user info from cache
func getUserInfo(providerID, state string) (string, error) {
key := fmt.Sprintf("oauth_user_info_%s_%s", providerID, state)
store, err := store.Get("__yao.oauth.cache")
if err != nil {
return "", err
}
key := userInfoKey(providerID, state)
store := oauth.OAuth.GetCache()
value, ok := store.Get(key)
if !ok || value == nil {
return "", fmt.Errorf("user info not found")
@ -467,18 +462,15 @@ func getUserInfo(providerID, state string) (string, error) {
// removeUserInfo removes the user info from cache
func removeUserInfo(providerID, state string) error {
key := fmt.Sprintf("oauth_user_info_%s_%s", providerID, state)
store, err := store.Get("__yao.oauth.cache")
if err != nil {
return err
}
store.Del(key)
return nil
key := userInfoKey(providerID, state)
store := oauth.OAuth.GetCache()
return store.Del(key)
}
// removeState removes the state from the session
func removeState(providerID, sid string) error {
// Get the state from the session
state, err := session.Global().ID(sid).Get(fmt.Sprintf("oauth_state_%s", providerID))
state, err := session.Global().ID(sid).Get(stateKey(providerID))
if err != nil {
return err
}
@ -493,12 +485,12 @@ func removeState(providerID, sid string) error {
removeRedirectURI(providerID, stateStr)
removeUserInfo(providerID, stateStr)
return session.Global().ID(sid).Del(fmt.Sprintf("oauth_state_%s", providerID))
return session.Global().ID(sid).Del(stateKey(providerID))
}
// validateState validates the state from the session
func validateState(providerID, sid, state string) error {
value, err := session.Global().ID(sid).Get(fmt.Sprintf("oauth_state_%s", providerID))
value, err := session.Global().ID(sid).Get(stateKey(providerID))
if err != nil {
return err
}