package oauth2 import ( "crypto/rand" "encoding/base64" "errors" "time" ) // In the oauth2 flow a "state" (arbitrary string) is passed // - from our server (writing the state in the query of the url where we want // the client to go) // - to the client (reading the response) // - to the oauth2 server (when the client follow the redirection) // - to our server (when the oauth2 callbacks) // // For this reason each oauth2.Authenticator must stores an object that manages // these states. // // States are used to ensure the oauth2 callback is valid and to switch // over the right behavior (currently: // - log-in // - sign-up // but this should be extended to arbitrary methods/endpoints). // // An object that manges the states for an oauth2.Authenticator needs to // implement the OAuth2StateProvider interface: type OAuth2StateProvider interface { GenerateState(scope string, exp time.Duration) (string, error) Validate(token string) error Delete(token string) error GetScope(token string) (string, error) } // We implement this interface as a mere map. type States map[string]stateValue // The "states" we exchange with the oauth2 servver are the keys and the values // represent the actual data. // This makes the request opaque for the oauth2 server, that only sees random // keys. // Right now each state is just an expiration data and a scope // later we should probably use jwt type stateValue struct { exp time.Time scope string } // What follows it the interface implementation: func (s States) GenerateState(scope string, exp time.Duration) (string, error) { var buf [32]byte if _, err := rand.Read(buf[:]); err != nil { return "", err } key := base64.RawURLEncoding.EncodeToString(buf[:]) s[key] = stateValue{ exp: time.Now().Add(exp), scope: scope, } return key, nil } func (s States) Validate(state string) error { val, ok := s[state] if !ok { return errors.New("ValidateState: no id") } if time.Now().After(val.exp) { return errors.New("ValidateState: expired") } return nil } func (s States) Delete(state string) error { delete(s, state) return nil } func (s States) GetScope(token string) (string, error) { state, ok := s[token] if !ok { return "", errors.New("GetScope: token not found") } return state.scope, nil }