removed intermediate /internal directory
This commit is contained in:
@@ -0,0 +1,190 @@
|
||||
package authentication
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgtype"
|
||||
|
||||
"ruben/inventory2/consts"
|
||||
)
|
||||
|
||||
// NewState creates a new state for logging in, saving it in the database.
|
||||
func (a *Authenticator) NewState(ctx context.Context, targetURI string) ([32]byte, error) {
|
||||
state, err := generateRandomState()
|
||||
if err != nil {
|
||||
return state, fmt.Errorf("failed to generate random state: %w", err)
|
||||
}
|
||||
|
||||
if _, err = a.db.Exec(
|
||||
ctx,
|
||||
`
|
||||
INSERT INTO
|
||||
oauth_login_states (
|
||||
state,
|
||||
target_uri
|
||||
)
|
||||
VALUES (
|
||||
@state,
|
||||
@target_uri
|
||||
)
|
||||
`,
|
||||
pgx.NamedArgs{
|
||||
"state": state[:],
|
||||
"target_uri": targetURI,
|
||||
},
|
||||
); err != nil {
|
||||
return state, fmt.Errorf("failed to execute query: %w", err)
|
||||
}
|
||||
|
||||
return state, nil
|
||||
}
|
||||
|
||||
func generateRandomState() ([32]byte, error) {
|
||||
var b [32]byte
|
||||
_, err := rand.Read(b[:])
|
||||
return b, err
|
||||
}
|
||||
|
||||
// GetStateExpirationAndURL get's the oauth state's expiration
|
||||
func (a *Authenticator) GetStateExpirationAndURL(ctx context.Context, state string) (time.Time, string, error) {
|
||||
rows, err := a.db.Query(
|
||||
ctx,
|
||||
`
|
||||
SELECT
|
||||
expiration,
|
||||
target_uri
|
||||
FROM
|
||||
oauth_login_states
|
||||
WHERE
|
||||
state = ('\x' || @state)::BYTEA`,
|
||||
pgx.NamedArgs{
|
||||
"state": state,
|
||||
},
|
||||
)
|
||||
if err != nil {
|
||||
return time.Time{}, "", fmt.Errorf("failed to perform query: %w", err)
|
||||
}
|
||||
|
||||
type Row struct {
|
||||
Expiration time.Time
|
||||
Target_uri pgtype.Text
|
||||
}
|
||||
|
||||
r, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByNameLax[Row])
|
||||
if err != nil {
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return time.Time{}, "", consts.ErrNotFound
|
||||
}
|
||||
|
||||
return time.Time{}, "", fmt.Errorf("failed to scan row: %w", err)
|
||||
}
|
||||
|
||||
return r.Expiration, r.Target_uri.String, nil
|
||||
}
|
||||
|
||||
// TODO: need to automatically clean up expired tokens
|
||||
func (a *Authenticator) DeleteOAuthTokens(ctx context.Context, accessToken string) error {
|
||||
_, err := a.db.Exec(ctx, "DELETE FROM oauth_tokens WHERE access_token = @access_token", pgx.NamedArgs{"access_token": accessToken})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to perform query: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (a *Authenticator) GetAccessTokenClaimsAndExpiration(ctx context.Context, accessToken string) (claims AccessTokenClaims, err error) {
|
||||
rows, err := a.db.Query(
|
||||
ctx,
|
||||
`
|
||||
SELECT
|
||||
expiry,
|
||||
|
||||
id_token_custom_claims_name,
|
||||
id_token_custom_claims_picture,
|
||||
id_token_custom_claims_nickname,
|
||||
id_token_custom_claims_given_name,
|
||||
id_token_custom_claims_family_name,
|
||||
id_token_custom_claims_updated_at
|
||||
|
||||
FROM
|
||||
oauth_tokens
|
||||
WHERE
|
||||
access_token = @access_token
|
||||
`,
|
||||
pgx.NamedArgs{
|
||||
"access_token": accessToken,
|
||||
},
|
||||
)
|
||||
if err != nil {
|
||||
return AccessTokenClaims{}, fmt.Errorf("failed to perform query: %w", err)
|
||||
}
|
||||
|
||||
type Row struct {
|
||||
Expiry time.Time
|
||||
Id_token_custom_claims_name string
|
||||
Id_token_custom_claims_picture string
|
||||
Id_token_custom_claims_nickname string
|
||||
Id_token_custom_claims_given_name string
|
||||
Id_token_custom_claims_family_name string
|
||||
Id_token_custom_claims_updated_at time.Time
|
||||
}
|
||||
|
||||
r, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByNameLax[Row])
|
||||
if err != nil {
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return AccessTokenClaims{}, consts.ErrNotFound
|
||||
}
|
||||
return AccessTokenClaims{}, fmt.Errorf("failed to scan row: %w", err)
|
||||
}
|
||||
|
||||
claims.Expires = r.Expiry.Unix()
|
||||
claims.Expiration = r.Expiry
|
||||
claims.Name = r.Id_token_custom_claims_name
|
||||
claims.Picture = r.Id_token_custom_claims_picture
|
||||
claims.Nickname = r.Id_token_custom_claims_nickname
|
||||
claims.GivenName = r.Id_token_custom_claims_given_name
|
||||
claims.FamilyName = r.Id_token_custom_claims_family_name
|
||||
claims.UpdatedAt = r.Id_token_custom_claims_updated_at
|
||||
|
||||
return claims, nil
|
||||
}
|
||||
|
||||
func (a *Authenticator) getRefreshTokenForAccessToken(ctx context.Context, accessToken string) (refreshToken, tokenType string, err error) {
|
||||
rows, err := a.db.Query(
|
||||
ctx,
|
||||
`
|
||||
SELECT
|
||||
refresh_token,
|
||||
token_type
|
||||
FROM
|
||||
oauth_tokens
|
||||
WHERE
|
||||
access_token = @access_token
|
||||
`,
|
||||
pgx.NamedArgs{
|
||||
"access_token": accessToken,
|
||||
},
|
||||
)
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("failed to perform query: %w", err)
|
||||
}
|
||||
|
||||
type Row struct {
|
||||
Refresh_token string
|
||||
Token_type string
|
||||
}
|
||||
|
||||
r, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByNameLax[Row])
|
||||
if err != nil {
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return "", "", consts.ErrNotFound
|
||||
}
|
||||
return "", "", fmt.Errorf("failed to scan row: %w", err)
|
||||
}
|
||||
|
||||
return r.Refresh_token, r.Token_type, nil
|
||||
}
|
||||
Reference in New Issue
Block a user