Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 28 additions & 0 deletions internal/models/sso.go
Original file line number Diff line number Diff line change
@@ -1,8 +1,10 @@
package models

import (
"crypto/sha256"
"database/sql"
"database/sql/driver"
"encoding/hex"
"encoding/json"
"net/url"
"reflect"
Expand All @@ -23,6 +25,8 @@ type SSOProvider struct {
SAMLProvider SAMLProvider `has_one:"saml_providers" fk_id:"sso_provider_id" json:"saml,omitempty"`
SSODomains []SSODomain `has_many:"sso_domains" fk_id:"sso_provider_id" json:"domains"`

SCIMTokenHash *string `db:"scim_token_hash" json:"-"`

CreatedAt time.Time `db:"created_at" json:"created_at"`
UpdatedAt time.Time `db:"updated_at" json:"updated_at"`
}
Expand All @@ -39,6 +43,16 @@ func (p SSOProvider) Type() string {
return "saml"
}

func (p *SSOProvider) UpdateSCIMToken(token string) {
hash := toSHA256(token)
p.SCIMTokenHash = &hash
}

func toSHA256(token string) string {
sum := sha256.Sum256([]byte(token))
return hex.EncodeToString(sum[:])
}

type SAMLAttribute struct {
Name string `json:"name,omitempty"`
Names []string `json:"names,omitempty"`
Expand Down Expand Up @@ -222,6 +236,20 @@ func FindSSOProviderByResourceID(tx *storage.Connection, id string) (*SSOProvide
return &ssoProvider, nil
}

func FindSSOProviderBySCIMToken(tx *storage.Connection, token string) (*SSOProvider, error) {
var ssoProvider SSOProvider

if err := tx.Q().Where("scim_token_hash = ?", toSHA256(token)).First(&ssoProvider); err != nil {
if errors.Cause(err) == sql.ErrNoRows {
return nil, SSOProviderNotFoundError{}
}

return nil, errors.Wrap(err, "error finding SSO provider by SCIM token")
}

return &ssoProvider, nil
}

func FindSSOProviderForEmailAddress(tx *storage.Connection, emailAddress string) (*SSOProvider, error) {
parts := strings.Split(emailAddress, "@")
emailDomain := strings.ToLower(parts[1])
Expand Down
73 changes: 73 additions & 0 deletions internal/models/sso_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -469,3 +469,76 @@ func (ts *SSOTestSuite) TestFindSSOProviderByResourceID() {
require.Nil(ts.T(), got)
}
}

func buildSSOProvider() *SSOProvider {
id := uuid.Must(uuid.NewV4()).String()

return &SSOProvider{
SAMLProvider: SAMLProvider{
EntityID: "https://example.com/saml/metadata/" + id,
MetadataXML: "<example />",
},
}
}

func (ts *SSOTestSuite) TestUpdateSCIMToken() {
hashes := map[string]string{
"scim_test_token": "dcbcd9ffd696ae1f2ee0f035fa17680d78175020a5fa1aadc758dbd681e0fe1d",
"scim_rotated_token": "289adb37f8946571bb4aea1e663281126c7f2d84d929ff09429fcaa1eb3f27bf",
}

provider := buildSSOProvider()
require.Nil(ts.T(), provider.SCIMTokenHash)

for token, hash := range hashes {
provider.UpdateSCIMToken(token)
require.NotNil(ts.T(), provider.SCIMTokenHash)
require.Equal(ts.T(), hash, *provider.SCIMTokenHash)
}
}

func (ts *SSOTestSuite) TestFindSSOProviderBySCIMToken() {
provider := buildSSOProvider()

token := "scim_test_token"
provider.UpdateSCIMToken(token)
require.NoError(ts.T(), ts.db.Eager().Create(provider))

withoutToken := buildSSOProvider()
require.NoError(ts.T(), ts.db.Eager().Create(withoutToken))

ts.Run("resolves the provider that owns the token", func() {
found, err := FindSSOProviderBySCIMToken(ts.db, token)

require.NoError(ts.T(), err)
require.Equal(ts.T(), provider.ID, found.ID)
})

ts.Run("an unknown token resolves nothing", func() {
found, err := FindSSOProviderBySCIMToken(ts.db, "scim_unknown_token")

require.Nil(ts.T(), found)
require.True(ts.T(), IsNotFoundError(err))
})

ts.Run("an empty token does not match a provider without one", func() {
found, err := FindSSOProviderBySCIMToken(ts.db, "")

require.Nil(ts.T(), found)
require.True(ts.T(), IsNotFoundError(err))
})

ts.Run("rotation stops the previous token from resolving", func() {
newToken := "scim_rotated_token"
provider.UpdateSCIMToken(newToken)
require.NoError(ts.T(), ts.db.Update(provider))

found, err := FindSSOProviderBySCIMToken(ts.db, newToken)
require.NoError(ts.T(), err)
require.Equal(ts.T(), provider.ID, found.ID)

found, err = FindSSOProviderBySCIMToken(ts.db, token)
require.Nil(ts.T(), found)
require.True(ts.T(), IsNotFoundError(err))
})
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,9 @@
-- Holds the SHA-256 hex digest of the provider's SCIM token.
/* auth_migration: 20260731000000 */
alter table only {{ index .Options "Namespace" }}.sso_providers
add column if not exists scim_token_hash text null;

/* auth_migration: 20260731000000 */
create unique index if not exists sso_providers_scim_token_hash_idx
on {{ index .Options "Namespace" }}.sso_providers (scim_token_hash)
where scim_token_hash is not null;
Loading