Skip to content
Closed
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
2 changes: 1 addition & 1 deletion api/admin_info.go
Original file line number Diff line number Diff line change
Expand Up @@ -65,7 +65,7 @@ func getAdminInfo(w http.ResponseWriter, r *http.Request) {
// Runners
runnersInfo := map[string]any{
"use_remote_runner": util.Config.IsUseRemoteRunner(),
"default_global_runner_mode": util.Config.Runners.DefaultGlobalRunnersMode,
"default_global_runner_mode": util.Config.DefaultGlobalRunnersMode(),
}

// Task settings
Expand Down
46 changes: 30 additions & 16 deletions api/login.go
Original file line number Diff line number Diff line change
Expand Up @@ -812,6 +812,35 @@ func claimOidcToken(idToken *oidc.IDToken, provider util.OidcProvider) (res clai
return
}

// oidcClaimsFromUserInfo builds login claims from a UserInfo response when the
// token has no id_token. The caller must pass the UserInfo error separately so
// a failed fetch is not followed by a nil dereference on userInfo.
func oidcClaimsFromUserInfo(userInfo *oidc.UserInfo, userInfoErr error, provider util.OidcProvider) (claimResult, error) {
if userInfoErr != nil {
return claimResult{}, userInfoErr
}

var claims claimResult
var err error
if userInfo.Email == "" {
claims, err = claimOidcUserInfo(userInfo, provider)
if err != nil {
return claimResult{}, err
}
} else {
claims.email = userInfo.Email
claims.name = userInfo.Profile
claims.sub = userInfo.Subject
claims.emailVerified = oidcEmailVerified(userInfo, provider)
}

claims.username = getRandomUsername()
if userInfo.Profile == "" {
claims.name = getRandomProfileName()
}
return claims, nil
}

func getRandomUsername() string {
return random.String(16)
}
Expand Down Expand Up @@ -922,22 +951,7 @@ func oidcRedirect(w http.ResponseWriter, r *http.Request) {
} else {
var userInfo *oidc.UserInfo
userInfo, err = _oidc.UserInfo(ctx, oauth2.StaticTokenSource(oauth2Token))

if err == nil {
if userInfo.Email == "" {
claims, err = claimOidcUserInfo(userInfo, provider)
} else {
claims.email = userInfo.Email
claims.name = userInfo.Profile
claims.sub = userInfo.Subject
claims.emailVerified = oidcEmailVerified(userInfo, provider)
}
}

claims.username = getRandomUsername()
if userInfo.Profile == "" {
claims.name = getRandomProfileName()
}
claims, err = oidcClaimsFromUserInfo(userInfo, err, provider)
}

if err != nil {
Expand Down
23 changes: 23 additions & 0 deletions api/login_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,12 +3,14 @@ package api
import (
"encoding/base64"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"net/url"
"testing"
"time"

"github.com/coreos/go-oidc/v3/oidc"
"github.com/gorilla/mux"
"github.com/semaphoreui/semaphore/util"
"github.com/stretchr/testify/assert"
Expand Down Expand Up @@ -328,3 +330,24 @@ func TestOidcSuccessRedirectURL_NotRelative(t *testing.T) {
callback.ResolveReference(location).String(),
"the browser must not resolve the redirect against the callback path")
}

func TestOidcClaimsFromUserInfo_UserInfoFailure(t *testing.T) {
_, err := oidcClaimsFromUserInfo(nil, errors.New("userinfo unavailable"), util.OidcProvider{})
require.Error(t, err)
assert.EqualError(t, err, "userinfo unavailable")
}

func TestOidcClaimsFromUserInfo_WithEmail(t *testing.T) {
userInfo := &oidc.UserInfo{
Subject: "sub-123",
Profile: "Jane Doe",
Email: "jane@example.com",
}

claims, err := oidcClaimsFromUserInfo(userInfo, nil, util.OidcProvider{})
require.NoError(t, err)
assert.Equal(t, "sub-123", claims.sub)
assert.Equal(t, "jane@example.com", claims.email)
assert.Equal(t, "Jane Doe", claims.name)
assert.NotEmpty(t, claims.username)
}
2 changes: 1 addition & 1 deletion db/sql/SqlDb.go
Original file line number Diff line number Diff line change
Expand Up @@ -116,7 +116,7 @@ func (d *SqlDbConnection) Connect() {
}

func (d *SqlDbConnection) Close() {
if d.sql.Db == nil {
if d.sql == nil || d.sql.Db == nil {
return
}
err := d.sql.Db.Close()
Expand Down
10 changes: 10 additions & 0 deletions db/sql/SqlDb_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,16 @@ import (
"github.com/go-gorp/gorp/v3"
)

func TestSqlDbConnection_CloseWithoutConnect(t *testing.T) {
var conn SqlDbConnection
conn.Close() // must not panic when never connected
}

func TestSqlDb_CloseWithoutConnect(t *testing.T) {
store := CreateDb("sqlite")
store.Close() // must not panic when never connected
}

func TestValidatePort(t *testing.T) {
d := SqlDb{}
q := d.connection.prepareQueryWithDialect("select * from `test` where id = ?, email = ?", gorp.PostgresDialect{})
Expand Down
3 changes: 2 additions & 1 deletion db/sql/migration.go
Original file line number Diff line number Diff line change
Expand Up @@ -312,7 +312,8 @@ func (d *SqlDb) TryRollbackMigration(version db.Migration) {
return
}

queries := getVersionSQL(d.GetDialect(), getVersionErrPath(version), false)
// Most migrations have no undo SQL; a missing .err.sql file must not panic.
queries := getVersionSQL(d.GetDialect(), getVersionErrPath(version), true)

for _, query := range queries {
fmt.Printf(" [ROLLBACK] > %v\n", query)
Expand Down
20 changes: 20 additions & 0 deletions db/sql/migration_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,20 @@
package sql

import (
"testing"

"github.com/stretchr/testify/assert"
)

func TestGetVersionSQL_MissingErrFileWithIgnoreErrors(t *testing.T) {
assert.NotPanics(t, func() {
queries := getVersionSQL("sqlite", "v9999.9.9.err.sql", true)
assert.Nil(t, queries)
})
}

func TestGetVersionSQL_MissingErrFileWithoutIgnoreErrorsPanics(t *testing.T) {
assert.Panics(t, func() {
getVersionSQL("sqlite", "v9999.9.9.err.sql", false)
})
}
26 changes: 16 additions & 10 deletions util/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -707,23 +707,29 @@ func (conf *ConfigType) GetSshConfigPath() string {
}

func (conf *ConfigType) GetRunnerRegistrationToken() string {
if conf.Runners.RegistrationToken != "" {
if conf.Runners != nil && conf.Runners.RegistrationToken != "" {
return conf.Runners.RegistrationToken
}
return conf.RunnerRegistrationToken
}

func (conf *ConfigType) IsUseRemoteRunner() bool {
switch conf.Runners.DefaultGlobalRunnersMode {
case DefultGlobalRunnerDisable:
return false
case DefultGlobalRunnerRequire:
return true
case DefultGlobalRunnerPrefer:
return true
default:
return conf.UseRemoteRunner
if conf.Runners != nil {
switch conf.Runners.DefaultGlobalRunnersMode {
case DefultGlobalRunnerDisable:
return false
case DefultGlobalRunnerRequire, DefultGlobalRunnerPrefer:
return true
}
}
return conf.UseRemoteRunner
}

func (conf *ConfigType) DefaultGlobalRunnersMode() DefultGlobalRunnerMode {
if conf.Runners != nil {
return conf.Runners.DefaultGlobalRunnersMode
}
return DefultGlobalRunnerNone
}

// RunnersOfflineTimeout returns the heartbeat staleness after which a runner
Expand Down
17 changes: 17 additions & 0 deletions util/config_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -585,3 +585,20 @@ func TestValidateConfig(t *testing.T) {
ensureConfigValidationFailure(t, "AccessKeyEncryption", Config.AccessKeyEncryption)
Config.AccessKeyEncryption = testCookieHash
}

func TestIsUseRemoteRunner_NilRunnersConfig(t *testing.T) {
conf := NewConfigType()
conf.UseRemoteRunner = true

assert.True(t, conf.IsUseRemoteRunner())
assert.Equal(t, DefultGlobalRunnerNone, conf.DefaultGlobalRunnersMode())
assert.Empty(t, conf.GetRunnerRegistrationToken())
}

func TestIsUseRemoteRunner_RunnersModeOverridesLegacyFlag(t *testing.T) {
conf := NewConfigType()
conf.UseRemoteRunner = true
conf.Runners = &RunnersConfig{DefaultGlobalRunnersMode: DefultGlobalRunnerDisable}

assert.False(t, conf.IsUseRemoteRunner())
}
Loading