diff --git a/oauth/helpers.go b/oauth/helpers.go index 51ab5b0..5a183c3 100644 --- a/oauth/helpers.go +++ b/oauth/helpers.go @@ -4,9 +4,11 @@ import ( "errors" "fmt" "net/url" + "time" "github.com/haileyok/cocoon/internal/helpers" "github.com/haileyok/cocoon/oauth/constants" + "github.com/haileyok/cocoon/oauth/provider" ) func GenerateCode() string { @@ -46,3 +48,33 @@ func DecodeRequestUri(reqUri string) (string, error) { return reqId, nil } + +type SessionAgeResult struct { + SessionAge time.Duration + RefreshAge time.Duration + SessionExpired bool + RefreshExpired bool +} + +func GetSessionAgeFromToken(t provider.OauthToken) SessionAgeResult { + sessionLifetime := constants.PublicClientSessionLifetime + refreshLifetime := constants.PublicClientRefreshLifetime + if t.ClientAuth.Method != "none" { + sessionLifetime = constants.ConfidentialClientSessionLifetime + refreshLifetime = constants.ConfidentialClientRefreshLifetime + } + + res := SessionAgeResult{} + + res.SessionAge = time.Since(t.CreatedAt) + if res.SessionAge > sessionLifetime { + res.SessionExpired = true + } + + refreshAge := time.Since(t.UpdatedAt) + if refreshAge > refreshLifetime { + res.RefreshExpired = true + } + + return res +} diff --git a/server/handle_account.go b/server/handle_account.go index ce17693..611d4ea 100644 --- a/server/handle_account.go +++ b/server/handle_account.go @@ -3,6 +3,8 @@ package server import ( "time" + "github.com/haileyok/cocoon/oauth" + "github.com/haileyok/cocoon/oauth/constants" "github.com/haileyok/cocoon/oauth/provider" "github.com/labstack/echo/v4" ) @@ -13,10 +15,10 @@ func (s *Server) handleAccount(e echo.Context) error { return e.Redirect(303, "/account/signin") } - now := time.Now() + oldestPossibleSession := time.Now().Add(constants.ConfidentialClientSessionLifetime) var tokens []provider.OauthToken - if err := s.db.Raw("SELECT * FROM oauth_tokens WHERE sub = ? AND expires_at >= ? ORDER BY created_at ASC", nil, repo.Repo.Did, now).Scan(&tokens).Error; err != nil { + if err := s.db.Raw("SELECT * FROM oauth_tokens WHERE sub = ? AND created_at < ? ORDER BY created_at ASC", nil, repo.Repo.Did, oldestPossibleSession).Scan(&tokens).Error; err != nil { s.logger.Error("couldnt fetch oauth sessions for account", "did", repo.Repo.Did, "error", err) sess.AddFlash("Unable to fetch sessions. See server logs for more details.", "error") sess.Save(e.Request(), e.Response()) @@ -25,6 +27,15 @@ func (s *Server) handleAccount(e echo.Context) error { }) } + var filtered []provider.OauthToken + for _, t := range tokens { + ageRes := oauth.GetSessionAgeFromToken(t) + if ageRes.SessionExpired { + continue + } + filtered = append(filtered, t) + } + tokenInfo := []map[string]string{} for _, t := range tokens { tokenInfo = append(tokenInfo, map[string]string{ diff --git a/server/handle_oauth_token.go b/server/handle_oauth_token.go index 1ca61db..3b68c93 100644 --- a/server/handle_oauth_token.go +++ b/server/handle_oauth_token.go @@ -203,20 +203,13 @@ func (s *Server) handleOauthToken(e echo.Context) error { return helpers.InputError(e, to.StringPtr("dpop proof does not match expected jkt")) } - sessionLifetime := constants.PublicClientSessionLifetime - refreshLifetime := constants.PublicClientRefreshLifetime - if clientAuth.Method != "none" { - sessionLifetime = constants.ConfidentialClientSessionLifetime - refreshLifetime = constants.ConfidentialClientRefreshLifetime - } + ageRes := oauth.GetSessionAgeFromToken(oauthToken) - sessionAge := time.Since(oauthToken.CreatedAt) - if sessionAge > sessionLifetime { + if ageRes.SessionExpired { return helpers.InputError(e, to.StringPtr("Session expired")) } - refreshAge := time.Since(oauthToken.UpdatedAt) - if refreshAge > refreshLifetime { + if ageRes.RefreshExpired { return helpers.InputError(e, to.StringPtr("Refresh token expired")) }