mirror of
https://tangled.org/evan.jarrett.net/at-container-registry
synced 2026-09-29 21:45:33 +00:00
222 lines
6.6 KiB
Go
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[:])
|
|
}
|