Merge pull request #1007 from trheyi/main
Remove deprecated OAuth interfaces and types
This commit is contained in:
commit
7636455d08
15 changed files with 1550 additions and 108 deletions
34
openapi/oauth/client.go
Normal file
34
openapi/oauth/client.go
Normal file
|
|
@ -0,0 +1,34 @@
|
||||||
|
package oauth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"github.com/yaoapp/yao/openapi/oauth/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Register registers a new OAuth client with the authorization server
|
||||||
|
func (s *Service) Register(ctx context.Context, clientInfo *types.ClientInfo) (*types.ClientInfo, error) {
|
||||||
|
return s.clientProvider.CreateClient(ctx, clientInfo)
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdateClient updates an existing OAuth client configuration
|
||||||
|
func (s *Service) UpdateClient(ctx context.Context, clientID string, clientInfo *types.ClientInfo) (*types.ClientInfo, error) {
|
||||||
|
return s.clientProvider.UpdateClient(ctx, clientID, clientInfo)
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteClient removes an OAuth client from the authorization server
|
||||||
|
func (s *Service) DeleteClient(ctx context.Context, clientID string) error {
|
||||||
|
return s.clientProvider.DeleteClient(ctx, clientID)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateScope validates requested scopes against available scopes
|
||||||
|
func (s *Service) ValidateScope(ctx context.Context, requestedScopes []string, clientID string) (*types.ValidationResult, error) {
|
||||||
|
return s.clientProvider.ValidateScope(ctx, clientID, requestedScopes)
|
||||||
|
}
|
||||||
|
|
||||||
|
// DynamicClientRegistration handles dynamic client registration
|
||||||
|
// This implements RFC 7591 for automatic client registration
|
||||||
|
func (s *Service) DynamicClientRegistration(ctx context.Context, request *types.DynamicClientRegistrationRequest) (*types.DynamicClientRegistrationResponse, error) {
|
||||||
|
// TODO: Implement dynamic client registration
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
52
openapi/oauth/core.go
Normal file
52
openapi/oauth/core.go
Normal file
|
|
@ -0,0 +1,52 @@
|
||||||
|
package oauth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"github.com/yaoapp/yao/openapi/oauth/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
// AuthorizationServer returns the authorization server endpoint URL
|
||||||
|
func (s *Service) AuthorizationServer(ctx context.Context) string {
|
||||||
|
return s.config.IssuerURL
|
||||||
|
}
|
||||||
|
|
||||||
|
// ProtectedResource returns the protected resource endpoint URL
|
||||||
|
func (s *Service) ProtectedResource(ctx context.Context) string {
|
||||||
|
return s.config.IssuerURL
|
||||||
|
}
|
||||||
|
|
||||||
|
// Authorize processes an authorization request and returns an authorization code
|
||||||
|
// The authorization code can be exchanged for an access token
|
||||||
|
func (s *Service) Authorize(ctx context.Context, request *types.AuthorizationRequest) (*types.AuthorizationResponse, error) {
|
||||||
|
// TODO: Implement authorization flow
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Token exchanges an authorization code for an access token
|
||||||
|
// This is the core token endpoint functionality
|
||||||
|
func (s *Service) Token(ctx context.Context, grantType string, code string, clientID string, codeVerifier string) (*types.Token, error) {
|
||||||
|
// TODO: Implement token exchange
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Revoke revokes an access token or refresh token
|
||||||
|
// Once revoked, the token cannot be used for accessing protected resources
|
||||||
|
func (s *Service) Revoke(ctx context.Context, token string, tokenTypeHint string) error {
|
||||||
|
// TODO: Implement token revocation
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// RefreshToken exchanges a refresh token for a new access token
|
||||||
|
// This allows clients to obtain fresh access tokens without user interaction
|
||||||
|
func (s *Service) RefreshToken(ctx context.Context, refreshToken string, scope string) (*types.RefreshTokenResponse, error) {
|
||||||
|
// TODO: Implement refresh token exchange
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// RotateRefreshToken rotates a refresh token and invalidates the old one
|
||||||
|
// This implements refresh token rotation for enhanced security
|
||||||
|
func (s *Service) RotateRefreshToken(ctx context.Context, oldToken string) (*types.RefreshTokenResponse, error) {
|
||||||
|
// TODO: Implement refresh token rotation
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
14
openapi/oauth/device.go
Normal file
14
openapi/oauth/device.go
Normal file
|
|
@ -0,0 +1,14 @@
|
||||||
|
package oauth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"github.com/yaoapp/yao/openapi/oauth/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
// DeviceAuthorization initiates the device authorization flow
|
||||||
|
// This is used for devices with limited input capabilities
|
||||||
|
func (s *Service) DeviceAuthorization(ctx context.Context, clientID string, scope string) (*types.DeviceAuthorizationResponse, error) {
|
||||||
|
// TODO: Implement device authorization flow
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
28
openapi/oauth/discovery.go
Normal file
28
openapi/oauth/discovery.go
Normal file
|
|
@ -0,0 +1,28 @@
|
||||||
|
package oauth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"github.com/yaoapp/yao/openapi/oauth/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
// JWKS returns the JSON Web Key Set for token verification
|
||||||
|
// This endpoint provides public keys for validating JWT tokens
|
||||||
|
func (s *Service) JWKS(ctx context.Context) (*types.JWKSResponse, error) {
|
||||||
|
// TODO: Implement JWKS endpoint
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Endpoints returns a map of all available OAuth endpoints
|
||||||
|
// This provides endpoint discovery for clients
|
||||||
|
func (s *Service) Endpoints(ctx context.Context) (map[string]string, error) {
|
||||||
|
// TODO: Implement endpoint discovery
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetServerMetadata returns OAuth 2.0 Authorization Server Metadata
|
||||||
|
// This implements RFC 8414 for server discovery
|
||||||
|
func (s *Service) GetServerMetadata(ctx context.Context) (*types.AuthorizationServerMetadata, error) {
|
||||||
|
// TODO: Implement server metadata
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
35
openapi/oauth/mcp.go
Normal file
35
openapi/oauth/mcp.go
Normal file
|
|
@ -0,0 +1,35 @@
|
||||||
|
package oauth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"github.com/yaoapp/yao/openapi/oauth/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ValidateResourceParameter validates an OAuth 2.0 resource parameter
|
||||||
|
// This ensures the resource parameter is valid and properly formatted
|
||||||
|
func (s *Service) ValidateResourceParameter(ctx context.Context, resource string) (*types.ValidationResult, error) {
|
||||||
|
// TODO: Implement resource parameter validation
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetCanonicalResourceURI returns the canonical form of a resource URI
|
||||||
|
// This normalizes resource URIs for consistent processing
|
||||||
|
func (s *Service) GetCanonicalResourceURI(ctx context.Context, serverURI string) (string, error) {
|
||||||
|
// TODO: Implement canonical resource URI generation
|
||||||
|
return "", nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetProtectedResourceMetadata returns OAuth 2.0 Protected Resource Metadata
|
||||||
|
// This implements RFC 9728 for MCP server discovery
|
||||||
|
func (s *Service) GetProtectedResourceMetadata(ctx context.Context) (*types.ProtectedResourceMetadata, error) {
|
||||||
|
// TODO: Implement protected resource metadata
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// HandleWWWAuthenticate processes WWW-Authenticate challenges
|
||||||
|
// This handles authentication challenges from protected resources
|
||||||
|
func (s *Service) HandleWWWAuthenticate(ctx context.Context, challenge string) (*types.WWWAuthenticateChallenge, error) {
|
||||||
|
// TODO: Implement WWW-Authenticate challenge handling
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
@ -1,17 +1,21 @@
|
||||||
package oauth
|
package oauth
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/yaoapp/gou/store"
|
"github.com/yaoapp/gou/store"
|
||||||
|
"github.com/yaoapp/yao/openapi/oauth/providers/client"
|
||||||
|
"github.com/yaoapp/yao/openapi/oauth/providers/user"
|
||||||
|
"github.com/yaoapp/yao/openapi/oauth/types"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Service OAuth service
|
// Service OAuth service
|
||||||
type Service struct {
|
type Service struct {
|
||||||
config *Config
|
config *Config
|
||||||
store store.Store
|
store store.Store
|
||||||
userProvider UserProvider
|
cache store.Store
|
||||||
|
userProvider types.UserProvider
|
||||||
|
clientProvider types.ClientProvider
|
||||||
}
|
}
|
||||||
|
|
||||||
// Config OAuth service configuration
|
// Config OAuth service configuration
|
||||||
|
|
@ -19,20 +23,26 @@ type Config struct {
|
||||||
// Core storage interface
|
// Core storage interface
|
||||||
Store store.Store `json:"-"`
|
Store store.Store `json:"-"`
|
||||||
|
|
||||||
|
// Cache store
|
||||||
|
Cache store.Store `json:"-"`
|
||||||
|
|
||||||
// User provider interface
|
// User provider interface
|
||||||
UserProvider UserProvider `json:"-"`
|
UserProvider types.UserProvider `json:"-"`
|
||||||
|
|
||||||
|
// Client provider interface
|
||||||
|
ClientProvider types.ClientProvider `json:"-"`
|
||||||
|
|
||||||
// Certificate and key management
|
// Certificate and key management
|
||||||
Signing SigningConfig `json:"signing"`
|
Signing types.SigningConfig `json:"signing"`
|
||||||
|
|
||||||
// Token management settings
|
// Token management settings
|
||||||
Token TokenConfig `json:"token"`
|
Token types.TokenConfig `json:"token"`
|
||||||
|
|
||||||
// Security configuration
|
// Security configuration
|
||||||
Security SecurityConfig `json:"security"`
|
Security types.SecurityConfig `json:"security"`
|
||||||
|
|
||||||
// Default client settings
|
// Default client settings
|
||||||
Client ClientConfig `json:"client"`
|
Client types.ClientConfig `json:"client"`
|
||||||
|
|
||||||
// Feature flags
|
// Feature flags
|
||||||
Features FeatureFlags `json:"features"`
|
Features FeatureFlags `json:"features"`
|
||||||
|
|
@ -72,7 +82,7 @@ type FeatureFlags struct {
|
||||||
// NewService creates a new OAuth service with the given configuration
|
// NewService creates a new OAuth service with the given configuration
|
||||||
func NewService(config *Config) (*Service, error) {
|
func NewService(config *Config) (*Service, error) {
|
||||||
if config == nil {
|
if config == nil {
|
||||||
return nil, ErrInvalidConfiguration
|
return nil, types.ErrInvalidConfiguration
|
||||||
}
|
}
|
||||||
|
|
||||||
// Set default values if not provided
|
// Set default values if not provided
|
||||||
|
|
@ -88,13 +98,29 @@ func NewService(config *Config) (*Service, error) {
|
||||||
// Use UserProvider from config, or create a default one if not provided
|
// Use UserProvider from config, or create a default one if not provided
|
||||||
userProvider := config.UserProvider
|
userProvider := config.UserProvider
|
||||||
if userProvider == nil {
|
if userProvider == nil {
|
||||||
userProvider = NewDefaultUserProvider(nil, nil, nil)
|
userProvider = user.NewDefaultUserProvider(nil, nil, nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Use ClientProvider from config, or create a default one if not provided
|
||||||
|
clientProvider := config.ClientProvider
|
||||||
|
if clientProvider == nil {
|
||||||
|
var err error
|
||||||
|
clientProvider, err = client.NewDefaultClient(&client.DefaultClientOptions{
|
||||||
|
Prefix: "__yao:",
|
||||||
|
Store: config.Store,
|
||||||
|
Cache: config.Cache,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
service := &Service{
|
service := &Service{
|
||||||
config: config,
|
config: config,
|
||||||
store: config.Store,
|
store: config.Store,
|
||||||
|
cache: config.Cache,
|
||||||
userProvider: userProvider,
|
userProvider: userProvider,
|
||||||
|
clientProvider: clientProvider,
|
||||||
}
|
}
|
||||||
|
|
||||||
return service, nil
|
return service, nil
|
||||||
|
|
@ -106,10 +132,15 @@ func (s *Service) GetConfig() *Config {
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetUserProvider returns the user provider for the service
|
// GetUserProvider returns the user provider for the service
|
||||||
func (s *Service) GetUserProvider() UserProvider {
|
func (s *Service) GetUserProvider() types.UserProvider {
|
||||||
return s.userProvider
|
return s.userProvider
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GetClientProvider returns the client provider for the service
|
||||||
|
func (s *Service) GetClientProvider() types.ClientProvider {
|
||||||
|
return s.clientProvider
|
||||||
|
}
|
||||||
|
|
||||||
// setConfigDefaults sets default values for configuration
|
// setConfigDefaults sets default values for configuration
|
||||||
func setConfigDefaults(config *Config) error {
|
func setConfigDefaults(config *Config) error {
|
||||||
// Certificate defaults
|
// Certificate defaults
|
||||||
|
|
@ -182,53 +213,28 @@ func setConfigDefaults(config *Config) error {
|
||||||
// validateConfig validates the configuration
|
// validateConfig validates the configuration
|
||||||
func validateConfig(config *Config) error {
|
func validateConfig(config *Config) error {
|
||||||
if config.Store == nil {
|
if config.Store == nil {
|
||||||
return ErrStoreMissing
|
return types.ErrStoreMissing
|
||||||
}
|
}
|
||||||
|
|
||||||
// Validate issuer URL
|
// Validate issuer URL
|
||||||
if config.IssuerURL == "" {
|
if config.IssuerURL == "" {
|
||||||
return ErrIssuerURLMissing
|
return types.ErrIssuerURLMissing
|
||||||
}
|
}
|
||||||
|
|
||||||
// Validate certificate configuration
|
// Validate certificate configuration
|
||||||
if config.Signing.SigningCertPath == "" || config.Signing.SigningKeyPath == "" {
|
if config.Signing.SigningCertPath == "" || config.Signing.SigningKeyPath == "" {
|
||||||
return ErrCertificateMissing
|
return types.ErrCertificateMissing
|
||||||
}
|
}
|
||||||
|
|
||||||
// Validate token configuration
|
// Validate token configuration
|
||||||
if config.Token.AccessTokenLifetime <= 0 {
|
if config.Token.AccessTokenLifetime <= 0 {
|
||||||
return ErrInvalidTokenLifetime
|
return types.ErrInvalidTokenLifetime
|
||||||
}
|
}
|
||||||
|
|
||||||
// Validate security configuration
|
// Validate security configuration
|
||||||
if config.Security.PKCERequired && len(config.Security.PKCECodeChallengeMethod) == 0 {
|
if config.Security.PKCERequired && len(config.Security.PKCECodeChallengeMethod) == 0 {
|
||||||
return ErrPKCEConfigurationInvalid
|
return types.ErrPKCEConfigurationInvalid
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Error definitions
|
|
||||||
var (
|
|
||||||
ErrInvalidConfiguration = &ErrorResponse{Code: "invalid_configuration", ErrorDescription: "Invalid OAuth service configuration"}
|
|
||||||
ErrStoreMissing = &ErrorResponse{Code: "store_missing", ErrorDescription: "Store is required for OAuth service"}
|
|
||||||
ErrIssuerURLMissing = &ErrorResponse{Code: "issuer_url_missing", ErrorDescription: "Issuer URL is required for OAuth service"}
|
|
||||||
ErrCertificateMissing = &ErrorResponse{Code: "certificate_missing", ErrorDescription: "JWT signing certificate and key are required"}
|
|
||||||
ErrInvalidTokenLifetime = &ErrorResponse{Code: "invalid_token_lifetime", ErrorDescription: "Token lifetime must be greater than 0"}
|
|
||||||
ErrPKCEConfigurationInvalid = &ErrorResponse{Code: "pkce_configuration_invalid", ErrorDescription: "PKCE configuration is invalid"}
|
|
||||||
)
|
|
||||||
|
|
||||||
// AuthorizationServer returns the authorization server endpoint URL
|
|
||||||
func (s *Service) AuthorizationServer(ctx context.Context) string {
|
|
||||||
return s.config.IssuerURL
|
|
||||||
}
|
|
||||||
|
|
||||||
// ProtectedResource returns the protected resource endpoint URL
|
|
||||||
func (s *Service) ProtectedResource(ctx context.Context) string {
|
|
||||||
return s.config.IssuerURL
|
|
||||||
}
|
|
||||||
|
|
||||||
// UserInfo returns user information for a given access token
|
|
||||||
func (s *Service) UserInfo(ctx context.Context, accessToken string) (interface{}, error) {
|
|
||||||
return s.userProvider.GetUserByAccessToken(ctx, accessToken)
|
|
||||||
}
|
|
||||||
|
|
|
||||||
602
openapi/oauth/providers/client/default.go
Normal file
602
openapi/oauth/providers/client/default.go
Normal file
|
|
@ -0,0 +1,602 @@
|
||||||
|
package client
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/yaoapp/gou/store"
|
||||||
|
"github.com/yaoapp/yao/openapi/oauth/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
// DefaultClient provides a default implementation of ClientProvider
|
||||||
|
type DefaultClient struct {
|
||||||
|
prefix string
|
||||||
|
cache store.Store
|
||||||
|
store store.Store
|
||||||
|
}
|
||||||
|
|
||||||
|
// DefaultClientOptions provides options for the DefaultClient
|
||||||
|
type DefaultClientOptions struct {
|
||||||
|
Prefix string
|
||||||
|
Cache store.Store
|
||||||
|
Store store.Store
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewDefaultClient creates a new DefaultClient
|
||||||
|
func NewDefaultClient(options *DefaultClientOptions) (*DefaultClient, error) {
|
||||||
|
if options == nil {
|
||||||
|
return nil, types.ErrInvalidConfiguration
|
||||||
|
}
|
||||||
|
|
||||||
|
if options.Store == nil {
|
||||||
|
return nil, types.ErrStoreMissing
|
||||||
|
}
|
||||||
|
|
||||||
|
if options.Prefix == "" {
|
||||||
|
options.Prefix = "__yao:"
|
||||||
|
}
|
||||||
|
|
||||||
|
return &DefaultClient{
|
||||||
|
prefix: options.Prefix,
|
||||||
|
cache: options.Cache,
|
||||||
|
store: options.Store,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Key generation methods
|
||||||
|
|
||||||
|
func (c *DefaultClient) clientKey(clientID string) string {
|
||||||
|
return fmt.Sprintf("%soauth:client:%s", c.prefix, clientID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *DefaultClient) clientListKey() string {
|
||||||
|
return fmt.Sprintf("%soauth:clients", c.prefix)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetClientByID retrieves client information using a client ID
|
||||||
|
func (c *DefaultClient) GetClientByID(ctx context.Context, clientID string) (*types.ClientInfo, error) {
|
||||||
|
// Try cache first if available
|
||||||
|
if c.cache != nil {
|
||||||
|
if cached, ok := c.cache.Get(c.clientKey(clientID)); ok {
|
||||||
|
if clientInfo, ok := cached.(*types.ClientInfo); ok {
|
||||||
|
return clientInfo, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fallback to store
|
||||||
|
key := c.clientKey(clientID)
|
||||||
|
data, ok := c.store.Get(key)
|
||||||
|
if !ok {
|
||||||
|
return nil, &types.ErrorResponse{
|
||||||
|
Code: types.ErrorInvalidClient,
|
||||||
|
ErrorDescription: "Client not found",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var clientInfo *types.ClientInfo
|
||||||
|
|
||||||
|
// Handle different data types returned by different stores
|
||||||
|
switch v := data.(type) {
|
||||||
|
case *types.ClientInfo:
|
||||||
|
// Direct object (from cache)
|
||||||
|
clientInfo = v
|
||||||
|
case map[string]interface{}:
|
||||||
|
// Map with JSON field names (standard format)
|
||||||
|
jsonData, err := json.Marshal(v)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal map data: %w", err)
|
||||||
|
}
|
||||||
|
clientInfo = &types.ClientInfo{}
|
||||||
|
if err := json.Unmarshal(jsonData, clientInfo); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to unmarshal client data: %w", err)
|
||||||
|
}
|
||||||
|
case []byte:
|
||||||
|
// Byte data (for backward compatibility)
|
||||||
|
clientInfo = &types.ClientInfo{}
|
||||||
|
if err := json.Unmarshal(v, clientInfo); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to unmarshal client data: %w", err)
|
||||||
|
}
|
||||||
|
case string:
|
||||||
|
// String data (for backward compatibility)
|
||||||
|
clientInfo = &types.ClientInfo{}
|
||||||
|
if err := json.Unmarshal([]byte(v), clientInfo); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to unmarshal client data: %w", err)
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
// Try JSON marshaling as fallback for unknown types
|
||||||
|
jsonData, err := json.Marshal(data)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal data to JSON: %w", err)
|
||||||
|
}
|
||||||
|
clientInfo = &types.ClientInfo{}
|
||||||
|
if err := json.Unmarshal(jsonData, clientInfo); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to unmarshal client data: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Cache the result if cache is available
|
||||||
|
if c.cache != nil {
|
||||||
|
c.cache.Set(c.clientKey(clientID), clientInfo, 5*time.Minute) // Cache for 5 minutes
|
||||||
|
}
|
||||||
|
|
||||||
|
return clientInfo, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetClientByCredentials retrieves and validates client using client credentials
|
||||||
|
func (c *DefaultClient) GetClientByCredentials(ctx context.Context, clientID string, clientSecret string) (*types.ClientInfo, error) {
|
||||||
|
clientInfo, err := c.GetClientByID(ctx, clientID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// For public clients, no secret validation required
|
||||||
|
if clientInfo.ClientType == types.ClientTypePublic {
|
||||||
|
return clientInfo, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// For confidential clients, validate secret
|
||||||
|
if clientInfo.ClientSecret != clientSecret {
|
||||||
|
return nil, &types.ErrorResponse{
|
||||||
|
Code: types.ErrorInvalidClient,
|
||||||
|
ErrorDescription: "Invalid client credentials",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return clientInfo, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CreateClient creates a new OAuth client and returns the client information
|
||||||
|
func (c *DefaultClient) CreateClient(ctx context.Context, clientInfo *types.ClientInfo) (*types.ClientInfo, error) {
|
||||||
|
// Validate required fields
|
||||||
|
if clientInfo.ClientID == "" {
|
||||||
|
return nil, &types.ErrorResponse{
|
||||||
|
Code: types.ErrorInvalidRequest,
|
||||||
|
ErrorDescription: "Client ID is required",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if client already exists
|
||||||
|
existing, err := c.GetClientByID(ctx, clientInfo.ClientID)
|
||||||
|
if err == nil && existing != nil {
|
||||||
|
return nil, &types.ErrorResponse{
|
||||||
|
Code: types.ErrorInvalidClient,
|
||||||
|
ErrorDescription: "Client already exists",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set timestamps
|
||||||
|
now := time.Now()
|
||||||
|
clientInfo.CreatedAt = now
|
||||||
|
clientInfo.UpdatedAt = now
|
||||||
|
|
||||||
|
// Set defaults
|
||||||
|
if clientInfo.ClientType == "" {
|
||||||
|
clientInfo.ClientType = types.ClientTypeConfidential
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate client
|
||||||
|
validationResult, err := c.ValidateClient(ctx, clientInfo)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if !validationResult.Valid {
|
||||||
|
return nil, &types.ErrorResponse{
|
||||||
|
Code: types.ErrorInvalidRequest,
|
||||||
|
ErrorDescription: strings.Join(validationResult.Errors, "; "),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Save client data
|
||||||
|
if err := c.saveClient(ctx, clientInfo); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add to client list
|
||||||
|
if err := c.addToClientList(ctx, clientInfo.ClientID); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Cache the client if cache is available
|
||||||
|
if c.cache != nil {
|
||||||
|
c.cache.Set(c.clientKey(clientInfo.ClientID), clientInfo, 5*time.Minute)
|
||||||
|
}
|
||||||
|
|
||||||
|
return clientInfo, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdateClient updates an existing OAuth client configuration
|
||||||
|
func (c *DefaultClient) UpdateClient(ctx context.Context, clientID string, clientInfo *types.ClientInfo) (*types.ClientInfo, error) {
|
||||||
|
// Check if client exists
|
||||||
|
existing, err := c.GetClientByID(ctx, clientID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update client ID if provided
|
||||||
|
if clientInfo.ClientID != "" && clientInfo.ClientID != clientID {
|
||||||
|
return nil, &types.ErrorResponse{
|
||||||
|
Code: types.ErrorInvalidRequest,
|
||||||
|
ErrorDescription: "Cannot change client ID",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set client ID and preserve creation time
|
||||||
|
clientInfo.ClientID = clientID
|
||||||
|
clientInfo.CreatedAt = existing.CreatedAt
|
||||||
|
clientInfo.UpdatedAt = time.Now()
|
||||||
|
|
||||||
|
// Validate client
|
||||||
|
validationResult, err := c.ValidateClient(ctx, clientInfo)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if !validationResult.Valid {
|
||||||
|
return nil, &types.ErrorResponse{
|
||||||
|
Code: types.ErrorInvalidRequest,
|
||||||
|
ErrorDescription: strings.Join(validationResult.Errors, "; "),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Save updated client data
|
||||||
|
if err := c.saveClient(ctx, clientInfo); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update cache if available
|
||||||
|
if c.cache != nil {
|
||||||
|
c.cache.Set(c.clientKey(clientID), clientInfo, 5*time.Minute)
|
||||||
|
}
|
||||||
|
|
||||||
|
return clientInfo, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteClient removes an OAuth client from the system
|
||||||
|
func (c *DefaultClient) DeleteClient(ctx context.Context, clientID string) error {
|
||||||
|
// Check if client exists
|
||||||
|
_, err := c.GetClientByID(ctx, clientID)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Remove from client list
|
||||||
|
if err := c.removeFromClientList(ctx, clientID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Delete client data
|
||||||
|
key := c.clientKey(clientID)
|
||||||
|
if err := c.store.Del(key); err != nil {
|
||||||
|
return fmt.Errorf("failed to delete client: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Clear cache if available
|
||||||
|
if c.cache != nil {
|
||||||
|
c.cache.Del(c.clientKey(clientID))
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateClient validates client information and configuration
|
||||||
|
func (c *DefaultClient) ValidateClient(ctx context.Context, clientInfo *types.ClientInfo) (*types.ValidationResult, error) {
|
||||||
|
result := &types.ValidationResult{Valid: true}
|
||||||
|
|
||||||
|
// Validate client ID
|
||||||
|
if clientInfo.ClientID == "" {
|
||||||
|
result.Valid = false
|
||||||
|
result.Errors = append(result.Errors, "Client ID is required")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate client type
|
||||||
|
if clientInfo.ClientType != types.ClientTypeConfidential &&
|
||||||
|
clientInfo.ClientType != types.ClientTypePublic &&
|
||||||
|
clientInfo.ClientType != types.ClientTypeCredentialed {
|
||||||
|
result.Valid = false
|
||||||
|
result.Errors = append(result.Errors, "Invalid client type")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate client secret for confidential clients
|
||||||
|
if clientInfo.ClientType == types.ClientTypeConfidential && clientInfo.ClientSecret == "" {
|
||||||
|
result.Valid = false
|
||||||
|
result.Errors = append(result.Errors, "Client secret is required for confidential clients")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate redirect URIs
|
||||||
|
if len(clientInfo.RedirectURIs) == 0 {
|
||||||
|
result.Valid = false
|
||||||
|
result.Errors = append(result.Errors, "At least one redirect URI is required")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate grant types
|
||||||
|
if len(clientInfo.GrantTypes) == 0 {
|
||||||
|
clientInfo.GrantTypes = []string{types.GrantTypeAuthorizationCode}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate response types
|
||||||
|
if len(clientInfo.ResponseTypes) == 0 {
|
||||||
|
clientInfo.ResponseTypes = []string{types.ResponseTypeCode}
|
||||||
|
}
|
||||||
|
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ListClients retrieves a list of clients with optional filtering
|
||||||
|
func (c *DefaultClient) ListClients(ctx context.Context, filters map[string]interface{}, limit int, offset int) ([]*types.ClientInfo, int, error) {
|
||||||
|
// Get client list
|
||||||
|
clientIDs, err := c.getClientList(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return nil, 0, err
|
||||||
|
}
|
||||||
|
|
||||||
|
var clients []*types.ClientInfo
|
||||||
|
|
||||||
|
// Load all clients
|
||||||
|
for _, clientID := range clientIDs {
|
||||||
|
client, err := c.GetClientByID(ctx, clientID)
|
||||||
|
if err != nil {
|
||||||
|
continue // Skip invalid clients
|
||||||
|
}
|
||||||
|
|
||||||
|
// Apply filters
|
||||||
|
if c.matchesFilters(client, filters) {
|
||||||
|
clients = append(clients, client)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
total := len(clients)
|
||||||
|
|
||||||
|
// Apply pagination
|
||||||
|
if offset > 0 {
|
||||||
|
if offset >= len(clients) {
|
||||||
|
return []*types.ClientInfo{}, total, nil
|
||||||
|
}
|
||||||
|
clients = clients[offset:]
|
||||||
|
}
|
||||||
|
|
||||||
|
if limit > 0 && len(clients) > limit {
|
||||||
|
clients = clients[:limit]
|
||||||
|
}
|
||||||
|
|
||||||
|
return clients, total, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateRedirectURI validates if a redirect URI is registered for the client
|
||||||
|
func (c *DefaultClient) ValidateRedirectURI(ctx context.Context, clientID string, redirectURI string) (*types.ValidationResult, error) {
|
||||||
|
client, err := c.GetClientByID(ctx, clientID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
result := &types.ValidationResult{Valid: false}
|
||||||
|
|
||||||
|
for _, uri := range client.RedirectURIs {
|
||||||
|
if uri == redirectURI {
|
||||||
|
result.Valid = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !result.Valid {
|
||||||
|
result.Errors = append(result.Errors, "Redirect URI not registered for this client")
|
||||||
|
}
|
||||||
|
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateScope validates if the client is authorized to request specific scopes
|
||||||
|
func (c *DefaultClient) ValidateScope(ctx context.Context, clientID string, scopes []string) (*types.ValidationResult, error) {
|
||||||
|
client, err := c.GetClientByID(ctx, clientID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
result := &types.ValidationResult{Valid: true}
|
||||||
|
|
||||||
|
// If client has no scope restrictions, allow all scopes
|
||||||
|
if client.Scope == "" {
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse client allowed scopes
|
||||||
|
allowedScopes := strings.Fields(client.Scope)
|
||||||
|
allowedScopeMap := make(map[string]bool)
|
||||||
|
for _, scope := range allowedScopes {
|
||||||
|
allowedScopeMap[scope] = true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check each requested scope
|
||||||
|
for _, scope := range scopes {
|
||||||
|
if !allowedScopeMap[scope] {
|
||||||
|
result.Valid = false
|
||||||
|
result.Errors = append(result.Errors, fmt.Sprintf("Scope '%s' not allowed for this client", scope))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsClientActive checks if a client is active and can be used for authentication
|
||||||
|
func (c *DefaultClient) IsClientActive(ctx context.Context, clientID string) (bool, error) {
|
||||||
|
client, err := c.GetClientByID(ctx, clientID)
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// For now, all existing clients are considered active
|
||||||
|
// This can be extended to check additional status fields
|
||||||
|
return client != nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Helper methods
|
||||||
|
|
||||||
|
func (c *DefaultClient) saveClient(ctx context.Context, clientInfo *types.ClientInfo) error {
|
||||||
|
key := c.clientKey(clientInfo.ClientID)
|
||||||
|
|
||||||
|
// Convert to map[string]interface{} using JSON serialization to ensure consistent field names
|
||||||
|
jsonData, err := json.Marshal(clientInfo)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to marshal client data: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var clientMap map[string]interface{}
|
||||||
|
if err := json.Unmarshal(jsonData, &clientMap); err != nil {
|
||||||
|
return fmt.Errorf("failed to unmarshal to map: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Store the map - this ensures JSON field names are used consistently
|
||||||
|
if err := c.store.Set(key, clientMap, 0); err != nil {
|
||||||
|
return fmt.Errorf("failed to save client: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *DefaultClient) getClientList(ctx context.Context) ([]string, error) {
|
||||||
|
// Try cache first if available
|
||||||
|
if c.cache != nil {
|
||||||
|
if cached, ok := c.cache.Get(c.clientListKey()); ok {
|
||||||
|
if clientIDs, ok := cached.([]string); ok {
|
||||||
|
return clientIDs, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fallback to store
|
||||||
|
key := c.clientListKey()
|
||||||
|
data, ok := c.store.Get(key)
|
||||||
|
if !ok {
|
||||||
|
return []string{}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var clientIDs []string
|
||||||
|
|
||||||
|
// Handle different data types returned by different stores
|
||||||
|
switch v := data.(type) {
|
||||||
|
case []string:
|
||||||
|
// Direct slice (from stores that preserve slice types)
|
||||||
|
clientIDs = v
|
||||||
|
case []interface{}:
|
||||||
|
// Interface slice (from stores that decode to interface slices)
|
||||||
|
for _, item := range v {
|
||||||
|
if str, ok := item.(string); ok {
|
||||||
|
clientIDs = append(clientIDs, str)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case []byte:
|
||||||
|
// Byte data (for backward compatibility)
|
||||||
|
if err := json.Unmarshal(v, &clientIDs); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to unmarshal client list: %w", err)
|
||||||
|
}
|
||||||
|
case string:
|
||||||
|
// String data (for backward compatibility)
|
||||||
|
if err := json.Unmarshal([]byte(v), &clientIDs); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to unmarshal client list: %w", err)
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
// Handle MongoDB primitive types and other BSON types
|
||||||
|
jsonData, err := json.Marshal(data)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal data to JSON: %w", err)
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(jsonData, &clientIDs); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to unmarshal client list: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Cache the result if cache is available
|
||||||
|
if c.cache != nil {
|
||||||
|
c.cache.Set(c.clientListKey(), clientIDs, 5*time.Minute)
|
||||||
|
}
|
||||||
|
|
||||||
|
return clientIDs, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *DefaultClient) saveClientList(ctx context.Context, clientIDs []string) error {
|
||||||
|
key := c.clientListKey()
|
||||||
|
|
||||||
|
// Store the slice directly - this should work consistently across stores
|
||||||
|
if err := c.store.Set(key, clientIDs, 0); err != nil {
|
||||||
|
return fmt.Errorf("failed to save client list: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update cache if available
|
||||||
|
if c.cache != nil {
|
||||||
|
c.cache.Set(c.clientListKey(), clientIDs, 5*time.Minute)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *DefaultClient) addToClientList(ctx context.Context, clientID string) error {
|
||||||
|
clientIDs, err := c.getClientList(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if already exists
|
||||||
|
for _, id := range clientIDs {
|
||||||
|
if id == clientID {
|
||||||
|
return nil // Already exists
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
clientIDs = append(clientIDs, clientID)
|
||||||
|
|
||||||
|
// Clear cache first to ensure consistency
|
||||||
|
if c.cache != nil {
|
||||||
|
c.cache.Del(c.clientListKey())
|
||||||
|
}
|
||||||
|
|
||||||
|
return c.saveClientList(ctx, clientIDs)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *DefaultClient) removeFromClientList(ctx context.Context, clientID string) error {
|
||||||
|
clientIDs, err := c.getClientList(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Remove client ID
|
||||||
|
var newClientIDs []string
|
||||||
|
for _, id := range clientIDs {
|
||||||
|
if id != clientID {
|
||||||
|
newClientIDs = append(newClientIDs, id)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Clear cache first to ensure consistency
|
||||||
|
if c.cache != nil {
|
||||||
|
c.cache.Del(c.clientListKey())
|
||||||
|
}
|
||||||
|
|
||||||
|
return c.saveClientList(ctx, newClientIDs)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *DefaultClient) matchesFilters(client *types.ClientInfo, filters map[string]interface{}) bool {
|
||||||
|
if filters == nil {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
for key, value := range filters {
|
||||||
|
switch key {
|
||||||
|
case "client_type":
|
||||||
|
if client.ClientType != value.(string) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
case "client_name":
|
||||||
|
if client.ClientName != value.(string) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
case "application_type":
|
||||||
|
if client.ApplicationType != value.(string) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return true
|
||||||
|
}
|
||||||
517
openapi/oauth/providers/client/default_test.go
Normal file
517
openapi/oauth/providers/client/default_test.go
Normal file
|
|
@ -0,0 +1,517 @@
|
||||||
|
package client
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"github.com/yaoapp/gou/connector"
|
||||||
|
"github.com/yaoapp/gou/store"
|
||||||
|
"github.com/yaoapp/gou/store/badger"
|
||||||
|
"github.com/yaoapp/gou/store/lru"
|
||||||
|
"github.com/yaoapp/yao/openapi/oauth/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Store configuration for parameterized tests
|
||||||
|
type StoreConfig struct {
|
||||||
|
Name string
|
||||||
|
GetFunc func(*testing.T) store.Store
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test helpers
|
||||||
|
func getMongoStore(t *testing.T) store.Store {
|
||||||
|
// Skip test if MongoDB is not available
|
||||||
|
host := os.Getenv("MONGO_TEST_HOST")
|
||||||
|
if host == "" {
|
||||||
|
t.Skip("MongoDB not available - set MONGO_TEST_HOST environment variable")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create MongoDB store using connector
|
||||||
|
mongoConnector, err := connector.New("mongo", "oauth_test", []byte(`{
|
||||||
|
"name": "OAuth Test MongoDB",
|
||||||
|
"type": "mongo",
|
||||||
|
"options": {
|
||||||
|
"db": "oauth_test",
|
||||||
|
"hosts": [{
|
||||||
|
"host": "`+host+`",
|
||||||
|
"port": "`+os.Getenv("MONGO_TEST_PORT")+`",
|
||||||
|
"user": "`+os.Getenv("MONGO_TEST_USER")+`",
|
||||||
|
"pass": "`+os.Getenv("MONGO_TEST_PASS")+`"
|
||||||
|
}]
|
||||||
|
}
|
||||||
|
}`))
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
mongoStore, err := store.New(mongoConnector, nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
return mongoStore
|
||||||
|
}
|
||||||
|
|
||||||
|
func getBadgerStore(t *testing.T) store.Store {
|
||||||
|
// Create temporary directory for test database
|
||||||
|
tempDir := t.TempDir()
|
||||||
|
dbPath := filepath.Join(tempDir, "test_oauth_badger")
|
||||||
|
|
||||||
|
badgerStore, err := badger.New(dbPath)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Clean up on test completion
|
||||||
|
t.Cleanup(func() {
|
||||||
|
badgerStore.Close()
|
||||||
|
})
|
||||||
|
|
||||||
|
return badgerStore
|
||||||
|
}
|
||||||
|
|
||||||
|
func getLRUCache(t *testing.T) store.Store {
|
||||||
|
cache, err := lru.New(1000)
|
||||||
|
require.NoError(t, err)
|
||||||
|
return cache
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get all available store configurations
|
||||||
|
func getStoreConfigs() []StoreConfig {
|
||||||
|
return []StoreConfig{
|
||||||
|
{Name: "MongoDB", GetFunc: getMongoStore},
|
||||||
|
{Name: "Badger", GetFunc: getBadgerStore},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func createTestClient(clientID string) *types.ClientInfo {
|
||||||
|
return &types.ClientInfo{
|
||||||
|
ClientID: clientID,
|
||||||
|
ClientSecret: "secret-" + clientID,
|
||||||
|
ClientName: "Test Client " + clientID,
|
||||||
|
ClientType: types.ClientTypeConfidential,
|
||||||
|
RedirectURIs: []string{"https://example.com/callback"},
|
||||||
|
GrantTypes: []string{types.GrantTypeAuthorizationCode},
|
||||||
|
ResponseTypes: []string{types.ResponseTypeCode},
|
||||||
|
Scope: "openid profile email",
|
||||||
|
CreatedAt: time.Now(),
|
||||||
|
UpdatedAt: time.Now(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewDefaultClient(t *testing.T) {
|
||||||
|
storeConfigs := getStoreConfigs()
|
||||||
|
|
||||||
|
for _, config := range storeConfigs {
|
||||||
|
t.Run(config.Name, func(t *testing.T) {
|
||||||
|
t.Run("valid options", func(t *testing.T) {
|
||||||
|
store := config.GetFunc(t)
|
||||||
|
cache := getLRUCache(t)
|
||||||
|
|
||||||
|
client, err := NewDefaultClient(&DefaultClientOptions{
|
||||||
|
Prefix: "test:",
|
||||||
|
Store: store,
|
||||||
|
Cache: cache,
|
||||||
|
})
|
||||||
|
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, client)
|
||||||
|
assert.Equal(t, "test:", client.prefix)
|
||||||
|
assert.Equal(t, store, client.store)
|
||||||
|
assert.Equal(t, cache, client.cache)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("nil options", func(t *testing.T) {
|
||||||
|
client, err := NewDefaultClient(nil)
|
||||||
|
|
||||||
|
assert.Error(t, err)
|
||||||
|
assert.Nil(t, client)
|
||||||
|
assert.Equal(t, types.ErrInvalidConfiguration, err)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("nil store", func(t *testing.T) {
|
||||||
|
client, err := NewDefaultClient(&DefaultClientOptions{
|
||||||
|
Prefix: "test:",
|
||||||
|
Store: nil,
|
||||||
|
})
|
||||||
|
|
||||||
|
assert.Error(t, err)
|
||||||
|
assert.Nil(t, client)
|
||||||
|
assert.Equal(t, types.ErrStoreMissing, err)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("empty prefix uses default", func(t *testing.T) {
|
||||||
|
store := config.GetFunc(t)
|
||||||
|
|
||||||
|
client, err := NewDefaultClient(&DefaultClientOptions{
|
||||||
|
Store: store,
|
||||||
|
})
|
||||||
|
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, client)
|
||||||
|
assert.Equal(t, "__yao:", client.prefix)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("without cache", func(t *testing.T) {
|
||||||
|
store := config.GetFunc(t)
|
||||||
|
|
||||||
|
client, err := NewDefaultClient(&DefaultClientOptions{
|
||||||
|
Prefix: "test:",
|
||||||
|
Store: store,
|
||||||
|
})
|
||||||
|
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, client)
|
||||||
|
assert.Nil(t, client.cache)
|
||||||
|
})
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestKeyGeneration(t *testing.T) {
|
||||||
|
storeConfigs := getStoreConfigs()
|
||||||
|
|
||||||
|
for _, config := range storeConfigs {
|
||||||
|
t.Run(config.Name, func(t *testing.T) {
|
||||||
|
store := config.GetFunc(t)
|
||||||
|
client, err := NewDefaultClient(&DefaultClientOptions{
|
||||||
|
Prefix: "test:",
|
||||||
|
Store: store,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
t.Run("client key", func(t *testing.T) {
|
||||||
|
key := client.clientKey("test-client")
|
||||||
|
expected := "test:oauth:client:test-client"
|
||||||
|
assert.Equal(t, expected, key)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("client list key", func(t *testing.T) {
|
||||||
|
key := client.clientListKey()
|
||||||
|
expected := "test:oauth:clients"
|
||||||
|
assert.Equal(t, expected, key)
|
||||||
|
})
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateClient(t *testing.T) {
|
||||||
|
storeConfigs := getStoreConfigs()
|
||||||
|
|
||||||
|
for _, config := range storeConfigs {
|
||||||
|
t.Run(config.Name, func(t *testing.T) {
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
t.Run("create client without cache", func(t *testing.T) {
|
||||||
|
store := config.GetFunc(t)
|
||||||
|
client, err := NewDefaultClient(&DefaultClientOptions{
|
||||||
|
Prefix: "test1:",
|
||||||
|
Store: store,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Clean up first
|
||||||
|
client.store.Clear()
|
||||||
|
|
||||||
|
testClient := createTestClient("test-client-1")
|
||||||
|
|
||||||
|
created, err := client.CreateClient(ctx, testClient)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, created)
|
||||||
|
assert.Equal(t, testClient.ClientID, created.ClientID)
|
||||||
|
assert.Equal(t, testClient.ClientSecret, created.ClientSecret)
|
||||||
|
assert.NotZero(t, created.CreatedAt)
|
||||||
|
assert.NotZero(t, created.UpdatedAt)
|
||||||
|
|
||||||
|
// Verify client is in store
|
||||||
|
retrieved, err := client.GetClientByID(ctx, testClient.ClientID)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, testClient.ClientID, retrieved.ClientID)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("create client with cache", func(t *testing.T) {
|
||||||
|
store := config.GetFunc(t)
|
||||||
|
cache := getLRUCache(t)
|
||||||
|
client, err := NewDefaultClient(&DefaultClientOptions{
|
||||||
|
Prefix: "test2:",
|
||||||
|
Store: store,
|
||||||
|
Cache: cache,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Clean up first
|
||||||
|
client.store.Clear()
|
||||||
|
client.cache.Clear()
|
||||||
|
|
||||||
|
testClient := createTestClient("test-client-2")
|
||||||
|
|
||||||
|
created, err := client.CreateClient(ctx, testClient)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, created)
|
||||||
|
|
||||||
|
// Verify client is cached
|
||||||
|
key := client.clientKey(testClient.ClientID)
|
||||||
|
cached, ok := client.cache.Get(key)
|
||||||
|
assert.True(t, ok)
|
||||||
|
assert.NotNil(t, cached)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("create client with empty ID", func(t *testing.T) {
|
||||||
|
store := config.GetFunc(t)
|
||||||
|
client, err := NewDefaultClient(&DefaultClientOptions{
|
||||||
|
Prefix: "test3:",
|
||||||
|
Store: store,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
testClient := createTestClient("")
|
||||||
|
|
||||||
|
created, err := client.CreateClient(ctx, testClient)
|
||||||
|
assert.Error(t, err)
|
||||||
|
assert.Nil(t, created)
|
||||||
|
assert.Contains(t, err.Error(), "Client ID is required")
|
||||||
|
})
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetClientByID(t *testing.T) {
|
||||||
|
storeConfigs := getStoreConfigs()
|
||||||
|
|
||||||
|
for _, config := range storeConfigs {
|
||||||
|
t.Run(config.Name, func(t *testing.T) {
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
t.Run("get client without cache", func(t *testing.T) {
|
||||||
|
store := config.GetFunc(t)
|
||||||
|
client, err := NewDefaultClient(&DefaultClientOptions{
|
||||||
|
Prefix: "test4:",
|
||||||
|
Store: store,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Clean up first
|
||||||
|
client.store.Clear()
|
||||||
|
|
||||||
|
testClient := createTestClient("test-client-4")
|
||||||
|
|
||||||
|
// Create client first
|
||||||
|
_, err = client.CreateClient(ctx, testClient)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Get client
|
||||||
|
retrieved, err := client.GetClientByID(ctx, testClient.ClientID)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, retrieved)
|
||||||
|
assert.Equal(t, testClient.ClientID, retrieved.ClientID)
|
||||||
|
assert.Equal(t, testClient.ClientSecret, retrieved.ClientSecret)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("get client with cache hit", func(t *testing.T) {
|
||||||
|
store := config.GetFunc(t)
|
||||||
|
cache := getLRUCache(t)
|
||||||
|
client, err := NewDefaultClient(&DefaultClientOptions{
|
||||||
|
Prefix: "test5:",
|
||||||
|
Store: store,
|
||||||
|
Cache: cache,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Clean up first
|
||||||
|
client.store.Clear()
|
||||||
|
client.cache.Clear()
|
||||||
|
|
||||||
|
testClient := createTestClient("test-client-5")
|
||||||
|
|
||||||
|
// Create client first
|
||||||
|
_, err = client.CreateClient(ctx, testClient)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Get client (should hit cache)
|
||||||
|
retrieved, err := client.GetClientByID(ctx, testClient.ClientID)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, retrieved)
|
||||||
|
assert.Equal(t, testClient.ClientID, retrieved.ClientID)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("get non-existent client", func(t *testing.T) {
|
||||||
|
store := config.GetFunc(t)
|
||||||
|
client, err := NewDefaultClient(&DefaultClientOptions{
|
||||||
|
Prefix: "test6:",
|
||||||
|
Store: store,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
retrieved, err := client.GetClientByID(ctx, "non-existent")
|
||||||
|
assert.Error(t, err)
|
||||||
|
assert.Nil(t, retrieved)
|
||||||
|
assert.Contains(t, err.Error(), "Client not found")
|
||||||
|
})
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDeleteClient(t *testing.T) {
|
||||||
|
storeConfigs := getStoreConfigs()
|
||||||
|
|
||||||
|
for _, config := range storeConfigs {
|
||||||
|
t.Run(config.Name, func(t *testing.T) {
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
t.Run("delete client with cache", func(t *testing.T) {
|
||||||
|
store := config.GetFunc(t)
|
||||||
|
cache := getLRUCache(t)
|
||||||
|
client, err := NewDefaultClient(&DefaultClientOptions{
|
||||||
|
Prefix: "test7:",
|
||||||
|
Store: store,
|
||||||
|
Cache: cache,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Clean up first
|
||||||
|
client.store.Clear()
|
||||||
|
client.cache.Clear()
|
||||||
|
|
||||||
|
testClient := createTestClient("test-client-7")
|
||||||
|
|
||||||
|
// Create client first
|
||||||
|
_, err = client.CreateClient(ctx, testClient)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Verify client is cached
|
||||||
|
key := client.clientKey(testClient.ClientID)
|
||||||
|
_, ok := client.cache.Get(key)
|
||||||
|
assert.True(t, ok)
|
||||||
|
|
||||||
|
// Delete client
|
||||||
|
err = client.DeleteClient(ctx, testClient.ClientID)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
|
||||||
|
// Verify cache is cleared
|
||||||
|
_, ok = client.cache.Get(key)
|
||||||
|
assert.False(t, ok)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("delete non-existent client", func(t *testing.T) {
|
||||||
|
store := config.GetFunc(t)
|
||||||
|
client, err := NewDefaultClient(&DefaultClientOptions{
|
||||||
|
Prefix: "test8:",
|
||||||
|
Store: store,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
err = client.DeleteClient(ctx, "non-existent")
|
||||||
|
assert.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "Client not found")
|
||||||
|
})
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestValidateClient(t *testing.T) {
|
||||||
|
storeConfigs := getStoreConfigs()
|
||||||
|
|
||||||
|
for _, config := range storeConfigs {
|
||||||
|
t.Run(config.Name, func(t *testing.T) {
|
||||||
|
store := config.GetFunc(t)
|
||||||
|
client, err := NewDefaultClient(&DefaultClientOptions{
|
||||||
|
Prefix: "test9:",
|
||||||
|
Store: store,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
t.Run("valid client", func(t *testing.T) {
|
||||||
|
testClient := createTestClient("test-client-9")
|
||||||
|
|
||||||
|
result, err := client.ValidateClient(ctx, testClient)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.True(t, result.Valid)
|
||||||
|
assert.Empty(t, result.Errors)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("client without ID", func(t *testing.T) {
|
||||||
|
testClient := createTestClient("")
|
||||||
|
|
||||||
|
result, err := client.ValidateClient(ctx, testClient)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.False(t, result.Valid)
|
||||||
|
assert.Contains(t, result.Errors, "Client ID is required")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("client with invalid type", func(t *testing.T) {
|
||||||
|
testClient := createTestClient("test-client-10")
|
||||||
|
testClient.ClientType = "invalid"
|
||||||
|
|
||||||
|
result, err := client.ValidateClient(ctx, testClient)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.False(t, result.Valid)
|
||||||
|
assert.Contains(t, result.Errors, "Invalid client type")
|
||||||
|
})
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCacheConsistency(t *testing.T) {
|
||||||
|
storeConfigs := getStoreConfigs()
|
||||||
|
|
||||||
|
for _, config := range storeConfigs {
|
||||||
|
t.Run(config.Name, func(t *testing.T) {
|
||||||
|
store := config.GetFunc(t)
|
||||||
|
cache := getLRUCache(t)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
client, err := NewDefaultClient(&DefaultClientOptions{
|
||||||
|
Prefix: "test10:",
|
||||||
|
Store: store,
|
||||||
|
Cache: cache,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Clean up first
|
||||||
|
client.store.Clear()
|
||||||
|
client.cache.Clear()
|
||||||
|
|
||||||
|
testClient := createTestClient("test-client-10")
|
||||||
|
|
||||||
|
// Create client
|
||||||
|
_, err = client.CreateClient(ctx, testClient)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Verify cache is updated
|
||||||
|
key := client.clientKey(testClient.ClientID)
|
||||||
|
cached, ok := client.cache.Get(key)
|
||||||
|
assert.True(t, ok)
|
||||||
|
assert.NotNil(t, cached)
|
||||||
|
|
||||||
|
// Update client
|
||||||
|
updateData := &types.ClientInfo{
|
||||||
|
ClientID: testClient.ClientID,
|
||||||
|
ClientSecret: "updated-secret",
|
||||||
|
ClientName: "Updated Client",
|
||||||
|
ClientType: types.ClientTypeConfidential,
|
||||||
|
RedirectURIs: []string{"https://updated.com/callback"},
|
||||||
|
GrantTypes: []string{types.GrantTypeAuthorizationCode},
|
||||||
|
ResponseTypes: []string{types.ResponseTypeCode},
|
||||||
|
Scope: "openid profile",
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err = client.UpdateClient(ctx, testClient.ClientID, updateData)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Verify cache is updated
|
||||||
|
cached, ok = client.cache.Get(key)
|
||||||
|
assert.True(t, ok)
|
||||||
|
cachedClient := cached.(*types.ClientInfo)
|
||||||
|
assert.Equal(t, "Updated Client", cachedClient.ClientName)
|
||||||
|
assert.Equal(t, "updated-secret", cachedClient.ClientSecret)
|
||||||
|
|
||||||
|
// Delete client
|
||||||
|
err = client.DeleteClient(ctx, testClient.ClientID)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Verify cache is cleared
|
||||||
|
_, ok = client.cache.Get(key)
|
||||||
|
assert.False(t, ok)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
52
openapi/oauth/providers/user/default.go
Normal file
52
openapi/oauth/providers/user/default.go
Normal file
|
|
@ -0,0 +1,52 @@
|
||||||
|
package user
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"github.com/yaoapp/yao/openapi/oauth/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
// DefaultUserProvider provides a default implementation of UserProvider
|
||||||
|
type DefaultUserProvider struct {
|
||||||
|
getUserByAccessTokenFunc func(ctx context.Context, accessToken string) (interface{}, error)
|
||||||
|
getUserBySubjectFunc func(ctx context.Context, subject string) (interface{}, error)
|
||||||
|
validateUserScopeFunc func(ctx context.Context, userID string, scopes []string) (bool, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewDefaultUserProvider creates a new DefaultUserProvider with the given functions
|
||||||
|
func NewDefaultUserProvider(
|
||||||
|
getUserByAccessTokenFunc func(ctx context.Context, accessToken string) (interface{}, error),
|
||||||
|
getUserBySubjectFunc func(ctx context.Context, subject string) (interface{}, error),
|
||||||
|
validateUserScopeFunc func(ctx context.Context, userID string, scopes []string) (bool, error),
|
||||||
|
) *DefaultUserProvider {
|
||||||
|
return &DefaultUserProvider{
|
||||||
|
getUserByAccessTokenFunc: getUserByAccessTokenFunc,
|
||||||
|
getUserBySubjectFunc: getUserBySubjectFunc,
|
||||||
|
validateUserScopeFunc: validateUserScopeFunc,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetUserByAccessToken retrieves user information using an access token
|
||||||
|
func (p *DefaultUserProvider) GetUserByAccessToken(ctx context.Context, accessToken string) (interface{}, error) {
|
||||||
|
if p.getUserByAccessTokenFunc == nil {
|
||||||
|
return nil, &types.ErrorResponse{Code: "not_implemented", ErrorDescription: "GetUserByAccessToken is not implemented"}
|
||||||
|
}
|
||||||
|
return p.getUserByAccessTokenFunc(ctx, accessToken)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetUserBySubject retrieves user information using a subject identifier
|
||||||
|
func (p *DefaultUserProvider) GetUserBySubject(ctx context.Context, subject string) (interface{}, error) {
|
||||||
|
if p.getUserBySubjectFunc == nil {
|
||||||
|
return nil, &types.ErrorResponse{Code: "not_implemented", ErrorDescription: "GetUserBySubject is not implemented"}
|
||||||
|
}
|
||||||
|
return p.getUserBySubjectFunc(ctx, subject)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateUserScope validates if a user has access to requested scopes
|
||||||
|
func (p *DefaultUserProvider) ValidateUserScope(ctx context.Context, userID string, scopes []string) (bool, error) {
|
||||||
|
if p.validateUserScopeFunc == nil {
|
||||||
|
// Default implementation: allow all scopes
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
|
return p.validateUserScopeFunc(ctx, userID, scopes)
|
||||||
|
}
|
||||||
50
openapi/oauth/security.go
Normal file
50
openapi/oauth/security.go
Normal file
|
|
@ -0,0 +1,50 @@
|
||||||
|
package oauth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"github.com/yaoapp/yao/openapi/oauth/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
// GenerateCodeChallenge generates a code challenge from a code verifier
|
||||||
|
// This is used for PKCE (Proof Key for Code Exchange) flow
|
||||||
|
func (s *Service) GenerateCodeChallenge(ctx context.Context, codeVerifier string, method string) (string, error) {
|
||||||
|
// TODO: Implement code challenge generation
|
||||||
|
return "", nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateCodeChallenge validates a code verifier against a code challenge
|
||||||
|
// This verifies the PKCE code challenge during token exchange
|
||||||
|
func (s *Service) ValidateCodeChallenge(ctx context.Context, codeVerifier string, codeChallenge string, method string) error {
|
||||||
|
// TODO: Implement code challenge validation
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateStateParameter validates OAuth state parameters
|
||||||
|
// This prevents CSRF attacks by verifying state parameters
|
||||||
|
func (s *Service) ValidateStateParameter(ctx context.Context, state string, clientID string) (*types.ValidationResult, error) {
|
||||||
|
// TODO: Implement state parameter validation
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GenerateStateParameter generates a secure state parameter
|
||||||
|
// This creates cryptographically secure state values for CSRF protection
|
||||||
|
func (s *Service) GenerateStateParameter(ctx context.Context, clientID string) (*types.StateParameter, error) {
|
||||||
|
// TODO: Implement state parameter generation
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateRedirectURI validates redirect URIs against registered URIs
|
||||||
|
func (s *Service) ValidateRedirectURI(ctx context.Context, redirectURI string, registeredURIs []string) (*types.ValidationResult, error) {
|
||||||
|
// This method signature doesn't match our ClientProvider interface
|
||||||
|
// We need the clientID to validate, so let's assume we can extract it from context
|
||||||
|
// or we need to modify the interface
|
||||||
|
return &types.ValidationResult{Valid: true}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// PushAuthorizationRequest processes a pushed authorization request
|
||||||
|
// This implements RFC 9126 for enhanced security
|
||||||
|
func (s *Service) PushAuthorizationRequest(ctx context.Context, request *types.PushedAuthorizationRequest) (*types.PushedAuthorizationResponse, error) {
|
||||||
|
// TODO: Implement pushed authorization request
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
35
openapi/oauth/token.go
Normal file
35
openapi/oauth/token.go
Normal file
|
|
@ -0,0 +1,35 @@
|
||||||
|
package oauth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"github.com/yaoapp/yao/openapi/oauth/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Introspect returns information about an access token
|
||||||
|
// This endpoint allows resource servers to validate tokens
|
||||||
|
func (s *Service) Introspect(ctx context.Context, token string) (*types.TokenIntrospectionResponse, error) {
|
||||||
|
// TODO: Implement token introspection
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// TokenExchange exchanges one token for another token
|
||||||
|
// This implements RFC 8693 for token exchange scenarios
|
||||||
|
func (s *Service) TokenExchange(ctx context.Context, subjectToken string, subjectTokenType string, audience string, scope string) (*types.TokenExchangeResponse, error) {
|
||||||
|
// TODO: Implement token exchange
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateTokenAudience validates token audience claims
|
||||||
|
// This ensures tokens are only used with their intended audiences
|
||||||
|
func (s *Service) ValidateTokenAudience(ctx context.Context, token string, expectedAudience string) (*types.ValidationResult, error) {
|
||||||
|
// TODO: Implement token audience validation
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ValidateTokenBinding validates token binding information
|
||||||
|
// This ensures tokens are bound to the correct client or device
|
||||||
|
func (s *Service) ValidateTokenBinding(ctx context.Context, token string, binding *types.TokenBinding) (*types.ValidationResult, error) {
|
||||||
|
// TODO: Implement token binding validation
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
11
openapi/oauth/types/errors.go
Normal file
11
openapi/oauth/types/errors.go
Normal file
|
|
@ -0,0 +1,11 @@
|
||||||
|
package types
|
||||||
|
|
||||||
|
// Error definitions
|
||||||
|
var (
|
||||||
|
ErrInvalidConfiguration = &ErrorResponse{Code: "invalid_configuration", ErrorDescription: "Invalid OAuth service configuration"}
|
||||||
|
ErrStoreMissing = &ErrorResponse{Code: "store_missing", ErrorDescription: "Store is required for OAuth service"}
|
||||||
|
ErrIssuerURLMissing = &ErrorResponse{Code: "issuer_url_missing", ErrorDescription: "Issuer URL is required for OAuth service"}
|
||||||
|
ErrCertificateMissing = &ErrorResponse{Code: "certificate_missing", ErrorDescription: "JWT signing certificate and key are required"}
|
||||||
|
ErrInvalidTokenLifetime = &ErrorResponse{Code: "invalid_token_lifetime", ErrorDescription: "Token lifetime must be greater than 0"}
|
||||||
|
ErrPKCEConfigurationInvalid = &ErrorResponse{Code: "pkce_configuration_invalid", ErrorDescription: "PKCE configuration is invalid"}
|
||||||
|
)
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
package oauth
|
package types
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
|
@ -146,3 +146,40 @@ type UserProvider interface {
|
||||||
// ValidateUserScope validates if a user has access to requested scopes
|
// ValidateUserScope validates if a user has access to requested scopes
|
||||||
ValidateUserScope(ctx context.Context, userID string, scopes []string) (bool, error)
|
ValidateUserScope(ctx context.Context, userID string, scopes []string) (bool, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ClientProvider interface for OAuth client management and persistence
|
||||||
|
type ClientProvider interface {
|
||||||
|
// GetClientByID retrieves client information using a client ID
|
||||||
|
GetClientByID(ctx context.Context, clientID string) (*ClientInfo, error)
|
||||||
|
|
||||||
|
// GetClientByCredentials retrieves and validates client using client credentials
|
||||||
|
// Used for client authentication in token requests
|
||||||
|
GetClientByCredentials(ctx context.Context, clientID string, clientSecret string) (*ClientInfo, error)
|
||||||
|
|
||||||
|
// CreateClient creates a new OAuth client and returns the client information
|
||||||
|
CreateClient(ctx context.Context, clientInfo *ClientInfo) (*ClientInfo, error)
|
||||||
|
|
||||||
|
// UpdateClient updates an existing OAuth client configuration
|
||||||
|
UpdateClient(ctx context.Context, clientID string, clientInfo *ClientInfo) (*ClientInfo, error)
|
||||||
|
|
||||||
|
// DeleteClient removes an OAuth client from the system
|
||||||
|
// This should also invalidate all associated tokens
|
||||||
|
DeleteClient(ctx context.Context, clientID string) error
|
||||||
|
|
||||||
|
// ValidateClient validates client information and configuration
|
||||||
|
// Returns validation result with any errors or warnings
|
||||||
|
ValidateClient(ctx context.Context, clientInfo *ClientInfo) (*ValidationResult, error)
|
||||||
|
|
||||||
|
// ListClients retrieves a list of clients with optional filtering
|
||||||
|
// Supports pagination and filtering by various criteria
|
||||||
|
ListClients(ctx context.Context, filters map[string]interface{}, limit int, offset int) ([]*ClientInfo, int, error)
|
||||||
|
|
||||||
|
// ValidateRedirectURI validates if a redirect URI is registered for the client
|
||||||
|
ValidateRedirectURI(ctx context.Context, clientID string, redirectURI string) (*ValidationResult, error)
|
||||||
|
|
||||||
|
// ValidateScope validates if the client is authorized to request specific scopes
|
||||||
|
ValidateScope(ctx context.Context, clientID string, scopes []string) (*ValidationResult, error)
|
||||||
|
|
||||||
|
// IsClientActive checks if a client is active and can be used for authentication
|
||||||
|
IsClientActive(ctx context.Context, clientID string) (bool, error)
|
||||||
|
}
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
package oauth
|
package types
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"time"
|
"time"
|
||||||
|
|
@ -218,6 +218,7 @@ type ClientInfo struct {
|
||||||
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method,omitempty"`
|
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method,omitempty"`
|
||||||
CreatedAt time.Time `json:"created_at,omitempty"`
|
CreatedAt time.Time `json:"created_at,omitempty"`
|
||||||
UpdatedAt time.Time `json:"updated_at,omitempty"`
|
UpdatedAt time.Time `json:"updated_at,omitempty"`
|
||||||
|
Extra map[string]interface{} `json:"extra,omitempty"` // Extra fields for custom client properties
|
||||||
}
|
}
|
||||||
|
|
||||||
// AuthorizationServerMetadata represents OAuth 2.0 Authorization Server Metadata (RFC 8414)
|
// AuthorizationServerMetadata represents OAuth 2.0 Authorization Server Metadata (RFC 8414)
|
||||||
|
|
@ -4,47 +4,15 @@ import (
|
||||||
"context"
|
"context"
|
||||||
)
|
)
|
||||||
|
|
||||||
// DefaultUserProvider provides a default implementation of UserProvider
|
// UserInfo returns user information for a given access token
|
||||||
type DefaultUserProvider struct {
|
func (s *Service) UserInfo(ctx context.Context, accessToken string) (interface{}, error) {
|
||||||
getUserByAccessTokenFunc func(ctx context.Context, accessToken string) (interface{}, error)
|
return s.userProvider.GetUserByAccessToken(ctx, accessToken)
|
||||||
getUserBySubjectFunc func(ctx context.Context, subject string) (interface{}, error)
|
|
||||||
validateUserScopeFunc func(ctx context.Context, userID string, scopes []string) (bool, error)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewDefaultUserProvider creates a new DefaultUserProvider with the given functions
|
// Additional user-related helper methods can be added here as needed
|
||||||
func NewDefaultUserProvider(
|
// For example:
|
||||||
getUserByAccessTokenFunc func(ctx context.Context, accessToken string) (interface{}, error),
|
// - User profile management
|
||||||
getUserBySubjectFunc func(ctx context.Context, subject string) (interface{}, error),
|
// - User consent handling
|
||||||
validateUserScopeFunc func(ctx context.Context, userID string, scopes []string) (bool, error),
|
// - User authentication verification
|
||||||
) *DefaultUserProvider {
|
// - User scope validation
|
||||||
return &DefaultUserProvider{
|
// etc.
|
||||||
getUserByAccessTokenFunc: getUserByAccessTokenFunc,
|
|
||||||
getUserBySubjectFunc: getUserBySubjectFunc,
|
|
||||||
validateUserScopeFunc: validateUserScopeFunc,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetUserByAccessToken retrieves user information using an access token
|
|
||||||
func (p *DefaultUserProvider) GetUserByAccessToken(ctx context.Context, accessToken string) (interface{}, error) {
|
|
||||||
if p.getUserByAccessTokenFunc == nil {
|
|
||||||
return nil, &ErrorResponse{Code: "not_implemented", ErrorDescription: "GetUserByAccessToken is not implemented"}
|
|
||||||
}
|
|
||||||
return p.getUserByAccessTokenFunc(ctx, accessToken)
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetUserBySubject retrieves user information using a subject identifier
|
|
||||||
func (p *DefaultUserProvider) GetUserBySubject(ctx context.Context, subject string) (interface{}, error) {
|
|
||||||
if p.getUserBySubjectFunc == nil {
|
|
||||||
return nil, &ErrorResponse{Code: "not_implemented", ErrorDescription: "GetUserBySubject is not implemented"}
|
|
||||||
}
|
|
||||||
return p.getUserBySubjectFunc(ctx, subject)
|
|
||||||
}
|
|
||||||
|
|
||||||
// ValidateUserScope validates if a user has access to requested scopes
|
|
||||||
func (p *DefaultUserProvider) ValidateUserScope(ctx context.Context, userID string, scopes []string) (bool, error) {
|
|
||||||
if p.validateUserScopeFunc == nil {
|
|
||||||
// Default implementation: allow all scopes
|
|
||||||
return true, nil
|
|
||||||
}
|
|
||||||
return p.validateUserScopeFunc(ctx, userID, scopes)
|
|
||||||
}
|
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue