package authentication import ( "context" "crypto/rand" "errors" "fmt" "ruben/inventory2/internal/consts" "time" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgtype" ) // 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, expiration time.Time, 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{}, time.Time{}, 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{}, time.Time{}, consts.ErrNotFound } return AccessTokenClaims{}, time.Time{}, fmt.Errorf("failed to scan row: %w", err) } 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, r.Expiry, 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 }