Fear: impl google oauth2

pull/21/head
zijiren233 3 years ago
parent 854e78c1f4
commit 9467046f7f

@ -9,7 +9,7 @@ import (
func InitProvider(ctx context.Context) error { func InitProvider(ctx context.Context) error {
for op, v := range conf.Conf.OAuth2 { for op, v := range conf.Conf.OAuth2 {
err := provider.InitProvider(op, v.ClientID, v.ClientSecret) err := provider.InitProvider(op, v.ClientID, v.ClientSecret, provider.WithRedirectURL(v.RedirectURL))
if err != nil { if err != nil {
return err return err
} }

@ -9,13 +9,15 @@ type OAuth2Config map[provider.OAuth2Provider]OAuth2ProviderConfig
type OAuth2ProviderConfig struct { type OAuth2ProviderConfig struct {
ClientID string `yaml:"client_id"` ClientID string `yaml:"client_id"`
ClientSecret string `yaml:"client_secret"` ClientSecret string `yaml:"client_secret"`
RedirectURL string `yaml:"redirect_url"`
} }
func DefaultOAuth2Config() OAuth2Config { func DefaultOAuth2Config() OAuth2Config {
return OAuth2Config{ return OAuth2Config{
(&provider.GithubProvider{}).Provider(): { (&provider.GithubProvider{}).Provider(): {
ClientID: "github_client_id", ClientID: "",
ClientSecret: "github_client_secret", ClientSecret: "",
RedirectURL: "",
}, },
} }
} }

@ -10,25 +10,29 @@ import (
) )
type GithubProvider struct { type GithubProvider struct {
ClientID, ClientSecret string config oauth2.Config
} }
func (p *GithubProvider) Init(ClientID, ClientSecret string) { func (p *GithubProvider) Init(ClientID, ClientSecret string, options ...Oauth2Option) {
p.ClientID = ClientID p.config.ClientID = ClientID
p.ClientSecret = ClientSecret p.config.ClientSecret = ClientSecret
p.config.Scopes = []string{"user"}
p.config.Endpoint = github.Endpoint
for _, o := range options {
o(&p.config)
}
} }
func (p *GithubProvider) Provider() OAuth2Provider { func (p *GithubProvider) Provider() OAuth2Provider {
return "github" return "github"
} }
func (p *GithubProvider) NewConfig() *oauth2.Config { func (p *GithubProvider) NewConfig(options ...Oauth2Option) *oauth2.Config {
return &oauth2.Config{ c := p.config
ClientID: p.ClientID, for _, o := range options {
ClientSecret: p.ClientSecret, o(&c)
Scopes: []string{"user"},
Endpoint: github.Endpoint,
} }
return &c
} }
func (p *GithubProvider) GetUserInfo(ctx context.Context, config *oauth2.Config, code string) (*UserInfo, error) { func (p *GithubProvider) GetUserInfo(ctx context.Context, config *oauth2.Config, code string) (*UserInfo, error) {

@ -9,25 +9,29 @@ import (
) )
type GitlabProvider struct { type GitlabProvider struct {
ClientID, ClientSecret string config oauth2.Config
} }
func (g *GitlabProvider) Init(ClientID, ClientSecret string) { func (g *GitlabProvider) Init(ClientID, ClientSecret string, options ...Oauth2Option) {
g.ClientID = ClientID g.config.ClientID = ClientID
g.ClientSecret = ClientSecret g.config.ClientSecret = ClientSecret
g.config.Scopes = []string{"read_user"}
g.config.Endpoint = gitlab.Endpoint
for _, o := range options {
o(&g.config)
}
} }
func (g *GitlabProvider) Provider() OAuth2Provider { func (g *GitlabProvider) Provider() OAuth2Provider {
return "gitlab" return "gitlab"
} }
func (g *GitlabProvider) NewConfig() *oauth2.Config { func (g *GitlabProvider) NewConfig(options ...Oauth2Option) *oauth2.Config {
return &oauth2.Config{ c := g.config
ClientID: g.ClientID, for _, o := range options {
ClientSecret: g.ClientSecret, o(&c)
Scopes: []string{"read_user"},
Endpoint: gitlab.Endpoint,
} }
return &c
} }
func (g *GitlabProvider) GetUserInfo(ctx context.Context, config *oauth2.Config, code string) (*UserInfo, error) { func (g *GitlabProvider) GetUserInfo(ctx context.Context, config *oauth2.Config, code string) (*UserInfo, error) {

@ -4,30 +4,35 @@ import (
"context" "context"
"net/http" "net/http"
json "github.com/json-iterator/go"
"golang.org/x/oauth2" "golang.org/x/oauth2"
"golang.org/x/oauth2/google" "golang.org/x/oauth2/google"
) )
type GoogleProvider struct { type GoogleProvider struct {
ClientID, ClientSecret string config oauth2.Config
} }
func (g *GoogleProvider) Init(ClientID, ClientSecret string) { func (g *GoogleProvider) Init(ClientID, ClientSecret string, options ...Oauth2Option) {
g.ClientID = ClientID g.config.ClientID = ClientID
g.ClientSecret = ClientSecret g.config.ClientSecret = ClientSecret
g.config.Scopes = []string{"profile"}
g.config.Endpoint = google.Endpoint
for _, o := range options {
o(&g.config)
}
} }
func (g *GoogleProvider) Provider() OAuth2Provider { func (g *GoogleProvider) Provider() OAuth2Provider {
return "google" return "google"
} }
func (g *GoogleProvider) NewConfig() *oauth2.Config { func (g *GoogleProvider) NewConfig(options ...Oauth2Option) *oauth2.Config {
return &oauth2.Config{ c := g.config
ClientID: g.ClientID, for _, o := range options {
ClientSecret: g.ClientSecret, o(&c)
Scopes: []string{"profile"},
Endpoint: google.Endpoint,
} }
return &c
} }
func (g *GoogleProvider) GetUserInfo(ctx context.Context, config *oauth2.Config, code string) (*UserInfo, error) { func (g *GoogleProvider) GetUserInfo(ctx context.Context, config *oauth2.Config, code string) (*UserInfo, error) {
@ -45,9 +50,22 @@ func (g *GoogleProvider) GetUserInfo(ctx context.Context, config *oauth2.Config,
return nil, err return nil, err
} }
defer resp.Body.Close() defer resp.Body.Close()
return nil, FormatErrNotImplemented("google") ui := googleUserInfo{}
err = json.NewDecoder(resp.Body).Decode(&ui)
if err != nil {
return nil, err
}
return &UserInfo{
Username: ui.Name,
ProviderUserID: ui.ID,
}, nil
} }
func init() { func init() {
registerProvider(new(GoogleProvider)) registerProvider(new(GoogleProvider))
} }
type googleUserInfo struct {
ID uint `json:"id,string"`
Name string `json:"name"`
}

@ -19,19 +19,27 @@ type UserInfo struct {
ProviderUserID uint ProviderUserID uint
} }
type Oauth2Option func(*oauth2.Config)
func WithRedirectURL(url string) Oauth2Option {
return func(c *oauth2.Config) {
c.RedirectURL = url
}
}
type ProviderInterface interface { type ProviderInterface interface {
Init(ClientID, ClientSecret string) Init(ClientID, ClientSecret string, options ...Oauth2Option)
Provider() OAuth2Provider Provider() OAuth2Provider
NewConfig() *oauth2.Config NewConfig(options ...Oauth2Option) *oauth2.Config
GetUserInfo(ctx context.Context, config *oauth2.Config, code string) (*UserInfo, error) GetUserInfo(ctx context.Context, config *oauth2.Config, code string) (*UserInfo, error)
} }
func InitProvider(p OAuth2Provider, ClientID, ClientSecret string) error { func InitProvider(p OAuth2Provider, ClientID, ClientSecret string, options ...Oauth2Option) error {
pi, ok := allowedProviders[p] pi, ok := allowedProviders[p]
if !ok { if !ok {
return FormatErrNotImplemented(p) return FormatErrNotImplemented(p)
} }
pi.Init(ClientID, ClientSecret) pi.Init(ClientID, ClientSecret, options...)
if enabledProviders == nil { if enabledProviders == nil {
enabledProviders = make(map[OAuth2Provider]ProviderInterface) enabledProviders = make(map[OAuth2Provider]ProviderInterface)
} }

@ -5,7 +5,6 @@ import (
"time" "time"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/sirupsen/logrus"
"github.com/synctv-org/synctv/internal/op" "github.com/synctv-org/synctv/internal/op"
"github.com/synctv-org/synctv/internal/provider" "github.com/synctv-org/synctv/internal/provider"
"github.com/synctv-org/synctv/server/middlewares" "github.com/synctv-org/synctv/server/middlewares"
@ -16,9 +15,9 @@ import (
// /oauth2/login/:type // /oauth2/login/:type
func OAuth2(ctx *gin.Context) { func OAuth2(ctx *gin.Context) {
p := provider.OAuth2Provider(ctx.Param("type")) t := ctx.Param("type")
pi, err := provider.GetProvider(p) pi, err := provider.GetProvider(provider.OAuth2Provider(t))
if err != nil { if err != nil {
ctx.AbortWithStatusJSON(http.StatusBadRequest, model.NewApiErrorResp(err)) ctx.AbortWithStatusJSON(http.StatusBadRequest, model.NewApiErrorResp(err))
return return
@ -31,9 +30,8 @@ func OAuth2(ctx *gin.Context) {
} }
func OAuth2Api(ctx *gin.Context) { func OAuth2Api(ctx *gin.Context) {
p := provider.OAuth2Provider(ctx.Param("type")) t := ctx.Param("type")
pi, err := provider.GetProvider(provider.OAuth2Provider(t))
pi, err := provider.GetProvider(p)
if err != nil { if err != nil {
ctx.AbortWithStatusJSON(http.StatusBadRequest, model.NewApiErrorResp(err)) ctx.AbortWithStatusJSON(http.StatusBadRequest, model.NewApiErrorResp(err))
} }
@ -48,8 +46,6 @@ func OAuth2Api(ctx *gin.Context) {
// /oauth2/callback/:type // /oauth2/callback/:type
func OAuth2Callback(ctx *gin.Context) { func OAuth2Callback(ctx *gin.Context) {
p := provider.OAuth2Provider(ctx.Param("type"))
code := ctx.Query("code") code := ctx.Query("code")
if code == "" { if code == "" {
ctx.AbortWithStatusJSON(http.StatusBadRequest, model.NewApiErrorStringResp("invalid oauth2 code")) ctx.AbortWithStatusJSON(http.StatusBadRequest, model.NewApiErrorStringResp("invalid oauth2 code"))
@ -68,6 +64,7 @@ func OAuth2Callback(ctx *gin.Context) {
return return
} }
p := provider.OAuth2Provider(ctx.Param("type"))
pi, err := provider.GetProvider(p) pi, err := provider.GetProvider(p)
if err != nil { if err != nil {
ctx.AbortWithStatusJSON(http.StatusBadRequest, model.NewApiErrorResp(err)) ctx.AbortWithStatusJSON(http.StatusBadRequest, model.NewApiErrorResp(err))
@ -91,15 +88,11 @@ func OAuth2Callback(ctx *gin.Context) {
return return
} }
logrus.Info("asdasd")
RenderToken(ctx, "/web/", token) RenderToken(ctx, "/web/", token)
} }
// /oauth2/callback/:type // /oauth2/callback/:type
func OAuth2CallbackApi(ctx *gin.Context) { func OAuth2CallbackApi(ctx *gin.Context) {
p := provider.OAuth2Provider(ctx.Param("type"))
req := model.OAuth2CallbackReq{} req := model.OAuth2CallbackReq{}
if err := req.Decode(ctx); err != nil { if err := req.Decode(ctx); err != nil {
ctx.AbortWithStatusJSON(http.StatusBadRequest, model.NewApiErrorResp(err)) ctx.AbortWithStatusJSON(http.StatusBadRequest, model.NewApiErrorResp(err))
@ -112,6 +105,7 @@ func OAuth2CallbackApi(ctx *gin.Context) {
return return
} }
p := provider.OAuth2Provider(ctx.Param("type"))
pi, err := provider.GetProvider(p) pi, err := provider.GetProvider(p)
if err != nil { if err != nil {
ctx.AbortWithStatusJSON(http.StatusBadRequest, model.NewApiErrorResp(err)) ctx.AbortWithStatusJSON(http.StatusBadRequest, model.NewApiErrorResp(err))

@ -8,9 +8,18 @@ import (
type SyncCache[K comparable, V any] struct { type SyncCache[K comparable, V any] struct {
cache rwmap.RWMap[K, *entry[V]] cache rwmap.RWMap[K, *entry[V]]
deletedCallback func(v V)
ticker *time.Ticker ticker *time.Ticker
} }
type SyncCacheConfig[K comparable, V any] func(sc *SyncCache[K, V])
func WithDeletedCallback[K comparable, V any](callback func(v V)) SyncCacheConfig[K, V] {
return func(sc *SyncCache[K, V]) {
sc.deletedCallback = callback
}
}
func NewSyncCache[K comparable, V any](trimTime time.Duration) *SyncCache[K, V] { func NewSyncCache[K comparable, V any](trimTime time.Duration) *SyncCache[K, V] {
sc := &SyncCache[K, V]{ sc := &SyncCache[K, V]{
ticker: time.NewTicker(trimTime), ticker: time.NewTicker(trimTime),
@ -31,7 +40,10 @@ func (sc *SyncCache[K, V]) Releases() {
func (sc *SyncCache[K, V]) trim() { func (sc *SyncCache[K, V]) trim() {
sc.cache.Range(func(key K, value *entry[V]) bool { sc.cache.Range(func(key K, value *entry[V]) bool {
if value.IsExpired() { if value.IsExpired() {
sc.cache.Delete(key) e, loaded := sc.cache.LoadAndDelete(key)
if loaded && sc.deletedCallback != nil {
sc.deletedCallback(e.value)
}
} }
return true return true
}) })

Loading…
Cancel
Save