diff --git a/oauthproxy.go b/oauthproxy.go index 88fcdfb0..6798b3ec 100644 --- a/oauthproxy.go +++ b/oauthproxy.go @@ -804,13 +804,15 @@ func (p *OAuthProxy) backendLogout(rw http.ResponseWriter, req *http.Request, si return } - resp, err = PicsSignOutAllSessions(providerData.BackendLogoutAllSessionsURL, session.IntrospectClaims, session.AccessToken) + resp, err := PicsSignOutAllSessions(providerData.BackendLogoutAllSessionsURL, session.IntrospectClaims, session.AccessToken) if err != nil { logger.Errorf("error while calling backend logout all sessions: %v", err) return } - defer resp.Body.Close() + if resp.StatusCode() != 200 { + logger.Errorf("error while calling backend logout url, returned error code %v", resp.StatusCode()) + } } else { if providerData.BackendLogoutURL == "" { return @@ -826,10 +828,9 @@ func (p *OAuthProxy) backendLogout(rw http.ResponseWriter, req *http.Request, si } defer resp.Body.Close() - } - - if resp.StatusCode != 200 { - logger.Errorf("error while calling backend logout url, returned error code %v", resp.StatusCode) + if resp.StatusCode != 200 { + logger.Errorf("error while calling backend logout url, returned error code %v", resp.StatusCode) + } } } diff --git a/pics_oauthproxy.go b/pics_oauthproxy.go index 6bfd1c0c..9464dc22 100644 --- a/pics_oauthproxy.go +++ b/pics_oauthproxy.go @@ -4,36 +4,31 @@ import ( "encoding/base64" "encoding/json" "fmt" - "net/http" "strings" "github.com/oauth2-proxy/oauth2-proxy/v7/pkg/logger" + "github.com/oauth2-proxy/oauth2-proxy/v7/pkg/requests" ) const ( picsSignOutAllDevicesPath = "/sign_out_all_sessions" ) -func PicsSignOutAllSessions(backendLogoutAllSessionsURL string, introspectClaims string, accessToken string) (resp *http.Response, err error) { +func PicsSignOutAllSessions(backendLogoutAllSessionsURL string, introspectClaims string, accessToken string) (resp requests.Result, err error) { userID, err := getUserID(introspectClaims) if err != nil { - return nil, fmt.Errorf("error getting userId from instrospect claims: %v", err) + return nil, fmt.Errorf("error getting userID from instrospect claims: %v", err) } backendLogoutURL := strings.ReplaceAll(backendLogoutAllSessionsURL, "{user_id}", userID) + resp = requests.New(backendLogoutURL). + WithMethod("POST"). + SetHeader("Authorization", "Bearer "+accessToken). + SetHeader("API-Version", "1"). + SetHeader("Accept", "application/json"). + Do() - dummyBody := strings.NewReader(`{}`) - req, err := http.NewRequest("POST", backendLogoutURL, dummyBody) - if err != nil { - return nil, fmt.Errorf("error creating post request: %v", err) - } - - req.Header.Set("Authorization", "Bearer "+accessToken) - req.Header.Set("API-Version", "1") - req.Header.Set("Accept", "application/json") - - resp, err = http.DefaultClient.Do(req) - if err != nil { + if resp.Error() != nil { return nil, fmt.Errorf("error logging out from IAM: %v", err) } diff --git a/pics_oauthproxy_test.go b/pics_oauthproxy_test.go new file mode 100644 index 00000000..09906794 --- /dev/null +++ b/pics_oauthproxy_test.go @@ -0,0 +1,55 @@ +package main + +import ( + "encoding/base64" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/assert" +) + +func createIntrospectClaims() string { + claims := map[string]interface{}{ + "sub": "1234567890", + } + claimsBytes, err := json.Marshal(claims) + if err != nil { + return "" + } + + return base64.StdEncoding.EncodeToString(claimsBytes) +} + +func Test_PicsSignOutAllSessionsReturnsErrorWhenUserIDIsNotFound(t *testing.T) { + _, err := PicsSignOutAllSessions("http://localhost:8080/test", "", "") + + assert.Error(t, err) +} + +func Test_getUserID(t *testing.T) { + introspectClaims := createIntrospectClaims() + userID, err := getUserID(introspectClaims) + + assert.NoError(t, err) + assert.Equal(t, "1234567890", userID) +} + +func Test_PicsSignOutAllSessionsReturns200Ok(t *testing.T) { + introspectClaims := createIntrospectClaims() + accessToken := "validAccessToken" + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "Bearer "+accessToken, r.Header.Get("Authorization")) + assert.Equal(t, "1", r.Header.Get("API-Version")) + assert.Equal(t, "application/json", r.Header.Get("Accept")) + w.WriteHeader(http.StatusOK) + })) + defer server.Close() + + resp, err := PicsSignOutAllSessions(server.URL+"/{user_id}", introspectClaims, accessToken) + + assert.NoError(t, err) + assert.Equal(t, http.StatusOK, resp.StatusCode()) +}