mirror of
https://github.com/netbirdio/dex.git
synced 2026-05-22 18:43:53 -07:00
Passing context storage (#3941)
Signed-off-by: Bob Maertz <1771054+bobmaertz@users.noreply.github.com>
This commit is contained in:
+14
-14
@@ -51,7 +51,7 @@ type dexAPI struct {
|
||||
}
|
||||
|
||||
func (d dexAPI) GetClient(ctx context.Context, req *api.GetClientReq) (*api.GetClientResp, error) {
|
||||
c, err := d.s.GetClient(req.Id)
|
||||
c, err := d.s.GetClient(ctx, req.Id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -108,7 +108,7 @@ func (d dexAPI) UpdateClient(ctx context.Context, req *api.UpdateClientReq) (*ap
|
||||
return nil, errors.New("update client: no client ID supplied")
|
||||
}
|
||||
|
||||
err := d.s.UpdateClient(req.Id, func(old storage.Client) (storage.Client, error) {
|
||||
err := d.s.UpdateClient(ctx, req.Id, func(old storage.Client) (storage.Client, error) {
|
||||
if req.RedirectUris != nil {
|
||||
old.RedirectURIs = req.RedirectUris
|
||||
}
|
||||
@@ -134,7 +134,7 @@ func (d dexAPI) UpdateClient(ctx context.Context, req *api.UpdateClientReq) (*ap
|
||||
}
|
||||
|
||||
func (d dexAPI) DeleteClient(ctx context.Context, req *api.DeleteClientReq) (*api.DeleteClientResp, error) {
|
||||
err := d.s.DeleteClient(req.Id)
|
||||
err := d.s.DeleteClient(ctx, req.Id)
|
||||
if err != nil {
|
||||
if err == storage.ErrNotFound {
|
||||
return &api.DeleteClientResp{NotFound: true}, nil
|
||||
@@ -219,7 +219,7 @@ func (d dexAPI) UpdatePassword(ctx context.Context, req *api.UpdatePasswordReq)
|
||||
return old, nil
|
||||
}
|
||||
|
||||
if err := d.s.UpdatePassword(req.Email, updater); err != nil {
|
||||
if err := d.s.UpdatePassword(ctx, req.Email, updater); err != nil {
|
||||
if err == storage.ErrNotFound {
|
||||
return &api.UpdatePasswordResp{NotFound: true}, nil
|
||||
}
|
||||
@@ -235,7 +235,7 @@ func (d dexAPI) DeletePassword(ctx context.Context, req *api.DeletePasswordReq)
|
||||
return nil, errors.New("no email supplied")
|
||||
}
|
||||
|
||||
err := d.s.DeletePassword(req.Email)
|
||||
err := d.s.DeletePassword(ctx, req.Email)
|
||||
if err != nil {
|
||||
if err == storage.ErrNotFound {
|
||||
return &api.DeletePasswordResp{NotFound: true}, nil
|
||||
@@ -268,7 +268,7 @@ func (d dexAPI) GetDiscovery(ctx context.Context, req *api.DiscoveryReq) (*api.D
|
||||
}
|
||||
|
||||
func (d dexAPI) ListPasswords(ctx context.Context, req *api.ListPasswordReq) (*api.ListPasswordResp, error) {
|
||||
passwordList, err := d.s.ListPasswords()
|
||||
passwordList, err := d.s.ListPasswords(ctx)
|
||||
if err != nil {
|
||||
d.logger.Error("failed to list passwords", "err", err)
|
||||
return nil, fmt.Errorf("list passwords: %v", err)
|
||||
@@ -298,7 +298,7 @@ func (d dexAPI) VerifyPassword(ctx context.Context, req *api.VerifyPasswordReq)
|
||||
return nil, errors.New("no password to verify supplied")
|
||||
}
|
||||
|
||||
password, err := d.s.GetPassword(req.Email)
|
||||
password, err := d.s.GetPassword(ctx, req.Email)
|
||||
if err != nil {
|
||||
if err == storage.ErrNotFound {
|
||||
return &api.VerifyPasswordResp{
|
||||
@@ -327,7 +327,7 @@ func (d dexAPI) ListRefresh(ctx context.Context, req *api.ListRefreshReq) (*api.
|
||||
return nil, err
|
||||
}
|
||||
|
||||
offlineSessions, err := d.s.GetOfflineSessions(id.UserId, id.ConnId)
|
||||
offlineSessions, err := d.s.GetOfflineSessions(ctx, id.UserId, id.ConnId)
|
||||
if err != nil {
|
||||
if err == storage.ErrNotFound {
|
||||
// This means that this user-client pair does not have a refresh token yet.
|
||||
@@ -381,7 +381,7 @@ func (d dexAPI) RevokeRefresh(ctx context.Context, req *api.RevokeRefreshReq) (*
|
||||
return old, nil
|
||||
}
|
||||
|
||||
if err := d.s.UpdateOfflineSessions(id.UserId, id.ConnId, updater); err != nil {
|
||||
if err := d.s.UpdateOfflineSessions(ctx, id.UserId, id.ConnId, updater); err != nil {
|
||||
if err == storage.ErrNotFound {
|
||||
return &api.RevokeRefreshResp{NotFound: true}, nil
|
||||
}
|
||||
@@ -397,7 +397,7 @@ func (d dexAPI) RevokeRefresh(ctx context.Context, req *api.RevokeRefreshReq) (*
|
||||
//
|
||||
// TODO(ericchiang): we don't have any good recourse if this call fails.
|
||||
// Consider garbage collection of refresh tokens with no associated ref.
|
||||
if err := d.s.DeleteRefresh(refreshID); err != nil {
|
||||
if err := d.s.DeleteRefresh(ctx, refreshID); err != nil {
|
||||
d.logger.Error("failed to delete refresh token", "err", err)
|
||||
return nil, err
|
||||
}
|
||||
@@ -448,7 +448,7 @@ func (d dexAPI) CreateConnector(ctx context.Context, req *api.CreateConnectorReq
|
||||
return &api.CreateConnectorResp{}, nil
|
||||
}
|
||||
|
||||
func (d dexAPI) UpdateConnector(_ context.Context, req *api.UpdateConnectorReq) (*api.UpdateConnectorResp, error) {
|
||||
func (d dexAPI) UpdateConnector(ctx context.Context, req *api.UpdateConnectorReq) (*api.UpdateConnectorResp, error) {
|
||||
if !featureflags.APIConnectorsCRUD.Enabled() {
|
||||
return nil, fmt.Errorf("%s feature flag is not enabled", featureflags.APIConnectorsCRUD.Name)
|
||||
}
|
||||
@@ -485,7 +485,7 @@ func (d dexAPI) UpdateConnector(_ context.Context, req *api.UpdateConnectorReq)
|
||||
return old, nil
|
||||
}
|
||||
|
||||
if err := d.s.UpdateConnector(req.Id, updater); err != nil {
|
||||
if err := d.s.UpdateConnector(ctx, req.Id, updater); err != nil {
|
||||
if err == storage.ErrNotFound {
|
||||
return &api.UpdateConnectorResp{NotFound: true}, nil
|
||||
}
|
||||
@@ -505,7 +505,7 @@ func (d dexAPI) DeleteConnector(ctx context.Context, req *api.DeleteConnectorReq
|
||||
return nil, errors.New("no id supplied")
|
||||
}
|
||||
|
||||
err := d.s.DeleteConnector(req.Id)
|
||||
err := d.s.DeleteConnector(ctx, req.Id)
|
||||
if err != nil {
|
||||
if err == storage.ErrNotFound {
|
||||
return &api.DeleteConnectorResp{NotFound: true}, nil
|
||||
@@ -521,7 +521,7 @@ func (d dexAPI) ListConnectors(ctx context.Context, req *api.ListConnectorReq) (
|
||||
return nil, fmt.Errorf("%s feature flag is not enabled", featureflags.APIConnectorsCRUD.Name)
|
||||
}
|
||||
|
||||
connectorList, err := d.s.ListConnectors()
|
||||
connectorList, err := d.s.ListConnectors(ctx)
|
||||
if err != nil {
|
||||
d.logger.Error("api: failed to list connectors", "err", err)
|
||||
return nil, fmt.Errorf("list connectors: %v", err)
|
||||
|
||||
+2
-2
@@ -149,7 +149,7 @@ func TestPassword(t *testing.T) {
|
||||
t.Fatalf("Unable to update password: %v", err)
|
||||
}
|
||||
|
||||
pass, err := s.GetPassword(updateReq.Email)
|
||||
pass, err := s.GetPassword(ctx, updateReq.Email)
|
||||
if err != nil {
|
||||
t.Fatalf("Unable to retrieve password: %v", err)
|
||||
}
|
||||
@@ -449,7 +449,7 @@ func TestUpdateClient(t *testing.T) {
|
||||
t.Errorf("expected in response NotFound: %t", tc.want.NotFound)
|
||||
}
|
||||
|
||||
client, err := s.GetClient(tc.req.Id)
|
||||
client, err := s.GetClient(ctx, tc.req.Id)
|
||||
if err != nil {
|
||||
t.Errorf("no client found in the storage: %v", err)
|
||||
}
|
||||
|
||||
@@ -199,6 +199,7 @@ func (s *Server) handleDeviceTokenDeprecated(w http.ResponseWriter, r *http.Requ
|
||||
}
|
||||
|
||||
func (s *Server) handleDeviceToken(w http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
deviceCode := r.Form.Get("device_code")
|
||||
if deviceCode == "" {
|
||||
s.tokenErrHelper(w, errInvalidRequest, "No device code received", http.StatusBadRequest)
|
||||
@@ -208,7 +209,7 @@ func (s *Server) handleDeviceToken(w http.ResponseWriter, r *http.Request) {
|
||||
now := s.now()
|
||||
|
||||
// Grab the device token, check validity
|
||||
deviceToken, err := s.storage.GetDeviceToken(deviceCode)
|
||||
deviceToken, err := s.storage.GetDeviceToken(ctx, deviceCode)
|
||||
if err != nil {
|
||||
if err != storage.ErrNotFound {
|
||||
s.logger.ErrorContext(r.Context(), "failed to get device code", "err", err)
|
||||
@@ -240,7 +241,7 @@ func (s *Server) handleDeviceToken(w http.ResponseWriter, r *http.Request) {
|
||||
return old, nil
|
||||
}
|
||||
// Update device token last request time in storage
|
||||
if err := s.storage.UpdateDeviceToken(deviceCode, updater); err != nil {
|
||||
if err := s.storage.UpdateDeviceToken(ctx, deviceCode, updater); err != nil {
|
||||
s.logger.ErrorContext(r.Context(), "failed to update device token", "err", err)
|
||||
s.renderError(r, w, http.StatusInternalServerError, "")
|
||||
return
|
||||
@@ -299,7 +300,7 @@ func (s *Server) handleDeviceCallback(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
authCode, err := s.storage.GetAuthCode(code)
|
||||
authCode, err := s.storage.GetAuthCode(ctx, code)
|
||||
if err != nil || s.now().After(authCode.Expiry) {
|
||||
errCode := http.StatusBadRequest
|
||||
if err != nil && err != storage.ErrNotFound {
|
||||
@@ -311,7 +312,7 @@ func (s *Server) handleDeviceCallback(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
// Grab the device request from storage
|
||||
deviceReq, err := s.storage.GetDeviceRequest(userCode)
|
||||
deviceReq, err := s.storage.GetDeviceRequest(ctx, userCode)
|
||||
if err != nil || s.now().After(deviceReq.Expiry) {
|
||||
errCode := http.StatusBadRequest
|
||||
if err != nil && err != storage.ErrNotFound {
|
||||
@@ -322,7 +323,7 @@ func (s *Server) handleDeviceCallback(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
client, err := s.storage.GetClient(deviceReq.ClientID)
|
||||
client, err := s.storage.GetClient(ctx, deviceReq.ClientID)
|
||||
if err != nil {
|
||||
if err != storage.ErrNotFound {
|
||||
s.logger.ErrorContext(r.Context(), "failed to get client", "err", err)
|
||||
@@ -345,7 +346,7 @@ func (s *Server) handleDeviceCallback(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
// Grab the device token from storage
|
||||
old, err := s.storage.GetDeviceToken(deviceReq.DeviceCode)
|
||||
old, err := s.storage.GetDeviceToken(ctx, deviceReq.DeviceCode)
|
||||
if err != nil || s.now().After(old.Expiry) {
|
||||
errCode := http.StatusBadRequest
|
||||
if err != nil && err != storage.ErrNotFound {
|
||||
@@ -373,7 +374,7 @@ func (s *Server) handleDeviceCallback(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
// Update refresh token in the storage, store the token and mark as complete
|
||||
if err := s.storage.UpdateDeviceToken(deviceReq.DeviceCode, updater); err != nil {
|
||||
if err := s.storage.UpdateDeviceToken(ctx, deviceReq.DeviceCode, updater); err != nil {
|
||||
s.logger.ErrorContext(r.Context(), "failed to update device token", "err", err)
|
||||
s.renderError(r, w, http.StatusBadRequest, "")
|
||||
return
|
||||
@@ -391,6 +392,7 @@ func (s *Server) handleDeviceCallback(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
func (s *Server) verifyUserCode(w http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
switch r.Method {
|
||||
case http.MethodPost:
|
||||
err := r.ParseForm()
|
||||
@@ -409,7 +411,7 @@ func (s *Server) verifyUserCode(w http.ResponseWriter, r *http.Request) {
|
||||
userCode = strings.ToUpper(userCode)
|
||||
|
||||
// Find the user code in the available requests
|
||||
deviceRequest, err := s.storage.GetDeviceRequest(userCode)
|
||||
deviceRequest, err := s.storage.GetDeviceRequest(ctx, userCode)
|
||||
if err != nil || s.now().After(deviceRequest.Expiry) {
|
||||
if err != nil && err != storage.ErrNotFound {
|
||||
s.logger.ErrorContext(r.Context(), "failed to get device request", "err", err)
|
||||
|
||||
+37
-33
@@ -32,8 +32,9 @@ const (
|
||||
)
|
||||
|
||||
func (s *Server) handlePublicKeys(w http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
// TODO(ericchiang): Cache this.
|
||||
keys, err := s.storage.GetKeys()
|
||||
keys, err := s.storage.GetKeys(ctx)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(r.Context(), "failed to get keys", "err", err)
|
||||
s.renderError(r, w, http.StatusInternalServerError, "Internal server error.")
|
||||
@@ -135,6 +136,7 @@ func (s *Server) constructDiscovery() discovery {
|
||||
|
||||
// handleAuthorization handles the OAuth2 auth endpoint.
|
||||
func (s *Server) handleAuthorization(w http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
// Extract the arguments
|
||||
if err := r.ParseForm(); err != nil {
|
||||
s.logger.ErrorContext(r.Context(), "failed to parse arguments", "err", err)
|
||||
@@ -144,8 +146,7 @@ func (s *Server) handleAuthorization(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
connectorID := r.Form.Get("connector_id")
|
||||
|
||||
connectors, err := s.storage.ListConnectors()
|
||||
connectors, err := s.storage.ListConnectors(ctx)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(r.Context(), "failed to get list of connectors", "err", err)
|
||||
s.renderError(r, w, http.StatusInternalServerError, "Failed to retrieve connector list.")
|
||||
@@ -219,7 +220,7 @@ func (s *Server) handleConnectorLogin(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
conn, err := s.getConnector(connID)
|
||||
conn, err := s.getConnector(ctx, connID)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(r.Context(), "Failed to get connector", "err", err)
|
||||
s.renderError(r, w, http.StatusBadRequest, "Requested resource does not exist")
|
||||
@@ -314,6 +315,7 @@ func (s *Server) handleConnectorLogin(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
func (s *Server) handlePasswordLogin(w http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
authID := r.URL.Query().Get("state")
|
||||
if authID == "" {
|
||||
s.renderError(r, w, http.StatusBadRequest, "User session error.")
|
||||
@@ -322,7 +324,7 @@ func (s *Server) handlePasswordLogin(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
backLink := r.URL.Query().Get("back")
|
||||
|
||||
authReq, err := s.storage.GetAuthRequest(authID)
|
||||
authReq, err := s.storage.GetAuthRequest(ctx, authID)
|
||||
if err != nil {
|
||||
if err == storage.ErrNotFound {
|
||||
s.logger.ErrorContext(r.Context(), "invalid 'state' parameter provided", "err", err)
|
||||
@@ -345,7 +347,7 @@ func (s *Server) handlePasswordLogin(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
conn, err := s.getConnector(authReq.ConnectorID)
|
||||
conn, err := s.getConnector(ctx, authReq.ConnectorID)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(r.Context(), "failed to get connector", "connector_id", authReq.ConnectorID, "err", err)
|
||||
s.renderError(r, w, http.StatusInternalServerError, "Requested resource does not exist.")
|
||||
@@ -390,7 +392,7 @@ func (s *Server) handlePasswordLogin(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
if canSkipApproval {
|
||||
authReq, err = s.storage.GetAuthRequest(authReq.ID)
|
||||
authReq, err = s.storage.GetAuthRequest(ctx, authReq.ID)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(r.Context(), "failed to get finalized auth request", "err", err)
|
||||
s.renderError(r, w, http.StatusInternalServerError, "Login error.")
|
||||
@@ -425,7 +427,7 @@ func (s *Server) handleConnectorCallback(w http.ResponseWriter, r *http.Request)
|
||||
return
|
||||
}
|
||||
|
||||
authReq, err := s.storage.GetAuthRequest(authID)
|
||||
authReq, err := s.storage.GetAuthRequest(ctx, authID)
|
||||
if err != nil {
|
||||
if err == storage.ErrNotFound {
|
||||
s.logger.ErrorContext(r.Context(), "invalid 'state' parameter provided", "err", err)
|
||||
@@ -448,7 +450,7 @@ func (s *Server) handleConnectorCallback(w http.ResponseWriter, r *http.Request)
|
||||
return
|
||||
}
|
||||
|
||||
conn, err := s.getConnector(authReq.ConnectorID)
|
||||
conn, err := s.getConnector(ctx, authReq.ConnectorID)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(r.Context(), "failed to get connector", "connector_id", authReq.ConnectorID, "err", err)
|
||||
s.renderError(r, w, http.StatusInternalServerError, "Requested resource does not exist.")
|
||||
@@ -490,7 +492,7 @@ func (s *Server) handleConnectorCallback(w http.ResponseWriter, r *http.Request)
|
||||
}
|
||||
|
||||
if canSkipApproval {
|
||||
authReq, err = s.storage.GetAuthRequest(authReq.ID)
|
||||
authReq, err = s.storage.GetAuthRequest(ctx, authReq.ID)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(r.Context(), "failed to get finalized auth request", "err", err)
|
||||
s.renderError(r, w, http.StatusInternalServerError, "Login error.")
|
||||
@@ -521,7 +523,7 @@ func (s *Server) finalizeLogin(ctx context.Context, identity connector.Identity,
|
||||
a.ConnectorData = identity.ConnectorData
|
||||
return a, nil
|
||||
}
|
||||
if err := s.storage.UpdateAuthRequest(authReq.ID, updater); err != nil {
|
||||
if err := s.storage.UpdateAuthRequest(ctx, authReq.ID, updater); err != nil {
|
||||
return "", false, fmt.Errorf("failed to update auth request: %v", err)
|
||||
}
|
||||
|
||||
@@ -545,7 +547,7 @@ func (s *Server) finalizeLogin(ctx context.Context, identity connector.Identity,
|
||||
|
||||
if offlineAccessRequested && canRefresh {
|
||||
// Try to retrieve an existing OfflineSession object for the corresponding user.
|
||||
session, err := s.storage.GetOfflineSessions(identity.UserID, authReq.ConnectorID)
|
||||
session, err := s.storage.GetOfflineSessions(ctx, identity.UserID, authReq.ConnectorID)
|
||||
switch {
|
||||
case err != nil && err == storage.ErrNotFound:
|
||||
offlineSessions := storage.OfflineSessions{
|
||||
@@ -563,7 +565,7 @@ func (s *Server) finalizeLogin(ctx context.Context, identity connector.Identity,
|
||||
}
|
||||
case err == nil:
|
||||
// Update existing OfflineSession obj with new RefreshTokenRef.
|
||||
if err := s.storage.UpdateOfflineSessions(session.UserID, session.ConnID, func(old storage.OfflineSessions) (storage.OfflineSessions, error) {
|
||||
if err := s.storage.UpdateOfflineSessions(ctx, session.UserID, session.ConnID, func(old storage.OfflineSessions) (storage.OfflineSessions, error) {
|
||||
if len(identity.ConnectorData) > 0 {
|
||||
old.ConnectorData = identity.ConnectorData
|
||||
}
|
||||
@@ -594,6 +596,7 @@ func (s *Server) finalizeLogin(ctx context.Context, identity connector.Identity,
|
||||
}
|
||||
|
||||
func (s *Server) handleApproval(w http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
macEncoded := r.FormValue("hmac")
|
||||
if macEncoded == "" {
|
||||
s.renderError(r, w, http.StatusUnauthorized, "Unauthorized request")
|
||||
@@ -605,7 +608,7 @@ func (s *Server) handleApproval(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
authReq, err := s.storage.GetAuthRequest(r.FormValue("req"))
|
||||
authReq, err := s.storage.GetAuthRequest(ctx, r.FormValue("req"))
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(r.Context(), "failed to get auth request", "err", err)
|
||||
s.renderError(r, w, http.StatusInternalServerError, "Database error.")
|
||||
@@ -629,7 +632,7 @@ func (s *Server) handleApproval(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
switch r.Method {
|
||||
case http.MethodGet:
|
||||
client, err := s.storage.GetClient(authReq.ClientID)
|
||||
client, err := s.storage.GetClient(ctx, authReq.ClientID)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(r.Context(), "Failed to get client", "client_id", authReq.ClientID, "err", err)
|
||||
s.renderError(r, w, http.StatusInternalServerError, "Failed to retrieve client.")
|
||||
@@ -654,7 +657,7 @@ func (s *Server) sendCodeResponse(w http.ResponseWriter, r *http.Request, authRe
|
||||
return
|
||||
}
|
||||
|
||||
if err := s.storage.DeleteAuthRequest(authReq.ID); err != nil {
|
||||
if err := s.storage.DeleteAuthRequest(ctx, authReq.ID); err != nil {
|
||||
if err != storage.ErrNotFound {
|
||||
s.logger.ErrorContext(r.Context(), "Failed to delete authorization request", "err", err)
|
||||
s.renderError(r, w, http.StatusInternalServerError, "Internal server error.")
|
||||
@@ -786,6 +789,7 @@ func (s *Server) sendCodeResponse(w http.ResponseWriter, r *http.Request, authRe
|
||||
}
|
||||
|
||||
func (s *Server) withClientFromStorage(w http.ResponseWriter, r *http.Request, handler func(http.ResponseWriter, *http.Request, storage.Client)) {
|
||||
ctx := r.Context()
|
||||
clientID, clientSecret, ok := r.BasicAuth()
|
||||
if ok {
|
||||
var err error
|
||||
@@ -802,7 +806,7 @@ func (s *Server) withClientFromStorage(w http.ResponseWriter, r *http.Request, h
|
||||
clientSecret = r.PostFormValue("client_secret")
|
||||
}
|
||||
|
||||
client, err := s.storage.GetClient(clientID)
|
||||
client, err := s.storage.GetClient(ctx, clientID)
|
||||
if err != nil {
|
||||
if err != storage.ErrNotFound {
|
||||
s.logger.ErrorContext(r.Context(), "failed to get client", "err", err)
|
||||
@@ -885,7 +889,7 @@ func (s *Server) handleAuthCode(w http.ResponseWriter, r *http.Request, client s
|
||||
return
|
||||
}
|
||||
|
||||
authCode, err := s.storage.GetAuthCode(code)
|
||||
authCode, err := s.storage.GetAuthCode(ctx, code)
|
||||
if err != nil || s.now().After(authCode.Expiry) || authCode.ClientID != client.ID {
|
||||
if err != storage.ErrNotFound {
|
||||
s.logger.ErrorContext(r.Context(), "failed to get auth code", "err", err)
|
||||
@@ -950,7 +954,7 @@ func (s *Server) exchangeAuthCode(ctx context.Context, w http.ResponseWriter, au
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := s.storage.DeleteAuthCode(authCode.ID); err != nil {
|
||||
if err := s.storage.DeleteAuthCode(ctx, authCode.ID); err != nil {
|
||||
s.logger.ErrorContext(ctx, "failed to delete auth code", "err", err)
|
||||
s.tokenErrHelper(w, errServerError, "", http.StatusInternalServerError)
|
||||
return nil, err
|
||||
@@ -960,7 +964,7 @@ func (s *Server) exchangeAuthCode(ctx context.Context, w http.ResponseWriter, au
|
||||
// Ensure the connector supports refresh tokens.
|
||||
//
|
||||
// Connectors like `saml` do not implement RefreshConnector.
|
||||
conn, err := s.getConnector(authCode.ConnectorID)
|
||||
conn, err := s.getConnector(ctx, authCode.ConnectorID)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(ctx, "connector not found", "connector_id", authCode.ConnectorID, "err", err)
|
||||
s.tokenErrHelper(w, errServerError, "", http.StatusInternalServerError)
|
||||
@@ -1016,7 +1020,7 @@ func (s *Server) exchangeAuthCode(ctx context.Context, w http.ResponseWriter, au
|
||||
defer func() {
|
||||
if deleteToken {
|
||||
// Delete newly created refresh token from storage.
|
||||
if err := s.storage.DeleteRefresh(refresh.ID); err != nil {
|
||||
if err := s.storage.DeleteRefresh(ctx, refresh.ID); err != nil {
|
||||
s.logger.ErrorContext(ctx, "failed to delete refresh token", "err", err)
|
||||
s.tokenErrHelper(w, errServerError, "", http.StatusInternalServerError)
|
||||
return
|
||||
@@ -1032,7 +1036,7 @@ func (s *Server) exchangeAuthCode(ctx context.Context, w http.ResponseWriter, au
|
||||
}
|
||||
|
||||
// Try to retrieve an existing OfflineSession object for the corresponding user.
|
||||
if session, err := s.storage.GetOfflineSessions(refresh.Claims.UserID, refresh.ConnectorID); err != nil {
|
||||
if session, err := s.storage.GetOfflineSessions(ctx, refresh.Claims.UserID, refresh.ConnectorID); err != nil {
|
||||
if err != storage.ErrNotFound {
|
||||
s.logger.ErrorContext(ctx, "failed to get offline session", "err", err)
|
||||
s.tokenErrHelper(w, errServerError, "", http.StatusInternalServerError)
|
||||
@@ -1057,7 +1061,7 @@ func (s *Server) exchangeAuthCode(ctx context.Context, w http.ResponseWriter, au
|
||||
} else {
|
||||
if oldTokenRef, ok := session.Refresh[tokenRef.ClientID]; ok {
|
||||
// Delete old refresh token from storage.
|
||||
if err := s.storage.DeleteRefresh(oldTokenRef.ID); err != nil && err != storage.ErrNotFound {
|
||||
if err := s.storage.DeleteRefresh(ctx, oldTokenRef.ID); err != nil && err != storage.ErrNotFound {
|
||||
s.logger.ErrorContext(ctx, "failed to delete refresh token", "err", err)
|
||||
s.tokenErrHelper(w, errServerError, "", http.StatusInternalServerError)
|
||||
deleteToken = true
|
||||
@@ -1066,7 +1070,7 @@ func (s *Server) exchangeAuthCode(ctx context.Context, w http.ResponseWriter, au
|
||||
}
|
||||
|
||||
// Update existing OfflineSession obj with new RefreshTokenRef.
|
||||
if err := s.storage.UpdateOfflineSessions(session.UserID, session.ConnID, func(old storage.OfflineSessions) (storage.OfflineSessions, error) {
|
||||
if err := s.storage.UpdateOfflineSessions(ctx, session.UserID, session.ConnID, func(old storage.OfflineSessions) (storage.OfflineSessions, error) {
|
||||
old.Refresh[tokenRef.ClientID] = &tokenRef
|
||||
return old, nil
|
||||
}); err != nil {
|
||||
@@ -1140,7 +1144,7 @@ func (s *Server) handlePasswordGrant(w http.ResponseWriter, r *http.Request, cli
|
||||
continue
|
||||
}
|
||||
|
||||
isTrusted, err := s.validateCrossClientTrust(r.Context(), client.ID, peerID)
|
||||
isTrusted, err := s.validateCrossClientTrust(ctx, client.ID, peerID)
|
||||
if err != nil {
|
||||
s.tokenErrHelper(w, errInvalidClient, fmt.Sprintf("Error validating cross client trust %v.", err), http.StatusBadRequest)
|
||||
return
|
||||
@@ -1165,7 +1169,7 @@ func (s *Server) handlePasswordGrant(w http.ResponseWriter, r *http.Request, cli
|
||||
|
||||
// Which connector
|
||||
connID := s.passwordConnector
|
||||
conn, err := s.getConnector(connID)
|
||||
conn, err := s.getConnector(ctx, connID)
|
||||
if err != nil {
|
||||
s.tokenErrHelper(w, errInvalidRequest, "Requested connector does not exist.", http.StatusBadRequest)
|
||||
return
|
||||
@@ -1201,14 +1205,14 @@ func (s *Server) handlePasswordGrant(w http.ResponseWriter, r *http.Request, cli
|
||||
Groups: identity.Groups,
|
||||
}
|
||||
|
||||
accessToken, _, err := s.newAccessToken(r.Context(), client.ID, claims, scopes, nonce, connID)
|
||||
accessToken, _, err := s.newAccessToken(ctx, client.ID, claims, scopes, nonce, connID)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(r.Context(), "password grant failed to create new access token", "err", err)
|
||||
s.tokenErrHelper(w, errServerError, "", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
idToken, expiry, err := s.newIDToken(r.Context(), client.ID, claims, scopes, nonce, accessToken, "", connID)
|
||||
idToken, expiry, err := s.newIDToken(ctx, client.ID, claims, scopes, nonce, accessToken, "", connID)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(r.Context(), "password grant failed to create new ID token", "err", err)
|
||||
s.tokenErrHelper(w, errServerError, "", http.StatusInternalServerError)
|
||||
@@ -1268,7 +1272,7 @@ func (s *Server) handlePasswordGrant(w http.ResponseWriter, r *http.Request, cli
|
||||
defer func() {
|
||||
if deleteToken {
|
||||
// Delete newly created refresh token from storage.
|
||||
if err := s.storage.DeleteRefresh(refresh.ID); err != nil {
|
||||
if err := s.storage.DeleteRefresh(ctx, refresh.ID); err != nil {
|
||||
s.logger.ErrorContext(r.Context(), "failed to delete refresh token", "err", err)
|
||||
s.tokenErrHelper(w, errServerError, "", http.StatusInternalServerError)
|
||||
return
|
||||
@@ -1284,7 +1288,7 @@ func (s *Server) handlePasswordGrant(w http.ResponseWriter, r *http.Request, cli
|
||||
}
|
||||
|
||||
// Try to retrieve an existing OfflineSession object for the corresponding user.
|
||||
if session, err := s.storage.GetOfflineSessions(refresh.Claims.UserID, refresh.ConnectorID); err != nil {
|
||||
if session, err := s.storage.GetOfflineSessions(ctx, refresh.Claims.UserID, refresh.ConnectorID); err != nil {
|
||||
if err != storage.ErrNotFound {
|
||||
s.logger.ErrorContext(r.Context(), "failed to get offline session", "err", err)
|
||||
s.tokenErrHelper(w, errServerError, "", http.StatusInternalServerError)
|
||||
@@ -1310,7 +1314,7 @@ func (s *Server) handlePasswordGrant(w http.ResponseWriter, r *http.Request, cli
|
||||
} else {
|
||||
if oldTokenRef, ok := session.Refresh[tokenRef.ClientID]; ok {
|
||||
// Delete old refresh token from storage.
|
||||
if err := s.storage.DeleteRefresh(oldTokenRef.ID); err != nil {
|
||||
if err := s.storage.DeleteRefresh(ctx, oldTokenRef.ID); err != nil {
|
||||
if err == storage.ErrNotFound {
|
||||
s.logger.Warn("database inconsistent, refresh token missing", "token_id", oldTokenRef.ID)
|
||||
} else {
|
||||
@@ -1323,7 +1327,7 @@ func (s *Server) handlePasswordGrant(w http.ResponseWriter, r *http.Request, cli
|
||||
}
|
||||
|
||||
// Update existing OfflineSession obj with new RefreshTokenRef.
|
||||
if err := s.storage.UpdateOfflineSessions(session.UserID, session.ConnID, func(old storage.OfflineSessions) (storage.OfflineSessions, error) {
|
||||
if err := s.storage.UpdateOfflineSessions(ctx, session.UserID, session.ConnID, func(old storage.OfflineSessions) (storage.OfflineSessions, error) {
|
||||
old.Refresh[tokenRef.ClientID] = &tokenRef
|
||||
old.ConnectorData = identity.ConnectorData
|
||||
return old, nil
|
||||
@@ -1371,7 +1375,7 @@ func (s *Server) handleTokenExchange(w http.ResponseWriter, r *http.Request, cli
|
||||
return
|
||||
}
|
||||
|
||||
conn, err := s.getConnector(connID)
|
||||
conn, err := s.getConnector(ctx, connID)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(r.Context(), "failed to get connector", "err", err)
|
||||
s.tokenErrHelper(w, errInvalidRequest, "Requested connector does not exist.", http.StatusBadRequest)
|
||||
|
||||
@@ -138,7 +138,7 @@ type emptyStorage struct {
|
||||
storage.Storage
|
||||
}
|
||||
|
||||
func (*emptyStorage) GetAuthRequest(string) (storage.AuthRequest, error) {
|
||||
func (*emptyStorage) GetAuthRequest(context.Context, string) (storage.AuthRequest, error) {
|
||||
return storage.AuthRequest{}, storage.ErrNotFound
|
||||
}
|
||||
|
||||
@@ -407,7 +407,7 @@ func TestHandlePassword(t *testing.T) {
|
||||
err := json.Unmarshal(rr.Body.Bytes(), &ref)
|
||||
require.NoError(t, err)
|
||||
|
||||
newSess, err := s.storage.GetOfflineSessions("0-385-28089-0", "test")
|
||||
newSess, err := s.storage.GetOfflineSessions(ctx, "0-385-28089-0", "test")
|
||||
if tc.offlineSessionCreated {
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, `{"test": "true"}`, string(newSess.ConnectorData))
|
||||
@@ -562,7 +562,7 @@ func TestHandlePasswordLoginWithSkipApproval(t *testing.T) {
|
||||
cb, _ := url.Parse(resp.Header.Get("Location"))
|
||||
require.Equal(t, tc.expectedRes, cb.Path)
|
||||
|
||||
offlineSession, err := s.storage.GetOfflineSessions("0-385-28089-0", connID)
|
||||
offlineSession, err := s.storage.GetOfflineSessions(ctx, "0-385-28089-0", connID)
|
||||
if tc.offlineSessionCreated {
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, offlineSession)
|
||||
@@ -701,7 +701,7 @@ func TestHandleConnectorCallbackWithSkipApproval(t *testing.T) {
|
||||
cb, _ := url.Parse(resp.Header.Get("Location"))
|
||||
require.Equal(t, tc.expectedRes, cb.Path)
|
||||
|
||||
offlineSession, err := s.storage.GetOfflineSessions("0-385-28089-0", connID)
|
||||
offlineSession, err := s.storage.GetOfflineSessions(ctx, "0-385-28089-0", connID)
|
||||
if tc.offlineSessionCreated {
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, offlineSession)
|
||||
|
||||
@@ -263,7 +263,7 @@ func (s *Server) introspectAccessToken(ctx context.Context, token string) (*Intr
|
||||
return nil, newIntrospectInternalServerError()
|
||||
}
|
||||
|
||||
client, err := s.storage.GetClient(clientID)
|
||||
client, err := s.storage.GetClient(ctx, clientID)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(ctx, "error while fetching client from storage", "err", err.Error())
|
||||
return nil, newIntrospectInternalServerError()
|
||||
|
||||
+7
-6
@@ -351,7 +351,7 @@ func genSubject(userID string, connID string) (string, error) {
|
||||
}
|
||||
|
||||
func (s *Server) newIDToken(ctx context.Context, clientID string, claims storage.Claims, scopes []string, nonce, accessToken, code, connID string) (idToken string, expiry time.Time, err error) {
|
||||
keys, err := s.storage.GetKeys()
|
||||
keys, err := s.storage.GetKeys(ctx)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(ctx, "failed to get keys", "err", err)
|
||||
return "", expiry, err
|
||||
@@ -453,6 +453,7 @@ func (s *Server) newIDToken(ctx context.Context, clientID string, claims storage
|
||||
|
||||
// parse the initial request from the OAuth2 client.
|
||||
func (s *Server) parseAuthorizationRequest(r *http.Request) (*storage.AuthRequest, error) {
|
||||
ctx := r.Context()
|
||||
if err := r.ParseForm(); err != nil {
|
||||
return nil, newDisplayedErr(http.StatusBadRequest, "Failed to parse request.")
|
||||
}
|
||||
@@ -477,7 +478,7 @@ func (s *Server) parseAuthorizationRequest(r *http.Request) (*storage.AuthReques
|
||||
codeChallengeMethod = codeChallengeMethodPlain
|
||||
}
|
||||
|
||||
client, err := s.storage.GetClient(clientID)
|
||||
client, err := s.storage.GetClient(ctx, clientID)
|
||||
if err != nil {
|
||||
if err == storage.ErrNotFound {
|
||||
return nil, newDisplayedErr(http.StatusNotFound, "Invalid client_id (%q).", clientID)
|
||||
@@ -499,7 +500,7 @@ func (s *Server) parseAuthorizationRequest(r *http.Request) (*storage.AuthReques
|
||||
}
|
||||
|
||||
if connectorID != "" {
|
||||
connectors, err := s.storage.ListConnectors()
|
||||
connectors, err := s.storage.ListConnectors(ctx)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(r.Context(), "failed to list connectors", "err", err)
|
||||
return nil, newRedirectedErr(errServerError, "Unable to retrieve connectors")
|
||||
@@ -634,7 +635,7 @@ func (s *Server) validateCrossClientTrust(ctx context.Context, clientID, peerID
|
||||
if peerID == clientID {
|
||||
return true, nil
|
||||
}
|
||||
peer, err := s.storage.GetClient(peerID)
|
||||
peer, err := s.storage.GetClient(ctx, peerID)
|
||||
if err != nil {
|
||||
if err != storage.ErrNotFound {
|
||||
s.logger.ErrorContext(ctx, "failed to get client", "err", err)
|
||||
@@ -707,7 +708,7 @@ type storageKeySet struct {
|
||||
storage.Storage
|
||||
}
|
||||
|
||||
func (s *storageKeySet) VerifySignature(_ context.Context, jwt string) (payload []byte, err error) {
|
||||
func (s *storageKeySet) VerifySignature(ctx context.Context, jwt string) (payload []byte, err error) {
|
||||
jws, err := jose.ParseSigned(jwt, []jose.SignatureAlgorithm{jose.RS256, jose.RS384, jose.RS512, jose.ES256, jose.ES384, jose.ES512})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -719,7 +720,7 @@ func (s *storageKeySet) VerifySignature(_ context.Context, jwt string) (payload
|
||||
break
|
||||
}
|
||||
|
||||
skeys, err := s.Storage.GetKeys()
|
||||
skeys, err := s.Storage.GetKeys(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -599,7 +599,7 @@ func TestValidRedirectURI(t *testing.T) {
|
||||
|
||||
func TestStorageKeySet(t *testing.T) {
|
||||
s := memory.New(logger)
|
||||
if err := s.UpdateKeys(func(keys storage.Keys) (storage.Keys, error) {
|
||||
if err := s.UpdateKeys(context.TODO(), func(keys storage.Keys) (storage.Keys, error) {
|
||||
keys.SigningKey = &jose.JSONWebKey{
|
||||
Key: testKey,
|
||||
KeyID: "testkey",
|
||||
|
||||
@@ -84,7 +84,7 @@ func (s *Server) getRefreshTokenFromStorage(ctx context.Context, clientID *strin
|
||||
refreshCtx := refreshContext{requestToken: token}
|
||||
|
||||
// Get RefreshToken
|
||||
refresh, err := s.storage.GetRefresh(token.RefreshId)
|
||||
refresh, err := s.storage.GetRefresh(ctx, token.RefreshId)
|
||||
if err != nil {
|
||||
if err != storage.ErrNotFound {
|
||||
s.logger.ErrorContext(ctx, "failed to get refresh token", "err", err)
|
||||
@@ -126,14 +126,14 @@ func (s *Server) getRefreshTokenFromStorage(ctx context.Context, clientID *strin
|
||||
refreshCtx.storageToken = &refresh
|
||||
|
||||
// Get Connector
|
||||
refreshCtx.connector, err = s.getConnector(refresh.ConnectorID)
|
||||
refreshCtx.connector, err = s.getConnector(ctx, refresh.ConnectorID)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(ctx, "connector not found", "connector_id", refresh.ConnectorID, "err", err)
|
||||
return nil, newInternalServerError()
|
||||
}
|
||||
|
||||
// Get Connector Data
|
||||
session, err := s.storage.GetOfflineSessions(refresh.Claims.UserID, refresh.ConnectorID)
|
||||
session, err := s.storage.GetOfflineSessions(ctx, refresh.Claims.UserID, refresh.ConnectorID)
|
||||
switch {
|
||||
case err != nil:
|
||||
if err != storage.ErrNotFound {
|
||||
@@ -223,7 +223,7 @@ func (s *Server) updateOfflineSession(ctx context.Context, refresh *storage.Refr
|
||||
|
||||
// Update LastUsed time stamp in refresh token reference object
|
||||
// in offline session for the user.
|
||||
err := s.storage.UpdateOfflineSessions(refresh.Claims.UserID, refresh.ConnectorID, offlineSessionUpdater)
|
||||
err := s.storage.UpdateOfflineSessions(ctx, refresh.Claims.UserID, refresh.ConnectorID, offlineSessionUpdater)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(ctx, "failed to update offline session", "err", err)
|
||||
return newInternalServerError()
|
||||
@@ -314,7 +314,7 @@ func (s *Server) updateRefreshToken(ctx context.Context, rCtx *refreshContext) (
|
||||
}
|
||||
|
||||
// Update refresh token in the storage.
|
||||
err := s.storage.UpdateRefreshToken(rCtx.storageToken.ID, refreshTokenUpdater)
|
||||
err := s.storage.UpdateRefreshToken(ctx, rCtx.storageToken.ID, refreshTokenUpdater)
|
||||
if err != nil {
|
||||
s.logger.ErrorContext(ctx, "failed to update refresh token", "err", err)
|
||||
return nil, ident, newInternalServerError()
|
||||
|
||||
+2
-2
@@ -95,7 +95,7 @@ func (s *Server) startKeyRotation(ctx context.Context, strategy rotationStrategy
|
||||
}
|
||||
|
||||
func (k keyRotator) rotate() error {
|
||||
keys, err := k.GetKeys()
|
||||
keys, err := k.GetKeys(context.Background())
|
||||
if err != nil && err != storage.ErrNotFound {
|
||||
return fmt.Errorf("get keys: %v", err)
|
||||
}
|
||||
@@ -128,7 +128,7 @@ func (k keyRotator) rotate() error {
|
||||
}
|
||||
|
||||
var nextRotation time.Time
|
||||
err = k.Storage.UpdateKeys(func(keys storage.Keys) (storage.Keys, error) {
|
||||
err = k.Storage.UpdateKeys(context.Background(), func(keys storage.Keys) (storage.Keys, error) {
|
||||
tNow := k.now()
|
||||
|
||||
// if you are running multiple instances of dex, another instance
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"log/slog"
|
||||
"sort"
|
||||
@@ -14,7 +15,7 @@ import (
|
||||
)
|
||||
|
||||
func signingKeyID(t *testing.T, s storage.Storage) string {
|
||||
keys, err := s.GetKeys()
|
||||
keys, err := s.GetKeys(context.TODO())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -22,7 +23,7 @@ func signingKeyID(t *testing.T, s storage.Storage) string {
|
||||
}
|
||||
|
||||
func verificationKeyIDs(t *testing.T, s storage.Storage) (ids []string) {
|
||||
keys, err := s.GetKeys()
|
||||
keys, err := s.GetKeys(context.TODO())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
+8
-8
@@ -316,7 +316,7 @@ func newServer(ctx context.Context, c Config, rotationStrategy rotationStrategy)
|
||||
|
||||
// Retrieves connector objects in backend storage. This list includes the static connectors
|
||||
// defined in the ConfigMap and dynamic connectors retrieved from the storage.
|
||||
storageConnectors, err := c.Storage.ListConnectors()
|
||||
storageConnectors, err := c.Storage.ListConnectors(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("server: failed to list connector objects from storage: %v", err)
|
||||
}
|
||||
@@ -535,7 +535,7 @@ type passwordDB struct {
|
||||
}
|
||||
|
||||
func (db passwordDB) Login(ctx context.Context, s connector.Scopes, email, password string) (connector.Identity, bool, error) {
|
||||
p, err := db.s.GetPassword(email)
|
||||
p, err := db.s.GetPassword(ctx, email)
|
||||
if err != nil {
|
||||
if err != storage.ErrNotFound {
|
||||
return connector.Identity{}, false, fmt.Errorf("get password: %v", err)
|
||||
@@ -560,7 +560,7 @@ func (db passwordDB) Login(ctx context.Context, s connector.Scopes, email, passw
|
||||
|
||||
func (db passwordDB) Refresh(ctx context.Context, s connector.Scopes, identity connector.Identity) (connector.Identity, error) {
|
||||
// If the user has been deleted, the refresh token will be rejected.
|
||||
p, err := db.s.GetPassword(identity.Email)
|
||||
p, err := db.s.GetPassword(ctx, identity.Email)
|
||||
if err != nil {
|
||||
if err == storage.ErrNotFound {
|
||||
return connector.Identity{}, errors.New("user not found")
|
||||
@@ -602,13 +602,13 @@ type keyCacher struct {
|
||||
keys atomic.Value // Always holds nil or type *storage.Keys.
|
||||
}
|
||||
|
||||
func (k *keyCacher) GetKeys() (storage.Keys, error) {
|
||||
func (k *keyCacher) GetKeys(ctx context.Context) (storage.Keys, error) {
|
||||
keys, ok := k.keys.Load().(*storage.Keys)
|
||||
if ok && keys != nil && k.now().Before(keys.NextRotation) {
|
||||
return *keys, nil
|
||||
}
|
||||
|
||||
storageKeys, err := k.Storage.GetKeys()
|
||||
storageKeys, err := k.Storage.GetKeys(ctx)
|
||||
if err != nil {
|
||||
return storageKeys, err
|
||||
}
|
||||
@@ -626,7 +626,7 @@ func (s *Server) startGarbageCollection(ctx context.Context, frequency time.Dura
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-time.After(frequency):
|
||||
if r, err := s.storage.GarbageCollect(now()); err != nil {
|
||||
if r, err := s.storage.GarbageCollect(ctx, now()); err != nil {
|
||||
s.logger.ErrorContext(ctx, "garbage collection failed", "err", err)
|
||||
} else if !r.IsEmpty() {
|
||||
s.logger.InfoContext(ctx, "garbage collection run, delete auth",
|
||||
@@ -719,8 +719,8 @@ func (s *Server) OpenConnector(conn storage.Connector) (Connector, error) {
|
||||
|
||||
// getConnector retrieves the connector object with the given id from the storage
|
||||
// and updates the connector list for server if necessary.
|
||||
func (s *Server) getConnector(id string) (Connector, error) {
|
||||
storageConnector, err := s.storage.GetConnector(id)
|
||||
func (s *Server) getConnector(ctx context.Context, id string) (Connector, error) {
|
||||
storageConnector, err := s.storage.GetConnector(ctx, id)
|
||||
if err != nil {
|
||||
return Connector{}, fmt.Errorf("failed to get connector object from storage: %v", err)
|
||||
}
|
||||
|
||||
@@ -875,7 +875,7 @@ func TestOAuth2CodeFlow(t *testing.T) {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
tokens, err := s.storage.ListRefreshTokens()
|
||||
tokens, err := s.storage.ListRefreshTokens(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to get existed refresh token: %v", err)
|
||||
}
|
||||
@@ -1369,15 +1369,15 @@ type storageWithKeysTrigger struct {
|
||||
f func()
|
||||
}
|
||||
|
||||
func (s storageWithKeysTrigger) GetKeys() (storage.Keys, error) {
|
||||
func (s storageWithKeysTrigger) GetKeys(ctx context.Context) (storage.Keys, error) {
|
||||
s.f()
|
||||
return s.Storage.GetKeys()
|
||||
return s.Storage.GetKeys(ctx)
|
||||
}
|
||||
|
||||
func TestKeyCacher(t *testing.T) {
|
||||
tNow := time.Now()
|
||||
now := func() time.Time { return tNow }
|
||||
|
||||
ctx := context.TODO()
|
||||
s := memory.New(logger)
|
||||
|
||||
tests := []struct {
|
||||
@@ -1390,7 +1390,7 @@ func TestKeyCacher(t *testing.T) {
|
||||
},
|
||||
{
|
||||
before: func() {
|
||||
s.UpdateKeys(func(old storage.Keys) (storage.Keys, error) {
|
||||
s.UpdateKeys(ctx, func(old storage.Keys) (storage.Keys, error) {
|
||||
old.NextRotation = tNow.Add(time.Minute)
|
||||
return old, nil
|
||||
})
|
||||
@@ -1410,7 +1410,7 @@ func TestKeyCacher(t *testing.T) {
|
||||
{
|
||||
before: func() {
|
||||
tNow = tNow.Add(time.Hour)
|
||||
s.UpdateKeys(func(old storage.Keys) (storage.Keys, error) {
|
||||
s.UpdateKeys(ctx, func(old storage.Keys) (storage.Keys, error) {
|
||||
old.NextRotation = tNow.Add(time.Minute)
|
||||
return old, nil
|
||||
})
|
||||
@@ -1428,7 +1428,7 @@ func TestKeyCacher(t *testing.T) {
|
||||
for i, tc := range tests {
|
||||
gotCall = false
|
||||
tc.before()
|
||||
s.GetKeys()
|
||||
s.GetKeys(context.TODO())
|
||||
if gotCall != tc.wantCallToStorage {
|
||||
t.Errorf("case %d: expected call to storage=%t got call to storage=%t", i, tc.wantCallToStorage, gotCall)
|
||||
}
|
||||
|
||||
@@ -148,7 +148,7 @@ func testAuthRequestCRUD(t *testing.T, s storage.Storage) {
|
||||
t.Fatalf("failed creating auth request: %v", err)
|
||||
}
|
||||
|
||||
if err := s.UpdateAuthRequest(a1.ID, func(old storage.AuthRequest) (storage.AuthRequest, error) {
|
||||
if err := s.UpdateAuthRequest(ctx, a1.ID, func(old storage.AuthRequest) (storage.AuthRequest, error) {
|
||||
old.Claims = identity
|
||||
old.ConnectorID = "connID"
|
||||
return old, nil
|
||||
@@ -156,7 +156,7 @@ func testAuthRequestCRUD(t *testing.T, s storage.Storage) {
|
||||
t.Fatalf("failed to update auth request: %v", err)
|
||||
}
|
||||
|
||||
got, err := s.GetAuthRequest(a1.ID)
|
||||
got, err := s.GetAuthRequest(ctx, a1.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to get auth req: %v", err)
|
||||
}
|
||||
@@ -168,15 +168,15 @@ func testAuthRequestCRUD(t *testing.T, s storage.Storage) {
|
||||
t.Fatalf("storage does not support PKCE, wanted challenge=%#v got %#v", codeChallenge, got.PKCE)
|
||||
}
|
||||
|
||||
if err := s.DeleteAuthRequest(a1.ID); err != nil {
|
||||
if err := s.DeleteAuthRequest(ctx, a1.ID); err != nil {
|
||||
t.Fatalf("failed to delete auth request: %v", err)
|
||||
}
|
||||
|
||||
if err := s.DeleteAuthRequest(a2.ID); err != nil {
|
||||
if err := s.DeleteAuthRequest(ctx, a2.ID); err != nil {
|
||||
t.Fatalf("failed to delete auth request: %v", err)
|
||||
}
|
||||
|
||||
_, err = s.GetAuthRequest(a1.ID)
|
||||
_, err = s.GetAuthRequest(ctx, a1.ID)
|
||||
mustBeErrNotFound(t, "auth request", err)
|
||||
}
|
||||
|
||||
@@ -234,7 +234,7 @@ func testAuthCodeCRUD(t *testing.T, s storage.Storage) {
|
||||
t.Fatalf("failed creating auth code: %v", err)
|
||||
}
|
||||
|
||||
got, err := s.GetAuthCode(a1.ID)
|
||||
got, err := s.GetAuthCode(ctx, a1.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to get auth code: %v", err)
|
||||
}
|
||||
@@ -246,15 +246,15 @@ func testAuthCodeCRUD(t *testing.T, s storage.Storage) {
|
||||
t.Errorf("auth code retrieved from storage did not match: %s", diff)
|
||||
}
|
||||
|
||||
if err := s.DeleteAuthCode(a1.ID); err != nil {
|
||||
if err := s.DeleteAuthCode(ctx, a1.ID); err != nil {
|
||||
t.Fatalf("delete auth code: %v", err)
|
||||
}
|
||||
|
||||
if err := s.DeleteAuthCode(a2.ID); err != nil {
|
||||
if err := s.DeleteAuthCode(ctx, a2.ID); err != nil {
|
||||
t.Fatalf("delete auth code: %v", err)
|
||||
}
|
||||
|
||||
_, err = s.GetAuthCode(a1.ID)
|
||||
_, err = s.GetAuthCode(ctx, a1.ID)
|
||||
mustBeErrNotFound(t, "auth code", err)
|
||||
}
|
||||
|
||||
@@ -268,7 +268,7 @@ func testClientCRUD(t *testing.T, s storage.Storage) {
|
||||
Name: "dex client",
|
||||
LogoURL: "https://goo.gl/JIyzIC",
|
||||
}
|
||||
err := s.DeleteClient(id1)
|
||||
err := s.DeleteClient(ctx, id1)
|
||||
mustBeErrNotFound(t, "client", err)
|
||||
|
||||
if err := s.CreateClient(ctx, c1); err != nil {
|
||||
@@ -293,7 +293,7 @@ func testClientCRUD(t *testing.T, s storage.Storage) {
|
||||
}
|
||||
|
||||
getAndCompare := func(_ string, want storage.Client) {
|
||||
gc, err := s.GetClient(id1)
|
||||
gc, err := s.GetClient(ctx, id1)
|
||||
if err != nil {
|
||||
t.Errorf("get client: %v", err)
|
||||
return
|
||||
@@ -306,7 +306,7 @@ func testClientCRUD(t *testing.T, s storage.Storage) {
|
||||
getAndCompare(id1, c1)
|
||||
|
||||
newSecret := "barfoo"
|
||||
err = s.UpdateClient(id1, func(old storage.Client) (storage.Client, error) {
|
||||
err = s.UpdateClient(ctx, id1, func(old storage.Client) (storage.Client, error) {
|
||||
old.Secret = newSecret
|
||||
return old, nil
|
||||
})
|
||||
@@ -316,15 +316,15 @@ func testClientCRUD(t *testing.T, s storage.Storage) {
|
||||
c1.Secret = newSecret
|
||||
getAndCompare(id1, c1)
|
||||
|
||||
if err := s.DeleteClient(id1); err != nil {
|
||||
if err := s.DeleteClient(ctx, id1); err != nil {
|
||||
t.Fatalf("delete client: %v", err)
|
||||
}
|
||||
|
||||
if err := s.DeleteClient(id2); err != nil {
|
||||
if err := s.DeleteClient(ctx, id2); err != nil {
|
||||
t.Fatalf("delete client: %v", err)
|
||||
}
|
||||
|
||||
_, err = s.GetClient(id1)
|
||||
_, err = s.GetClient(ctx, id1)
|
||||
mustBeErrNotFound(t, "client", err)
|
||||
}
|
||||
|
||||
@@ -359,7 +359,7 @@ func testRefreshTokenCRUD(t *testing.T, s storage.Storage) {
|
||||
mustBeErrAlreadyExists(t, "refresh token", err)
|
||||
|
||||
getAndCompare := func(id string, want storage.RefreshToken) {
|
||||
gr, err := s.GetRefresh(id)
|
||||
gr, err := s.GetRefresh(ctx, id)
|
||||
if err != nil {
|
||||
t.Errorf("get refresh: %v", err)
|
||||
return
|
||||
@@ -419,7 +419,7 @@ func testRefreshTokenCRUD(t *testing.T, s storage.Storage) {
|
||||
r.LastUsed = updatedAt
|
||||
return r, nil
|
||||
}
|
||||
if err := s.UpdateRefreshToken(id, updater); err != nil {
|
||||
if err := s.UpdateRefreshToken(ctx, id, updater); err != nil {
|
||||
t.Errorf("failed to update refresh token: %v", err)
|
||||
}
|
||||
refresh.Token = "spam"
|
||||
@@ -429,15 +429,15 @@ func testRefreshTokenCRUD(t *testing.T, s storage.Storage) {
|
||||
// Ensure that updating the first token doesn't impact the second. Issue #847.
|
||||
getAndCompare(id2, refresh2)
|
||||
|
||||
if err := s.DeleteRefresh(id); err != nil {
|
||||
if err := s.DeleteRefresh(ctx, id); err != nil {
|
||||
t.Fatalf("failed to delete refresh request: %v", err)
|
||||
}
|
||||
|
||||
if err := s.DeleteRefresh(id2); err != nil {
|
||||
if err := s.DeleteRefresh(ctx, id2); err != nil {
|
||||
t.Fatalf("failed to delete refresh request: %v", err)
|
||||
}
|
||||
|
||||
_, err = s.GetRefresh(id)
|
||||
_, err = s.GetRefresh(ctx, id)
|
||||
mustBeErrNotFound(t, "refresh token", err)
|
||||
}
|
||||
|
||||
@@ -485,7 +485,7 @@ func testPasswordCRUD(t *testing.T, s storage.Storage) {
|
||||
}
|
||||
|
||||
getAndCompare := func(id string, want storage.Password) {
|
||||
gr, err := s.GetPassword(id)
|
||||
gr, err := s.GetPassword(ctx, id)
|
||||
if err != nil {
|
||||
t.Errorf("get password %q: %v", id, err)
|
||||
return
|
||||
@@ -498,7 +498,7 @@ func testPasswordCRUD(t *testing.T, s storage.Storage) {
|
||||
getAndCompare("jane@example.com", password1)
|
||||
getAndCompare("JANE@example.com", password1) // Emails should be case insensitive
|
||||
|
||||
if err := s.UpdatePassword(password1.Email, func(old storage.Password) (storage.Password, error) {
|
||||
if err := s.UpdatePassword(ctx, password1.Email, func(old storage.Password) (storage.Password, error) {
|
||||
old.Username = "jane doe"
|
||||
return old, nil
|
||||
}); err != nil {
|
||||
@@ -512,7 +512,7 @@ func testPasswordCRUD(t *testing.T, s storage.Storage) {
|
||||
passwordList = append(passwordList, password1, password2)
|
||||
|
||||
listAndCompare := func(want []storage.Password) {
|
||||
passwords, err := s.ListPasswords()
|
||||
passwords, err := s.ListPasswords(ctx)
|
||||
if err != nil {
|
||||
t.Errorf("list password: %v", err)
|
||||
return
|
||||
@@ -526,15 +526,15 @@ func testPasswordCRUD(t *testing.T, s storage.Storage) {
|
||||
|
||||
listAndCompare(passwordList)
|
||||
|
||||
if err := s.DeletePassword(password1.Email); err != nil {
|
||||
if err := s.DeletePassword(ctx, password1.Email); err != nil {
|
||||
t.Fatalf("failed to delete password: %v", err)
|
||||
}
|
||||
|
||||
if err := s.DeletePassword(password2.Email); err != nil {
|
||||
if err := s.DeletePassword(ctx, password2.Email); err != nil {
|
||||
t.Fatalf("failed to delete password: %v", err)
|
||||
}
|
||||
|
||||
_, err = s.GetPassword(password1.Email)
|
||||
_, err = s.GetPassword(ctx, password1.Email)
|
||||
mustBeErrNotFound(t, "password", err)
|
||||
}
|
||||
|
||||
@@ -571,7 +571,7 @@ func testOfflineSessionCRUD(t *testing.T, s storage.Storage) {
|
||||
}
|
||||
|
||||
getAndCompare := func(userID string, connID string, want storage.OfflineSessions) {
|
||||
gr, err := s.GetOfflineSessions(userID, connID)
|
||||
gr, err := s.GetOfflineSessions(ctx, userID, connID)
|
||||
if err != nil {
|
||||
t.Errorf("get offline session: %v", err)
|
||||
return
|
||||
@@ -592,7 +592,7 @@ func testOfflineSessionCRUD(t *testing.T, s storage.Storage) {
|
||||
}
|
||||
session1.Refresh[tokenRef.ClientID] = &tokenRef
|
||||
|
||||
if err := s.UpdateOfflineSessions(session1.UserID, session1.ConnID, func(old storage.OfflineSessions) (storage.OfflineSessions, error) {
|
||||
if err := s.UpdateOfflineSessions(ctx, session1.UserID, session1.ConnID, func(old storage.OfflineSessions) (storage.OfflineSessions, error) {
|
||||
old.Refresh[tokenRef.ClientID] = &tokenRef
|
||||
return old, nil
|
||||
}); err != nil {
|
||||
@@ -601,15 +601,15 @@ func testOfflineSessionCRUD(t *testing.T, s storage.Storage) {
|
||||
|
||||
getAndCompare(userID1, "Conn1", session1)
|
||||
|
||||
if err := s.DeleteOfflineSessions(session1.UserID, session1.ConnID); err != nil {
|
||||
if err := s.DeleteOfflineSessions(ctx, session1.UserID, session1.ConnID); err != nil {
|
||||
t.Fatalf("failed to delete offline session: %v", err)
|
||||
}
|
||||
|
||||
if err := s.DeleteOfflineSessions(session2.UserID, session2.ConnID); err != nil {
|
||||
if err := s.DeleteOfflineSessions(ctx, session2.UserID, session2.ConnID); err != nil {
|
||||
t.Fatalf("failed to delete offline session: %v", err)
|
||||
}
|
||||
|
||||
_, err = s.GetOfflineSessions(session1.UserID, session1.ConnID)
|
||||
_, err = s.GetOfflineSessions(ctx, session1.UserID, session1.ConnID)
|
||||
mustBeErrNotFound(t, "offline session", err)
|
||||
}
|
||||
|
||||
@@ -646,7 +646,7 @@ func testConnectorCRUD(t *testing.T, s storage.Storage) {
|
||||
}
|
||||
|
||||
getAndCompare := func(id string, want storage.Connector) {
|
||||
gr, err := s.GetConnector(id)
|
||||
gr, err := s.GetConnector(ctx, id)
|
||||
if err != nil {
|
||||
t.Errorf("get connector: %v", err)
|
||||
return
|
||||
@@ -660,7 +660,7 @@ func testConnectorCRUD(t *testing.T, s storage.Storage) {
|
||||
|
||||
getAndCompare(id1, c1)
|
||||
|
||||
if err := s.UpdateConnector(c1.ID, func(old storage.Connector) (storage.Connector, error) {
|
||||
if err := s.UpdateConnector(ctx, c1.ID, func(old storage.Connector) (storage.Connector, error) {
|
||||
old.Type = "oidc"
|
||||
return old, nil
|
||||
}); err != nil {
|
||||
@@ -672,7 +672,7 @@ func testConnectorCRUD(t *testing.T, s storage.Storage) {
|
||||
|
||||
connectorList := []storage.Connector{c1, c2}
|
||||
listAndCompare := func(want []storage.Connector) {
|
||||
connectors, err := s.ListConnectors()
|
||||
connectors, err := s.ListConnectors(ctx)
|
||||
if err != nil {
|
||||
t.Errorf("list connectors: %v", err)
|
||||
return
|
||||
@@ -690,21 +690,23 @@ func testConnectorCRUD(t *testing.T, s storage.Storage) {
|
||||
}
|
||||
listAndCompare(connectorList)
|
||||
|
||||
if err := s.DeleteConnector(c1.ID); err != nil {
|
||||
if err := s.DeleteConnector(ctx, c1.ID); err != nil {
|
||||
t.Fatalf("failed to delete connector: %v", err)
|
||||
}
|
||||
|
||||
if err := s.DeleteConnector(c2.ID); err != nil {
|
||||
if err := s.DeleteConnector(ctx, c2.ID); err != nil {
|
||||
t.Fatalf("failed to delete connector: %v", err)
|
||||
}
|
||||
|
||||
_, err = s.GetConnector(c1.ID)
|
||||
_, err = s.GetConnector(ctx, c1.ID)
|
||||
mustBeErrNotFound(t, "connector", err)
|
||||
}
|
||||
|
||||
func testKeysCRUD(t *testing.T, s storage.Storage) {
|
||||
ctx := context.TODO()
|
||||
|
||||
updateAndCompare := func(k storage.Keys) {
|
||||
err := s.UpdateKeys(func(oldKeys storage.Keys) (storage.Keys, error) {
|
||||
err := s.UpdateKeys(ctx, func(oldKeys storage.Keys) (storage.Keys, error) {
|
||||
return k, nil
|
||||
})
|
||||
if err != nil {
|
||||
@@ -712,7 +714,7 @@ func testKeysCRUD(t *testing.T, s storage.Storage) {
|
||||
return
|
||||
}
|
||||
|
||||
if got, err := s.GetKeys(); err != nil {
|
||||
if got, err := s.GetKeys(ctx); err != nil {
|
||||
t.Errorf("failed to get keys: %v", err)
|
||||
} else {
|
||||
got.NextRotation = got.NextRotation.UTC()
|
||||
@@ -786,24 +788,24 @@ func testGC(t *testing.T, s storage.Storage) {
|
||||
}
|
||||
|
||||
for _, tz := range []*time.Location{time.UTC, est, pst} {
|
||||
result, err := s.GarbageCollect(expiry.Add(-time.Hour).In(tz))
|
||||
result, err := s.GarbageCollect(ctx, expiry.Add(-time.Hour).In(tz))
|
||||
if err != nil {
|
||||
t.Errorf("garbage collection failed: %v", err)
|
||||
} else if result.AuthCodes != 0 || result.AuthRequests != 0 {
|
||||
t.Errorf("expected no garbage collection results, got %#v", result)
|
||||
}
|
||||
if _, err := s.GetAuthCode(c.ID); err != nil {
|
||||
if _, err := s.GetAuthCode(ctx, c.ID); err != nil {
|
||||
t.Errorf("expected to be able to get auth code after GC: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
if r, err := s.GarbageCollect(expiry.Add(time.Hour)); err != nil {
|
||||
if r, err := s.GarbageCollect(ctx, expiry.Add(time.Hour)); err != nil {
|
||||
t.Errorf("garbage collection failed: %v", err)
|
||||
} else if r.AuthCodes != 1 {
|
||||
t.Errorf("expected to garbage collect 1 objects, got %d", r.AuthCodes)
|
||||
}
|
||||
|
||||
if _, err := s.GetAuthCode(c.ID); err == nil {
|
||||
if _, err := s.GetAuthCode(ctx, c.ID); err == nil {
|
||||
t.Errorf("expected auth code to be GC'd")
|
||||
} else if err != storage.ErrNotFound {
|
||||
t.Errorf("expected storage.ErrNotFound, got %v", err)
|
||||
@@ -837,24 +839,24 @@ func testGC(t *testing.T, s storage.Storage) {
|
||||
}
|
||||
|
||||
for _, tz := range []*time.Location{time.UTC, est, pst} {
|
||||
result, err := s.GarbageCollect(expiry.Add(-time.Hour).In(tz))
|
||||
result, err := s.GarbageCollect(ctx, expiry.Add(-time.Hour).In(tz))
|
||||
if err != nil {
|
||||
t.Errorf("garbage collection failed: %v", err)
|
||||
} else if result.AuthCodes != 0 || result.AuthRequests != 0 {
|
||||
t.Errorf("expected no garbage collection results, got %#v", result)
|
||||
}
|
||||
if _, err := s.GetAuthRequest(a.ID); err != nil {
|
||||
if _, err := s.GetAuthRequest(ctx, a.ID); err != nil {
|
||||
t.Errorf("expected to be able to get auth request after GC: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
if r, err := s.GarbageCollect(expiry.Add(time.Hour)); err != nil {
|
||||
if r, err := s.GarbageCollect(ctx, expiry.Add(time.Hour)); err != nil {
|
||||
t.Errorf("garbage collection failed: %v", err)
|
||||
} else if r.AuthRequests != 1 {
|
||||
t.Errorf("expected to garbage collect 1 objects, got %d", r.AuthRequests)
|
||||
}
|
||||
|
||||
if _, err := s.GetAuthRequest(a.ID); err == nil {
|
||||
if _, err := s.GetAuthRequest(ctx, a.ID); err == nil {
|
||||
t.Errorf("expected auth request to be GC'd")
|
||||
} else if err != storage.ErrNotFound {
|
||||
t.Errorf("expected storage.ErrNotFound, got %v", err)
|
||||
@@ -874,23 +876,23 @@ func testGC(t *testing.T, s storage.Storage) {
|
||||
}
|
||||
|
||||
for _, tz := range []*time.Location{time.UTC, est, pst} {
|
||||
result, err := s.GarbageCollect(expiry.Add(-time.Hour).In(tz))
|
||||
result, err := s.GarbageCollect(ctx, expiry.Add(-time.Hour).In(tz))
|
||||
if err != nil {
|
||||
t.Errorf("garbage collection failed: %v", err)
|
||||
} else if result.DeviceRequests != 0 {
|
||||
t.Errorf("expected no device garbage collection results, got %#v", result)
|
||||
}
|
||||
if _, err := s.GetDeviceRequest(d.UserCode); err != nil {
|
||||
if _, err := s.GetDeviceRequest(ctx, d.UserCode); err != nil {
|
||||
t.Errorf("expected to be able to get auth request after GC: %v", err)
|
||||
}
|
||||
}
|
||||
if r, err := s.GarbageCollect(expiry.Add(time.Hour)); err != nil {
|
||||
if r, err := s.GarbageCollect(ctx, expiry.Add(time.Hour)); err != nil {
|
||||
t.Errorf("garbage collection failed: %v", err)
|
||||
} else if r.DeviceRequests != 1 {
|
||||
t.Errorf("expected to garbage collect 1 device request, got %d", r.DeviceRequests)
|
||||
}
|
||||
|
||||
if _, err := s.GetDeviceRequest(d.UserCode); err == nil {
|
||||
if _, err := s.GetDeviceRequest(ctx, d.UserCode); err == nil {
|
||||
t.Errorf("expected device request to be GC'd")
|
||||
} else if err != storage.ErrNotFound {
|
||||
t.Errorf("expected storage.ErrNotFound, got %v", err)
|
||||
@@ -914,23 +916,23 @@ func testGC(t *testing.T, s storage.Storage) {
|
||||
}
|
||||
|
||||
for _, tz := range []*time.Location{time.UTC, est, pst} {
|
||||
result, err := s.GarbageCollect(expiry.Add(-time.Hour).In(tz))
|
||||
result, err := s.GarbageCollect(ctx, expiry.Add(-time.Hour).In(tz))
|
||||
if err != nil {
|
||||
t.Errorf("garbage collection failed: %v", err)
|
||||
} else if result.DeviceTokens != 0 {
|
||||
t.Errorf("expected no device token garbage collection results, got %#v", result)
|
||||
}
|
||||
if _, err := s.GetDeviceToken(dt.DeviceCode); err != nil {
|
||||
if _, err := s.GetDeviceToken(ctx, dt.DeviceCode); err != nil {
|
||||
t.Errorf("expected to be able to get device token after GC: %v", err)
|
||||
}
|
||||
}
|
||||
if r, err := s.GarbageCollect(expiry.Add(time.Hour)); err != nil {
|
||||
if r, err := s.GarbageCollect(ctx, expiry.Add(time.Hour)); err != nil {
|
||||
t.Errorf("garbage collection failed: %v", err)
|
||||
} else if r.DeviceTokens != 1 {
|
||||
t.Errorf("expected to garbage collect 1 device token, got %d", r.DeviceTokens)
|
||||
}
|
||||
|
||||
if _, err := s.GetDeviceToken(dt.DeviceCode); err == nil {
|
||||
if _, err := s.GetDeviceToken(ctx, dt.DeviceCode); err == nil {
|
||||
t.Errorf("expected device token to be GC'd")
|
||||
} else if err != storage.ErrNotFound {
|
||||
t.Errorf("expected storage.ErrNotFound, got %v", err)
|
||||
@@ -969,7 +971,7 @@ func testTimezones(t *testing.T, s storage.Storage) {
|
||||
if err := s.CreateAuthCode(ctx, c); err != nil {
|
||||
t.Fatalf("failed creating auth code: %v", err)
|
||||
}
|
||||
got, err := s.GetAuthCode(c.ID)
|
||||
got, err := s.GetAuthCode(ctx, c.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to get auth code: %v", err)
|
||||
}
|
||||
@@ -1003,7 +1005,7 @@ func testDeviceRequestCRUD(t *testing.T, s storage.Storage) {
|
||||
err := s.CreateDeviceRequest(ctx, d1)
|
||||
mustBeErrAlreadyExists(t, "device request", err)
|
||||
|
||||
got, err := s.GetDeviceRequest(d1.UserCode)
|
||||
got, err := s.GetDeviceRequest(ctx, d1.UserCode)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to get device request: %v", err)
|
||||
}
|
||||
@@ -1041,7 +1043,7 @@ func testDeviceTokenCRUD(t *testing.T, s storage.Storage) {
|
||||
mustBeErrAlreadyExists(t, "device token", err)
|
||||
|
||||
// Update the device token, simulate a redemption
|
||||
if err := s.UpdateDeviceToken(d1.DeviceCode, func(old storage.DeviceToken) (storage.DeviceToken, error) {
|
||||
if err := s.UpdateDeviceToken(ctx, d1.DeviceCode, func(old storage.DeviceToken) (storage.DeviceToken, error) {
|
||||
old.Token = "token data"
|
||||
old.Status = "complete"
|
||||
return old, nil
|
||||
@@ -1050,7 +1052,7 @@ func testDeviceTokenCRUD(t *testing.T, s storage.Storage) {
|
||||
}
|
||||
|
||||
// Retrieve the device token
|
||||
got, err := s.GetDeviceToken(d1.DeviceCode)
|
||||
got, err := s.GetDeviceToken(ctx, d1.DeviceCode)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to get device token: %v", err)
|
||||
}
|
||||
|
||||
@@ -42,9 +42,9 @@ func testClientConcurrentUpdate(t *testing.T, s storage.Storage) {
|
||||
|
||||
var err1, err2 error
|
||||
|
||||
err1 = s.UpdateClient(c.ID, func(old storage.Client) (storage.Client, error) {
|
||||
err1 = s.UpdateClient(ctx, c.ID, func(old storage.Client) (storage.Client, error) {
|
||||
old.Secret = "new secret 1"
|
||||
err2 = s.UpdateClient(c.ID, func(old storage.Client) (storage.Client, error) {
|
||||
err2 = s.UpdateClient(ctx, c.ID, func(old storage.Client) (storage.Client, error) {
|
||||
old.Secret = "new secret 2"
|
||||
return old, nil
|
||||
})
|
||||
@@ -87,9 +87,9 @@ func testAuthRequestConcurrentUpdate(t *testing.T, s storage.Storage) {
|
||||
|
||||
var err1, err2 error
|
||||
|
||||
err1 = s.UpdateAuthRequest(a.ID, func(old storage.AuthRequest) (storage.AuthRequest, error) {
|
||||
err1 = s.UpdateAuthRequest(ctx, a.ID, func(old storage.AuthRequest) (storage.AuthRequest, error) {
|
||||
old.State = "state 1"
|
||||
err2 = s.UpdateAuthRequest(a.ID, func(old storage.AuthRequest) (storage.AuthRequest, error) {
|
||||
err2 = s.UpdateAuthRequest(ctx, a.ID, func(old storage.AuthRequest) (storage.AuthRequest, error) {
|
||||
old.State = "state 2"
|
||||
return old, nil
|
||||
})
|
||||
@@ -121,9 +121,9 @@ func testPasswordConcurrentUpdate(t *testing.T, s storage.Storage) {
|
||||
|
||||
var err1, err2 error
|
||||
|
||||
err1 = s.UpdatePassword(password.Email, func(old storage.Password) (storage.Password, error) {
|
||||
err1 = s.UpdatePassword(ctx, password.Email, func(old storage.Password) (storage.Password, error) {
|
||||
old.Username = "user 1"
|
||||
err2 = s.UpdatePassword(password.Email, func(old storage.Password) (storage.Password, error) {
|
||||
err2 = s.UpdatePassword(ctx, password.Email, func(old storage.Password) (storage.Password, error) {
|
||||
old.Username = "user 2"
|
||||
return old, nil
|
||||
})
|
||||
@@ -163,8 +163,9 @@ func testKeysConcurrentUpdate(t *testing.T, s storage.Storage) {
|
||||
|
||||
var err1, err2 error
|
||||
|
||||
err1 = s.UpdateKeys(func(old storage.Keys) (storage.Keys, error) {
|
||||
err2 = s.UpdateKeys(func(old storage.Keys) (storage.Keys, error) {
|
||||
ctx := context.TODO()
|
||||
err1 = s.UpdateKeys(ctx, func(old storage.Keys) (storage.Keys, error) {
|
||||
err2 = s.UpdateKeys(ctx, func(old storage.Keys) (storage.Keys, error) {
|
||||
return keys1, nil
|
||||
})
|
||||
return keys2, nil
|
||||
|
||||
@@ -34,8 +34,8 @@ func (d *Database) CreateAuthCode(ctx context.Context, code storage.AuthCode) er
|
||||
}
|
||||
|
||||
// GetAuthCode extracts an auth code from the database by id.
|
||||
func (d *Database) GetAuthCode(id string) (storage.AuthCode, error) {
|
||||
authCode, err := d.client.AuthCode.Get(context.TODO(), id)
|
||||
func (d *Database) GetAuthCode(ctx context.Context, id string) (storage.AuthCode, error) {
|
||||
authCode, err := d.client.AuthCode.Get(ctx, id)
|
||||
if err != nil {
|
||||
return storage.AuthCode{}, convertDBError("get auth code: %w", err)
|
||||
}
|
||||
@@ -43,8 +43,8 @@ func (d *Database) GetAuthCode(id string) (storage.AuthCode, error) {
|
||||
}
|
||||
|
||||
// DeleteAuthCode deletes an auth code from the database by id.
|
||||
func (d *Database) DeleteAuthCode(id string) error {
|
||||
err := d.client.AuthCode.DeleteOneID(id).Exec(context.TODO())
|
||||
func (d *Database) DeleteAuthCode(ctx context.Context, id string) error {
|
||||
err := d.client.AuthCode.DeleteOneID(id).Exec(ctx)
|
||||
if err != nil {
|
||||
return convertDBError("delete auth code: %w", err)
|
||||
}
|
||||
|
||||
@@ -40,8 +40,8 @@ func (d *Database) CreateAuthRequest(ctx context.Context, authRequest storage.Au
|
||||
}
|
||||
|
||||
// GetAuthRequest extracts an auth request from the database by id.
|
||||
func (d *Database) GetAuthRequest(id string) (storage.AuthRequest, error) {
|
||||
authRequest, err := d.client.AuthRequest.Get(context.TODO(), id)
|
||||
func (d *Database) GetAuthRequest(ctx context.Context, id string) (storage.AuthRequest, error) {
|
||||
authRequest, err := d.client.AuthRequest.Get(ctx, id)
|
||||
if err != nil {
|
||||
return storage.AuthRequest{}, convertDBError("get auth request: %w", err)
|
||||
}
|
||||
@@ -49,8 +49,8 @@ func (d *Database) GetAuthRequest(id string) (storage.AuthRequest, error) {
|
||||
}
|
||||
|
||||
// DeleteAuthRequest deletes an auth request from the database by id.
|
||||
func (d *Database) DeleteAuthRequest(id string) error {
|
||||
err := d.client.AuthRequest.DeleteOneID(id).Exec(context.TODO())
|
||||
func (d *Database) DeleteAuthRequest(ctx context.Context, id string) error {
|
||||
err := d.client.AuthRequest.DeleteOneID(id).Exec(ctx)
|
||||
if err != nil {
|
||||
return convertDBError("delete auth request: %w", err)
|
||||
}
|
||||
@@ -58,8 +58,8 @@ func (d *Database) DeleteAuthRequest(id string) error {
|
||||
}
|
||||
|
||||
// UpdateAuthRequest changes an auth request by id using an updater function and saves it to the database.
|
||||
func (d *Database) UpdateAuthRequest(id string, updater func(old storage.AuthRequest) (storage.AuthRequest, error)) error {
|
||||
tx, err := d.BeginTx(context.TODO())
|
||||
func (d *Database) UpdateAuthRequest(ctx context.Context, id string, updater func(old storage.AuthRequest) (storage.AuthRequest, error)) error {
|
||||
tx, err := d.BeginTx(ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("update auth request tx: %w", err)
|
||||
}
|
||||
|
||||
@@ -24,8 +24,8 @@ func (d *Database) CreateClient(ctx context.Context, client storage.Client) erro
|
||||
}
|
||||
|
||||
// ListClients extracts an array of oauth2 clients from the database.
|
||||
func (d *Database) ListClients() ([]storage.Client, error) {
|
||||
clients, err := d.client.OAuth2Client.Query().All(context.TODO())
|
||||
func (d *Database) ListClients(ctx context.Context) ([]storage.Client, error) {
|
||||
clients, err := d.client.OAuth2Client.Query().All(ctx)
|
||||
if err != nil {
|
||||
return nil, convertDBError("list clients: %w", err)
|
||||
}
|
||||
@@ -38,8 +38,8 @@ func (d *Database) ListClients() ([]storage.Client, error) {
|
||||
}
|
||||
|
||||
// GetClient extracts an oauth2 client from the database by id.
|
||||
func (d *Database) GetClient(id string) (storage.Client, error) {
|
||||
client, err := d.client.OAuth2Client.Get(context.TODO(), id)
|
||||
func (d *Database) GetClient(ctx context.Context, id string) (storage.Client, error) {
|
||||
client, err := d.client.OAuth2Client.Get(ctx, id)
|
||||
if err != nil {
|
||||
return storage.Client{}, convertDBError("get client: %w", err)
|
||||
}
|
||||
@@ -47,8 +47,8 @@ func (d *Database) GetClient(id string) (storage.Client, error) {
|
||||
}
|
||||
|
||||
// DeleteClient deletes an oauth2 client from the database by id.
|
||||
func (d *Database) DeleteClient(id string) error {
|
||||
err := d.client.OAuth2Client.DeleteOneID(id).Exec(context.TODO())
|
||||
func (d *Database) DeleteClient(ctx context.Context, id string) error {
|
||||
err := d.client.OAuth2Client.DeleteOneID(id).Exec(ctx)
|
||||
if err != nil {
|
||||
return convertDBError("delete client: %w", err)
|
||||
}
|
||||
@@ -56,13 +56,13 @@ func (d *Database) DeleteClient(id string) error {
|
||||
}
|
||||
|
||||
// UpdateClient changes an oauth2 client by id using an updater function and saves it to the database.
|
||||
func (d *Database) UpdateClient(id string, updater func(old storage.Client) (storage.Client, error)) error {
|
||||
tx, err := d.BeginTx(context.TODO())
|
||||
func (d *Database) UpdateClient(ctx context.Context, id string, updater func(old storage.Client) (storage.Client, error)) error {
|
||||
tx, err := d.BeginTx(ctx)
|
||||
if err != nil {
|
||||
return convertDBError("update client tx: %w", err)
|
||||
}
|
||||
|
||||
client, err := tx.OAuth2Client.Get(context.TODO(), id)
|
||||
client, err := tx.OAuth2Client.Get(ctx, id)
|
||||
if err != nil {
|
||||
return rollback(tx, "update client database: %w", err)
|
||||
}
|
||||
@@ -79,7 +79,7 @@ func (d *Database) UpdateClient(id string, updater func(old storage.Client) (sto
|
||||
SetLogoURL(newClient.LogoURL).
|
||||
SetRedirectUris(newClient.RedirectURIs).
|
||||
SetTrustedPeers(newClient.TrustedPeers).
|
||||
Save(context.TODO())
|
||||
Save(ctx)
|
||||
if err != nil {
|
||||
return rollback(tx, "update client uploading: %w", err)
|
||||
}
|
||||
|
||||
@@ -22,8 +22,8 @@ func (d *Database) CreateConnector(ctx context.Context, connector storage.Connec
|
||||
}
|
||||
|
||||
// ListConnectors extracts an array of connectors from the database.
|
||||
func (d *Database) ListConnectors() ([]storage.Connector, error) {
|
||||
connectors, err := d.client.Connector.Query().All(context.TODO())
|
||||
func (d *Database) ListConnectors(ctx context.Context) ([]storage.Connector, error) {
|
||||
connectors, err := d.client.Connector.Query().All(ctx)
|
||||
if err != nil {
|
||||
return nil, convertDBError("list connectors: %w", err)
|
||||
}
|
||||
@@ -36,8 +36,8 @@ func (d *Database) ListConnectors() ([]storage.Connector, error) {
|
||||
}
|
||||
|
||||
// GetConnector extracts a connector from the database by id.
|
||||
func (d *Database) GetConnector(id string) (storage.Connector, error) {
|
||||
connector, err := d.client.Connector.Get(context.TODO(), id)
|
||||
func (d *Database) GetConnector(ctx context.Context, id string) (storage.Connector, error) {
|
||||
connector, err := d.client.Connector.Get(ctx, id)
|
||||
if err != nil {
|
||||
return storage.Connector{}, convertDBError("get connector: %w", err)
|
||||
}
|
||||
@@ -45,8 +45,8 @@ func (d *Database) GetConnector(id string) (storage.Connector, error) {
|
||||
}
|
||||
|
||||
// DeleteConnector deletes a connector from the database by id.
|
||||
func (d *Database) DeleteConnector(id string) error {
|
||||
err := d.client.Connector.DeleteOneID(id).Exec(context.TODO())
|
||||
func (d *Database) DeleteConnector(ctx context.Context, id string) error {
|
||||
err := d.client.Connector.DeleteOneID(id).Exec(ctx)
|
||||
if err != nil {
|
||||
return convertDBError("delete connector: %w", err)
|
||||
}
|
||||
@@ -54,13 +54,13 @@ func (d *Database) DeleteConnector(id string) error {
|
||||
}
|
||||
|
||||
// UpdateConnector changes a connector by id using an updater function and saves it to the database.
|
||||
func (d *Database) UpdateConnector(id string, updater func(old storage.Connector) (storage.Connector, error)) error {
|
||||
tx, err := d.BeginTx(context.TODO())
|
||||
func (d *Database) UpdateConnector(ctx context.Context, id string, updater func(old storage.Connector) (storage.Connector, error)) error {
|
||||
tx, err := d.BeginTx(ctx)
|
||||
if err != nil {
|
||||
return convertDBError("update connector tx: %w", err)
|
||||
}
|
||||
|
||||
connector, err := tx.Connector.Get(context.TODO(), id)
|
||||
connector, err := tx.Connector.Get(ctx, id)
|
||||
if err != nil {
|
||||
return rollback(tx, "update connector database: %w", err)
|
||||
}
|
||||
@@ -75,7 +75,7 @@ func (d *Database) UpdateConnector(id string, updater func(old storage.Connector
|
||||
SetType(newConnector.Type).
|
||||
SetResourceVersion(newConnector.ResourceVersion).
|
||||
SetConfig(newConnector.Config).
|
||||
Save(context.TODO())
|
||||
Save(ctx)
|
||||
if err != nil {
|
||||
return rollback(tx, "update connector uploading: %w", err)
|
||||
}
|
||||
|
||||
@@ -25,10 +25,10 @@ func (d *Database) CreateDeviceRequest(ctx context.Context, request storage.Devi
|
||||
}
|
||||
|
||||
// GetDeviceRequest extracts a device request from the database by user code.
|
||||
func (d *Database) GetDeviceRequest(userCode string) (storage.DeviceRequest, error) {
|
||||
func (d *Database) GetDeviceRequest(ctx context.Context, userCode string) (storage.DeviceRequest, error) {
|
||||
deviceRequest, err := d.client.DeviceRequest.Query().
|
||||
Where(devicerequest.UserCode(userCode)).
|
||||
Only(context.TODO())
|
||||
Only(ctx)
|
||||
if err != nil {
|
||||
return storage.DeviceRequest{}, convertDBError("get device request: %w", err)
|
||||
}
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user