Files
at-container-registry/pkg/auth/oauth/client.go
T
2025-10-04 13:50:28 -05:00

222 lines
6.6 KiB
Go

package oauth
import (
"context"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"fmt"
"net/http"
atprotoclient "atcr.io/pkg/atproto"
"authelia.com/client/oauth2"
)
// Client is an OAuth client for ATProto with DPoP support
type Client struct {
config *oauth2.Config
dpopKey *ecdsa.PrivateKey
dpopTransport *DPoPTransport
resolver *atprotoclient.Resolver
clientID string
redirectURI string
metadata *AuthServerMetadata
}
// NewClient creates a new OAuth client for ATProto
func NewClient(clientID, redirectURI string) (*Client, error) {
// Generate DPoP key
dpopKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
return nil, fmt.Errorf("failed to generate DPoP key: %w", err)
}
return &Client{
dpopKey: dpopKey,
dpopTransport: NewDPoPTransport(http.DefaultTransport, dpopKey),
resolver: atprotoclient.NewResolver(),
clientID: clientID,
redirectURI: redirectURI,
}, nil
}
// InitializeForHandle discovers the authorization server for a given handle/DID
func (c *Client) InitializeForHandle(ctx context.Context, handle string) error {
// Resolve handle to DID and PDS
_, pdsEndpoint, err := c.resolver.ResolveIdentity(ctx, handle)
if err != nil {
return fmt.Errorf("failed to resolve identity: %w", err)
}
// Discover authorization server metadata
metadata, err := DiscoverAuthServer(ctx, pdsEndpoint)
if err != nil {
return fmt.Errorf("failed to discover authorization server: %w", err)
}
c.metadata = metadata
// Configure OAuth2 client
// Note: Both localhost and production need redirect_uri and scopes in the config
// For localhost: client_id contains these (query-based) AND they're sent as params
// For production: client_id is metadata URL, params come from config
c.config = &oauth2.Config{
ClientID: c.clientID,
Endpoint: oauth2.Endpoint{
AuthURL: metadata.AuthorizationEndpoint,
TokenURL: metadata.TokenEndpoint,
PushedAuthURL: metadata.PushedAuthorizationRequestEndpoint,
},
RedirectURL: c.redirectURI,
Scopes: []string{"atproto"},
}
return nil
}
// SetScopes sets custom OAuth scopes (must be called after InitializeForHandle)
func (c *Client) SetScopes(scopes []string) {
if c.config != nil {
c.config.Scopes = scopes
}
}
// AuthorizeURL generates the authorization URL with PKCE
func (c *Client) AuthorizeURL(state string) (authURL string, codeVerifier string, err error) {
if c.config == nil {
return "", "", fmt.Errorf("client not initialized - call InitializeForHandle first")
}
// Generate PKCE code verifier
codeVerifier, err = generateCodeVerifier()
if err != nil {
return "", "", fmt.Errorf("failed to generate code verifier: %w", err)
}
// Generate code challenge
codeChallenge := generateCodeChallenge(codeVerifier)
// Use PAR (Pushed Authorization Request) if supported
if c.metadata.PushedAuthorizationRequestEndpoint != "" {
authURL, err = c.authorizeURLWithPAR(state, codeChallenge)
if err != nil {
return "", "", fmt.Errorf("PAR failed: %w", err)
}
} else {
// Fallback to standard authorization
authURL = c.config.AuthCodeURL(state,
oauth2.SetAuthURLParam("code_challenge", codeChallenge),
oauth2.SetAuthURLParam("code_challenge_method", "S256"),
)
}
return authURL, codeVerifier, nil
}
// authorizeURLWithPAR uses Pushed Authorization Request
func (c *Client) authorizeURLWithPAR(state, codeChallenge string) (string, error) {
fmt.Printf("DEBUG [oauth/client]: Starting PAR request\n")
fmt.Printf("DEBUG [oauth/client]: - client_id: %s\n", c.config.ClientID)
fmt.Printf("DEBUG [oauth/client]: - redirect_uri: %s\n", c.config.RedirectURL)
fmt.Printf("DEBUG [oauth/client]: - scope: %v\n", c.config.Scopes)
fmt.Printf("DEBUG [oauth/client]: - state: %s\n", state)
fmt.Printf("DEBUG [oauth/client]: - code_challenge_method: S256\n")
fmt.Printf("DEBUG [oauth/client]: - PAR endpoint: %s\n", c.config.Endpoint.PushedAuthURL)
// Create HTTP client with DPoP transport
ctx := context.WithValue(context.Background(), oauth2.HTTPClient, &http.Client{
Transport: c.dpopTransport,
})
// Use authelia's PushedAuth method
authURL, _, err := c.config.PushedAuth(ctx, state,
oauth2.SetAuthURLParam("code_challenge", codeChallenge),
oauth2.SetAuthURLParam("code_challenge_method", "S256"),
)
if err != nil {
fmt.Printf("ERROR [oauth/client]: PAR request failed: %v\n", err)
return "", err
}
fmt.Printf("DEBUG [oauth/client]: PAR successful, authURL: %s\n", authURL.String())
return authURL.String(), nil
}
// Exchange exchanges an authorization code for an access token
func (c *Client) Exchange(ctx context.Context, code, codeVerifier string) (*oauth2.Token, error) {
if c.config == nil {
return nil, fmt.Errorf("client not initialized")
}
// Create HTTP client with DPoP transport
ctx = context.WithValue(ctx, oauth2.HTTPClient, &http.Client{
Transport: c.dpopTransport,
})
// Exchange the code for a token
token, err := c.config.Exchange(ctx, code,
oauth2.SetAuthURLParam("code_verifier", codeVerifier),
)
if err != nil {
return nil, fmt.Errorf("failed to exchange code: %w", err)
}
return token, nil
}
// RefreshToken refreshes an access token using a refresh token
func (c *Client) RefreshToken(ctx context.Context, refreshToken string) (*oauth2.Token, error) {
if c.config == nil {
return nil, fmt.Errorf("client not initialized")
}
// Create HTTP client with DPoP transport
ctx = context.WithValue(ctx, oauth2.HTTPClient, &http.Client{
Transport: c.dpopTransport,
})
// Create a token source with the refresh token
token := &oauth2.Token{
RefreshToken: refreshToken,
}
// Refresh the token
newToken, err := c.config.TokenSource(ctx, token).Token()
if err != nil {
return nil, fmt.Errorf("failed to refresh token: %w", err)
}
return newToken, nil
}
// DPoPKey returns the DPoP private key
func (c *Client) DPoPKey() *ecdsa.PrivateKey {
return c.dpopKey
}
// SetDPoPKey sets the DPoP private key (useful when loading from storage)
func (c *Client) SetDPoPKey(key *ecdsa.PrivateKey) {
c.dpopKey = key
c.dpopTransport = NewDPoPTransport(http.DefaultTransport, key)
}
// generateCodeVerifier generates a PKCE code verifier
func generateCodeVerifier() (string, error) {
// Generate 32 random bytes
bytes := make([]byte, 32)
if _, err := rand.Read(bytes); err != nil {
return "", err
}
// Base64 URL encode
return base64.RawURLEncoding.EncodeToString(bytes), nil
}
// generateCodeChallenge generates a PKCE code challenge from a verifier
func generateCodeChallenge(verifier string) string {
hash := sha256.Sum256([]byte(verifier))
return base64.RawURLEncoding.EncodeToString(hash[:])
}