Skip to content

Commit 31dd89f

Browse files
authored
Merge pull request #32 from thand-io/embed-provider-in-token
Add the provider into the session
2 parents 215b29a + e8aebd5 commit 31dd89f

12 files changed

Lines changed: 101 additions & 60 deletions

File tree

internal/daemon/auth.go

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -160,7 +160,7 @@ func (s *Server) getAuthCallbackPage(c *gin.Context, auth models.AuthWrapper) {
160160
state := c.Query("state")
161161
code := c.Query("code")
162162

163-
session, err := provider.GetClient().CreateSession(context.TODO(), &models.AuthorizeUser{
163+
session, err := provider.GetClient().CreateSession(c, &models.AuthorizeUser{
164164
State: state,
165165
Code: code,
166166
RedirectUri: s.GetConfig().GetAuthCallbackUrl(auth.Provider),
@@ -171,8 +171,14 @@ func (s *Server) getAuthCallbackPage(c *gin.Context, auth models.AuthWrapper) {
171171
return
172172
}
173173

174+
exportableSession := &models.ExportableSession{
175+
Session: session,
176+
Provider: auth.Provider,
177+
}
178+
174179
// Covert our sensitive session to one we can store on the users local system
175-
localSession := session.ToLocalSession(s.Config.GetServices().GetEncryption())
180+
localSession := exportableSession.ToLocalSession(
181+
s.Config.GetServices().GetEncryption())
176182

177183
data := AuthCallbackPageData{
178184
TemplateData: s.GetTemplateData(c),

internal/daemon/elevate.go

Lines changed: 39 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -86,33 +86,26 @@ func (s *Server) postElevateJSON(c *gin.Context) {
8686
return
8787
}
8888

89-
// Parse as raw JSON to detect request type
90-
var rawData map[string]any
91-
if err := json.Unmarshal(body, &rawData); err != nil {
92-
s.getErrorPage(c, http.StatusBadRequest, "Invalid JSON payload", err)
89+
// This is a standard elevation request
90+
var request models.ElevateRequest
91+
if err := json.Unmarshal(body, &request); err != nil {
92+
s.getErrorPage(c, http.StatusBadRequest, "Invalid standard request payload", err)
9393
return
9494
}
9595

96-
// Check if this is a dynamic request (has providers array, permissions array, etc.)
97-
if providers, hasProviders := rawData["providers"].([]any); hasProviders && len(providers) > 0 {
98-
// This is a dynamic request
99-
var dynamicRequest models.ElevateDynamicRequest
100-
if err := json.Unmarshal(body, &dynamicRequest); err != nil {
101-
s.getErrorPage(c, http.StatusBadRequest, "Invalid dynamic request payload", err)
102-
return
103-
}
104-
s.handleDynamicRequest(c, dynamicRequest)
96+
if request.IsValid() {
97+
s.elevate(c, request)
10598
return
10699
}
107100

108-
// This is a standard elevation request
109-
var request models.ElevateRequest
110-
if err := json.Unmarshal(body, &request); err != nil {
111-
s.getErrorPage(c, http.StatusBadRequest, "Invalid standard request payload", err)
101+
// Parse as raw JSON to detect request type
102+
var dynamicRequest models.ElevateDynamicRequest
103+
if err := json.Unmarshal(body, &dynamicRequest); err != nil {
104+
s.getErrorPage(c, http.StatusBadRequest, "Invalid dynamic request payload", err)
112105
return
113106
}
107+
s.handleDynamicRequest(c, dynamicRequest)
114108

115-
s.elevate(c, request)
116109
}
117110

118111
func (s *Server) handleDynamicRequest(c *gin.Context, dynamicRequest models.ElevateDynamicRequest) {
@@ -182,27 +175,35 @@ func (s *Server) elevate(c *gin.Context, request models.ElevateRequest) {
182175
// lets attach a user session to the request.
183176
if s.Config.IsServer() {
184177

185-
// Get the auth provider from the workflow if set
186-
authProvider := []string{}
178+
if len(request.Workflow) == 0 {
179+
s.getErrorPage(c, http.StatusBadRequest, "No workflow specified for elevation request")
180+
return
181+
}
187182

188-
if len(request.Workflow) > 0 {
189-
workflowDef, err := s.Config.GetWorkflowByName(request.Workflow)
190-
if err != nil {
191-
s.getErrorPage(c, http.StatusBadRequest, "Invalid workflow specified", err)
192-
return
193-
}
194-
authProvider = []string{workflowDef.GetAuthentication()}
183+
workflowDef, err := s.Config.GetWorkflowByName(request.Workflow)
184+
185+
if err != nil {
186+
s.getErrorPage(c, http.StatusBadRequest, "Invalid workflow specified", err)
187+
return
195188
}
196189

197-
foundUser, err := s.getUser(c, authProvider...)
190+
authProvider := workflowDef.GetAuthentication()
191+
192+
foundUser, err := s.getUser(c, authProvider)
198193

199194
if err != nil {
200195
s.getErrorPage(c, http.StatusUnauthorized, "Unauthorized: unable to get user for list of available roles", err)
201196
return
202197
}
203198

204199
if foundUser != nil {
205-
request.Session = foundUser.ToLocalSession(s.Config.GetServices().GetEncryption())
200+
201+
exportableSession := &models.ExportableSession{
202+
Session: foundUser,
203+
Provider: authProvider,
204+
}
205+
206+
request.Session = exportableSession.ToLocalSession(s.Config.GetServices().GetEncryption())
206207
}
207208

208209
}
@@ -323,7 +324,8 @@ func (s *Server) getElevateAuthOAuth2(c *gin.Context) {
323324
})
324325

325326
if err != nil {
326-
s.getErrorPage(c, http.StatusInternalServerError, "Failed to create session for elevation request", err)
327+
s.getErrorPage(c, http.StatusInternalServerError,
328+
"Failed to create session for elevation request", err)
327329
return
328330
}
329331

@@ -332,7 +334,13 @@ func (s *Server) getElevateAuthOAuth2(c *gin.Context) {
332334

333335
workflowTask.SetUser(session.User)
334336

335-
localSession := session.ToLocalSession(s.Config.GetServices().GetEncryption())
337+
exportableSession := &models.ExportableSession{
338+
Session: session,
339+
Provider: authProvider,
340+
}
341+
342+
localSession := exportableSession.ToLocalSession(
343+
s.Config.GetServices().GetEncryption())
336344

337345
if err := s.setAuthCookie(c, authProvider, localSession); err != nil {
338346
s.getErrorPage(c, http.StatusInternalServerError, "Failed to set auth cookie", err)

internal/daemon/middleware.go

Lines changed: 29 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -48,7 +48,11 @@ func (s *Server) AuthMiddleware() gin.HandlerFunc {
4848
}
4949

5050
// processProviderCookies extracts sessions from provider cookies
51-
func (s *Server) processProviderCookies(cookie sessions.Session, encryptionServer models.EncryptionImpl, foundSessions map[string]*models.Session) {
51+
func (s *Server) processProviderCookies(
52+
cookie sessions.Session,
53+
encryptionServer models.EncryptionImpl,
54+
foundSessions map[string]*models.Session,
55+
) {
5256
allProviders := s.Config.GetProvidersByCapability(models.ProviderCapabilityAuthorizor)
5357

5458
for providerName := range allProviders {
@@ -65,12 +69,16 @@ func (s *Server) processProviderCookies(cookie sessions.Session, encryptionServe
6569
continue
6670
}
6771

68-
foundSessions[providerName] = decodedSession
72+
foundSessions[providerName] = decodedSession.Session
6973
}
7074
}
7175

7276
// processBearerToken extracts session from Authorization Bearer token
73-
func (s *Server) processBearerToken(c *gin.Context, encryptionServer models.EncryptionImpl, foundSessions map[string]*models.Session) {
77+
func (s *Server) processBearerToken(
78+
c *gin.Context,
79+
encryptionServer models.EncryptionImpl,
80+
foundSessions map[string]*models.Session,
81+
) {
7482
authHeader := c.GetHeader("Authorization")
7583
if len(authHeader) == 0 || !strings.HasPrefix(authHeader, "Bearer ") {
7684
return
@@ -83,11 +91,20 @@ func (s *Server) processBearerToken(c *gin.Context, encryptionServer models.Encr
8391
return
8492
}
8593

86-
foundSessions["todo"] = decodedSession
94+
if len(decodedSession.Provider) == 0 {
95+
logrus.Warnln("Decoded session from bearer token has no provider information")
96+
return
97+
}
98+
99+
foundSessions[decodedSession.Provider] = decodedSession.Session
87100
}
88101

89102
// processAPIKey extracts session from X-API-Key header
90-
func (s *Server) processAPIKey(c *gin.Context, encryptionServer models.EncryptionImpl, foundSessions map[string]*models.Session) {
103+
func (s *Server) processAPIKey(
104+
c *gin.Context,
105+
encryptionServer models.EncryptionImpl,
106+
foundSessions map[string]*models.Session,
107+
) {
91108
apiHeader := c.GetHeader("X-API-Key")
92109
if len(apiHeader) == 0 {
93110
return
@@ -99,7 +116,12 @@ func (s *Server) processAPIKey(c *gin.Context, encryptionServer models.Encryptio
99116
return
100117
}
101118

102-
foundSessions["todo"] = decodedSession
119+
if len(decodedSession.Provider) == 0 {
120+
logrus.Warnln("Decoded session from API key has no provider information")
121+
return
122+
}
123+
124+
foundSessions[decodedSession.Provider] = decodedSession.Session
103125
}
104126

105127
// handleAgentMode processes sessions for agent/client mode
@@ -126,7 +148,7 @@ func (s *Server) handleAgentMode(c *gin.Context, cookie sessions.Session) {
126148
c.Redirect(http.StatusTemporaryRedirect, c.Request.RequestURI)
127149
}
128150

129-
func getDecodedSession(encryptor models.EncryptionImpl, session string) (*models.Session, error) {
151+
func getDecodedSession(encryptor models.EncryptionImpl, session string) (*models.ExportableSession, error) {
130152

131153
localSession, err := models.DecodedLocalSession(session)
132154

internal/daemon/static/user.html

Lines changed: 3 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -28,7 +28,7 @@ <h3>Account Summary</h3>
2828
</div>
2929
<div class="info-item">
3030
<span class="form-label">Provider</span>
31-
<span class="info-value">{{if .User}}{{.User.Provider}}{{else}}N/A{{end}}</span>
31+
<span class="info-value">{{if .User}}{{.User.Source}}{{else}}N/A{{end}}</span>
3232
</div>
3333
<div class="info-item">
3434
<span class="form-label">Verified</span>
@@ -295,9 +295,8 @@ <h1>User Profile</h1>
295295
async function logout() {
296296
if (confirm('Are you sure you want to logout?')) {
297297
try {
298-
const response = await fetch('/api/v1/auth/logout', {
299-
method: 'POST',
300-
credentials: 'include'
298+
const response = await fetch('{{.Config.GetApiBasePath}}/auth/logout', {
299+
method: 'GET',
301300
});
302301

303302
if (response.ok) {

internal/models/common.go

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -15,8 +15,9 @@ var ENCODED_SESSION = "session"
1515
var ENCODED_SESSION_LOCAL = "session_local"
1616

1717
type EncodingWrapper struct {
18-
Type string `json:"type"`
19-
Data any `json:"data"`
18+
Type string `json:"type"`
19+
Identifier string `json:"identifier,omitempty"`
20+
Data any `json:"data"`
2021
}
2122

2223
func (e EncodingWrapper) EncodeAndEncrypt(encryptor EncryptionImpl) string {

internal/models/session.go

Lines changed: 9 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -25,15 +25,20 @@ type Session struct {
2525
Expiry time.Time `json:"expiry"`
2626
}
2727

28+
type ExportableSession struct {
29+
*Session
30+
Provider string `json:"provider"`
31+
}
32+
2833
// Encode the remote session from the local session
29-
func (s *Session) GetEncodedSession(encryptor EncryptionImpl) string {
34+
func (s *ExportableSession) GetEncodedSession(encryptor EncryptionImpl) string {
3035
return EncodingWrapper{
3136
Type: ENCODED_SESSION,
3237
Data: s,
3338
}.EncodeAndEncrypt(encryptor)
3439
}
3540

36-
func (s *Session) ToLocalSession(encryptor EncryptionImpl) *LocalSession {
41+
func (s *ExportableSession) ToLocalSession(encryptor EncryptionImpl) *LocalSession {
3742
return &LocalSession{
3843
Version: 1,
3944
Expiry: s.Expiry,
@@ -42,7 +47,7 @@ func (s *Session) ToLocalSession(encryptor EncryptionImpl) *LocalSession {
4247
}
4348

4449
// Decode the remote session from the local session
45-
func (s *LocalSession) GetDecodedSession(decryptor EncryptionImpl) (*Session, error) {
50+
func (s *LocalSession) GetDecodedSession(decryptor EncryptionImpl) (*ExportableSession, error) {
4651
decoded, err := EncodingWrapper{}.DecodeAndDecrypt(s.Session, decryptor)
4752

4853
if err != nil {
@@ -53,7 +58,7 @@ func (s *LocalSession) GetDecodedSession(decryptor EncryptionImpl) (*Session, er
5358
return nil, fmt.Errorf("invalid session type: %s", decoded.Type)
5459
}
5560

56-
var session *Session
61+
var session *ExportableSession
5762
common.ConvertMapToInterface(decoded.Data.(map[string]any), &session)
5863

5964
return session, nil

internal/models/user.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@ type User struct {
1111
Email string `json:"email"`
1212
Name string `json:"name"`
1313
Verified *bool `json:"verified,omitempty"`
14-
Provider string `json:"provider"`
14+
Source string `json:"source,omitempty"`
1515
Groups []string `json:"groups,omitempty"`
1616
}
1717

internal/providers/github/sessions.go

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -28,10 +28,10 @@ func (p *githubProvider) CreateSession(ctx context.Context, authRequest *models.
2828
session := &models.Session{
2929
UUID: uuid.New(),
3030
User: &models.User{
31-
ID: fmt.Sprintf("%d", user.ID),
32-
Email: user.Email,
33-
Name: user.Name,
34-
Provider: ProviderName,
31+
ID: fmt.Sprintf("%d", user.ID),
32+
Email: user.Email,
33+
Name: user.Name,
34+
Source: ProviderName,
3535
},
3636
AccessToken: accessToken,
3737
Expiry: time.Now().Add(24 * time.Hour), // GitHub tokens don't expire, but we set session expiry

internal/providers/oauth2.google/main.go

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -113,7 +113,7 @@ func (p *oauth2Provider) CreateSession(ctx context.Context, authRequest *models.
113113
}
114114

115115
// Use a new context with a secure http client
116-
secureContext := context.WithValue(context.TODO(), oauth2.HTTPClient, &http.Client{
116+
secureContext := context.WithValue(ctx, oauth2.HTTPClient, &http.Client{
117117
Transport: &http.Transport{
118118
TLSClientConfig: &tls.Config{
119119
InsecureSkipVerify: false, // Ensure this is false for production
@@ -146,7 +146,7 @@ func (p *oauth2Provider) CreateSession(ctx context.Context, authRequest *models.
146146
Email: userInfo.Email,
147147
Name: userInfo.Name,
148148
Verified: userInfo.VerifiedEmail,
149-
Provider: "google",
149+
Source: "google",
150150
},
151151
AccessToken: token.AccessToken,
152152
RefreshToken: token.RefreshToken,

internal/providers/saml/main.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -133,7 +133,7 @@ func (p *samlProvider) CreateSession(ctx context.Context, authRequest *models.Au
133133
user := &models.User{
134134
Username: "saml_user", // Extract from SAML assertion
135135
Email: "user@example.com", // Extract from SAML assertion
136-
Provider: "saml",
136+
Source: "saml",
137137
Groups: []string{}, // Extract groups from SAML assertion
138138
}
139139

0 commit comments

Comments
 (0)