diff --git a/backend/cmd/server/main.go b/backend/cmd/server/main.go index f02612f..65c3152 100644 --- a/backend/cmd/server/main.go +++ b/backend/cmd/server/main.go @@ -8,13 +8,11 @@ import ( "echsen.club/radio/internal/config" "echsen.club/radio/internal/db" - "echsen.club/radio/internal/radio" + "echsen.club/radio/internal/state" ) - func main() { config := config.Setup() - port := config.Port database, err := db.New(config.PostgresUrl) if err != nil { @@ -27,18 +25,12 @@ func main() { slog.Error("Migration failed", "error", err) os.Exit(1) } - - mux := http.NewServeMux() - - svc := radio.NewService(database, config) - hdl := radio.NewHandler(svc) - hdl.RegisterHandlers(mux) - protected_mux := hdl.AuthMiddleware(mux) + state := state.New(database, config) log.Println("Server starting on port ", config.Port) - if err := http.ListenAndServe(port, protected_mux); err != nil { - log.Fatalf("Listen error: %v", err) - } + if err := http.ListenAndServe(config.Port, state.Handler()); err != nil { + log.Fatalf("Listen error: %v", err) + } } diff --git a/backend/internal/radio/channel.go b/backend/internal/channel/channel.go similarity index 87% rename from backend/internal/radio/channel.go rename to backend/internal/channel/channel.go index a1930c3..302e518 100644 --- a/backend/internal/radio/channel.go +++ b/backend/internal/channel/channel.go @@ -1,4 +1,4 @@ -package radio +package channel type Channel struct { } diff --git a/backend/internal/channel/handler.go b/backend/internal/channel/handler.go new file mode 100644 index 0000000..c1c0e46 --- /dev/null +++ b/backend/internal/channel/handler.go @@ -0,0 +1,42 @@ +package channel + +import ( + "context" + "fmt" + "log/slog" + + "connectrpc.com/connect" + radiov1 "echsen.club/radio/gen/radio/v1" + "echsen.club/radio/internal/db" +) + +type StateProvider interface { + GetUserByID(ctx context.Context, id string) (*db.User, error) + UserFromContext(ctx context.Context) (*db.User, bool) +} + +type ChannelHandler struct { + state StateProvider +} + +func NewHandler(s StateProvider) *ChannelHandler { + return &ChannelHandler{ + state: s, + } +} + +func (handler *ChannelHandler) CreateChannel(ctx context.Context, req *radiov1.CreateChannelRequest) (*radiov1.CreateChannelResponse, error) { + + _, ok := handler.state.UserFromContext(ctx) + if !ok { + slog.Error("User missing from context in protected route") + return nil, connect.NewError(connect.CodeUnauthenticated, fmt.Errorf("User not authenticated")) + } + + channel := &radiov1.Channel{Id: "1", Frequency: "97.1", Description: ""} + return &radiov1.CreateChannelResponse{Channel: channel}, nil +} + +func (handler *ChannelHandler) DeleteChannel(ctx context.Context, req *radiov1.DeleteChannelRequest) (*radiov1.DeleteChannelResponse, error) { + return &radiov1.DeleteChannelResponse{}, nil +} diff --git a/backend/internal/db/user.go b/backend/internal/db/user.go new file mode 100644 index 0000000..44c826e --- /dev/null +++ b/backend/internal/db/user.go @@ -0,0 +1,12 @@ +package db + +import ( + "github.com/google/uuid" +) + +type User struct { + ID uuid.UUID + SubjectID string + Email string + DisplayName string +} diff --git a/backend/internal/auth/oidc.go b/backend/internal/oauth/oidc.go similarity index 54% rename from backend/internal/auth/oidc.go rename to backend/internal/oauth/oidc.go index 3ca478a..0064bb2 100644 --- a/backend/internal/auth/oidc.go +++ b/backend/internal/oauth/oidc.go @@ -1,4 +1,4 @@ -package auth +package oauth import ( "context" @@ -9,6 +9,7 @@ import ( "io" "log/slog" "net/http" + "net/url" "time" "echsen.club/radio/internal/config" @@ -18,14 +19,14 @@ import ( "golang.org/x/oauth2" ) -type OauthService struct { - Store sessions.Store +type Oauth struct { + Store sessions.Store Oauth2Config oauth2.Config - Verifier *oidc.IDTokenVerifier - LogoutUrl string + Verifier *oidc.IDTokenVerifier + LogoutUrl string } -func NewOauthService(db *sql.DB, config config.AuthConfig) (*OauthService, error) { +func New(db *sql.DB, config config.AuthConfig) (*Oauth, error) { ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) defer cancel() @@ -35,11 +36,11 @@ func NewOauthService(db *sql.DB, config config.AuthConfig) (*OauthService, error } oauth2Config := oauth2.Config{ - ClientID: config.ClientId, + ClientID: config.ClientId, ClientSecret: config.ClientSecret, RedirectURL: config.RedirectUrl, - Endpoint: provider.Endpoint(), - Scopes: []string{oidc.ScopeOpenID, "profile", "email"}, + Endpoint: provider.Endpoint(), + Scopes: []string{oidc.ScopeOpenID, "profile", "email"}, } var claims struct { @@ -61,11 +62,11 @@ func NewOauthService(db *sql.DB, config config.AuthConfig) (*OauthService, error SameSite: http.SameSiteLaxMode, } - service := OauthService{ - Store: store, + service := Oauth{ + Store: store, Oauth2Config: oauth2Config, - Verifier: provider.Verifier(&oidc.Config{ClientID: config.ClientId }), - LogoutUrl: claims.EndSessionEndpoint, + Verifier: provider.Verifier(&oidc.Config{ClientID: config.ClientId}), + LogoutUrl: claims.EndSessionEndpoint, } slog.Info("Created OIDC Config") @@ -91,7 +92,7 @@ func setCallbackCookie(w http.ResponseWriter, r *http.Request, name, value strin http.SetCookie(w, c) } -func (s *OauthService) LoginHandler(w http.ResponseWriter, r *http.Request) { +func (s *Oauth) LoginHandler(w http.ResponseWriter, r *http.Request) { state, err := randString(16) if err != nil { http.Error(w, "Internal error", http.StatusInternalServerError) @@ -111,45 +112,45 @@ func (s *OauthService) LoginHandler(w http.ResponseWriter, r *http.Request) { } type UserClaims struct { - Subject string - Email string - Name string + Subject string + Email string + Name string } -func (s *OauthService) CallbackHandler(onSuccess func(w http.ResponseWriter, r *http.Request, claims UserClaims)) http.HandlerFunc { +func (s *Oauth) CallbackHandler(onSuccess func(w http.ResponseWriter, r *http.Request, claims UserClaims)) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { - stateCookie, err := r.Cookie("state") - if err != nil || r.URL.Query().Get("state") != stateCookie.Value { - http.Error(w, "State invalid", http.StatusBadRequest) - return - } - - oauth2Token, err := s.Oauth2Config.Exchange(r.Context(), r.URL.Query().Get("code")) - if err != nil { - http.Error(w, "Failed to exchange token", http.StatusInternalServerError) - return - } - - rawIDToken, ok := oauth2Token.Extra("id_token").(string) - if !ok { - http.Error(w, "No id_token", http.StatusInternalServerError) - return - } - - idToken, err := s.Verifier.Verify(r.Context(), rawIDToken) - if err != nil { - http.Error(w, "Failed to verify ID token", http.StatusInternalServerError) - return - } - - nonceCookie, _ := r.Cookie("nonce") - if idToken.Nonce != nonceCookie.Value { - http.Error(w, "Nonce invalid", http.StatusBadRequest) - return - } + stateCookie, err := r.Cookie("state") + if err != nil || r.URL.Query().Get("state") != stateCookie.Value { + http.Error(w, "State invalid", http.StatusBadRequest) + return + } + + oauth2Token, err := s.Oauth2Config.Exchange(r.Context(), r.URL.Query().Get("code")) + if err != nil { + http.Error(w, "Failed to exchange token", http.StatusInternalServerError) + return + } + + rawIDToken, ok := oauth2Token.Extra("id_token").(string) + if !ok { + http.Error(w, "No id_token", http.StatusInternalServerError) + return + } + + idToken, err := s.Verifier.Verify(r.Context(), rawIDToken) + if err != nil { + http.Error(w, "Failed to verify ID token", http.StatusInternalServerError) + return + } + + nonceCookie, _ := r.Cookie("nonce") + if idToken.Nonce != nonceCookie.Value { + http.Error(w, "Nonce invalid", http.StatusBadRequest) + return + } var claims struct { - Subject string `json:"sub"` + Subject string `json:"sub"` Email string `json:"email"` Name string `json:"name"` } @@ -159,7 +160,7 @@ func (s *OauthService) CallbackHandler(onSuccess func(w http.ResponseWriter, r * } http.SetCookie(w, &http.Cookie{Name: "state", MaxAge: -1, Path: "/"}) - http.SetCookie(w, &http.Cookie{Name: "nonce", MaxAge: -1, Path: "/"}) + http.SetCookie(w, &http.Cookie{Name: "nonce", MaxAge: -1, Path: "/"}) onSuccess(w, r, UserClaims{ Subject: claims.Subject, @@ -168,3 +169,27 @@ func (s *OauthService) CallbackHandler(onSuccess func(w http.ResponseWriter, r * }) } } + +func (o *Oauth) LogoutHandler(w http.ResponseWriter, r *http.Request) { + session, _ := o.Store.Get(r, "auth-session") + session.Options.MaxAge = -1 + session.Save(r, w) + + if o.LogoutUrl == "" { + http.Redirect(w, r, "/", http.StatusFound) + return + } + + target, err := url.Parse(o.LogoutUrl) + if err != nil { + http.Redirect(w, r, "/", http.StatusFound) + return + } + + q := target.Query() + q.Set("post_logout_redirect_uri", "http://localhost:8080/") + q.Set("client_id", o.Oauth2Config.ClientID) + target.RawQuery = q.Encode() + + http.Redirect(w, r, target.String(), http.StatusFound) +} diff --git a/backend/internal/radio/auth_handlers.go b/backend/internal/radio/auth_handlers.go deleted file mode 100644 index 79e6fa6..0000000 --- a/backend/internal/radio/auth_handlers.go +++ /dev/null @@ -1,85 +0,0 @@ -package radio - -import ( - "context" - "log/slog" - "net/http" - "net/url" - - "echsen.club/radio/internal/auth" -) - -func (h *Handler) AuthMiddleware(next http.Handler) http.Handler { - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path == "/auth/login" || r.URL.Path == "/auth/callback" { - next.ServeHTTP(w, r) - return - } - - session, _ := h.service.oauth.Store.Get(r, "auth-session") - userID, ok := session.Values["user_id"].(string) - - if !ok || userID == "" { - http.Redirect(w, r, "/auth/login", http.StatusTemporaryRedirect) - return - } - - user, err := h.service.GetUserByID(r.Context(), userID) - if err != nil { - slog.Warn("Session user not found in DB", "user_id", userID) - http.Redirect(w, r, "/auth/login", http.StatusTemporaryRedirect) - return - } - - ctx := context.WithValue(r.Context(), userKey, user) - - next.ServeHTTP(w, r.WithContext(ctx)) - }) -} - -func (h *Handler) handleAuthSuccess(w http.ResponseWriter, r *http.Request, claims auth.UserClaims) { - userID, err := h.service.UpsertUser(r.Context(), claims.Subject, claims.Email, claims.Name) - // upserting on every endpoint may be a little wasteful but this app is made for a versy specific and tiny - // audience so it shouldnt matter. If it does end up mattering at some point you know what to change :) - if err != nil { - slog.Error("Database sync failed", "error", err) - http.Error(w, "Failed to sync user", http.StatusInternalServerError) - return - } - - session, _ := h.service.oauth.Store.Get(r, "auth-session") - session.Values["user_id"] = userID.String() - session.Values["authenticated"] = true - - if err := session.Save(r, w); err != nil { - slog.Error("Could not save session", "error", err) - http.Error(w, "Internal error", http.StatusInternalServerError) - return - } - - http.Redirect(w, r, "/", http.StatusFound) -} - -func (h *Handler) LogoutHandler(w http.ResponseWriter, r *http.Request) { - session, _ := h.service.oauth.Store.Get(r, "auth-session") - session.Options.MaxAge = -1 - session.Save(r, w) - - if h.service.oauth.LogoutUrl == "" { - http.Redirect(w, r, "/", http.StatusFound) - return - } - - target, err := url.Parse(h.service.oauth.LogoutUrl) - if err != nil { - http.Redirect(w, r, "/", http.StatusFound) - return - } - - q := target.Query() - q.Set("post_logout_redirect_uri", "http://localhost:8080/") - q.Set("client_id", h.service.oauth.Oauth2Config.ClientID) - target.RawQuery = q.Encode() - - http.Redirect(w, r, target.String(), http.StatusFound) -} diff --git a/backend/internal/radio/http_handlers.go b/backend/internal/radio/http_handlers.go deleted file mode 100644 index 8d7448d..0000000 --- a/backend/internal/radio/http_handlers.go +++ /dev/null @@ -1,60 +0,0 @@ -package radio - -import ( - "context" - "fmt" - "log/slog" - "net/http" - - "connectrpc.com/connect" - "connectrpc.com/validate" - - radiov1 "echsen.club/radio/gen/radio/v1" - "echsen.club/radio/gen/radio/v1/radiov1connect" -) - -type Handler struct { - service *Service -} - -func NewHandler(svc *Service) *Handler { - return &Handler{ - service: svc, - } -} - -func (h *Handler) RegisterHandlers(mux *http.ServeMux) { - mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != "/" { - http.NotFound(w, r) - return - } - w.Write([]byte(`

Welcome to Radio

Login with OIDC`)) - }) - mux.HandleFunc("/profile", h.ProfileHandler) - mux.HandleFunc("/auth/login", h.service.oauth.LoginHandler) - mux.HandleFunc("/auth/callback", h.service.oauth.CallbackHandler(h.handleAuthSuccess)) - - mux.Handle(radiov1connect.NewChannelServiceHandler(h, connect.WithInterceptors(validate.NewInterceptor()))) -} - -func (h *Handler) CreateChannel(ctx context.Context, req *radiov1.CreateChannelRequest) (*radiov1.CreateChannelResponse, error) { - channel := &radiov1.Channel{Id: "1", Frequency: "97.1", Description: ""} - return &radiov1.CreateChannelResponse{Channel: channel}, nil -} - -func (h *Handler) DeleteChannel(ctx context.Context, req *radiov1.DeleteChannelRequest) (*radiov1.DeleteChannelResponse, error) { - return &radiov1.DeleteChannelResponse{}, nil -} - -func (h *Handler) ProfileHandler(w http.ResponseWriter, r *http.Request) { - user, ok := UserFromContext(r.Context()) - if !ok { - slog.Error("User missing from context in protected route") - http.Error(w, "Unauthorized", http.StatusUnauthorized) - return - } - - fmt.Fprintf(w, "

Profile

Logged in as: %v

Home", user.SubjectID) -} - diff --git a/backend/internal/radio/service.go b/backend/internal/radio/service.go deleted file mode 100644 index cbcc672..0000000 --- a/backend/internal/radio/service.go +++ /dev/null @@ -1,26 +0,0 @@ -package radio - -import ( - "database/sql" - "log/slog" - - "echsen.club/radio/internal/auth" - "echsen.club/radio/internal/config" -) - -type Service struct { - db *sql.DB - oauth *auth.OauthService -} - -func NewService(db *sql.DB, config config.Config) *Service { - oauthService, err := auth.NewOauthService(db, config.AuthConfig) - if err != nil { - slog.Error("Error while creating oidc service", "error", err) - } - - return &Service{ - db: db, - oauth: oauthService, - } -} diff --git a/backend/internal/radio/user.go b/backend/internal/radio/user.go deleted file mode 100644 index 77fd1a7..0000000 --- a/backend/internal/radio/user.go +++ /dev/null @@ -1,57 +0,0 @@ -package radio - -import ( - "context" - "fmt" - - "github.com/google/uuid" -) - -type contextKey string -const userKey contextKey = "user" - -type User struct { - ID uuid.UUID - SubjectID string - Email string - DisplayName string -} - -func (s *Service) GetUserByID(ctx context.Context, id string) (*User, error) { - var u User - query := `SELECT id, subject_id, email, display_name FROM users WHERE id = $1` - - err := s.db.QueryRowContext(ctx, query, id).Scan(&u.ID, &u.SubjectID, &u.Email, &u.DisplayName) - if err != nil { - return nil, err - } - return &u, nil -} - -func UserFromContext(ctx context.Context) (*User, bool) { - u, ok := ctx.Value(userKey).(*User) - return u, ok -} - -func (s *Service) UpsertUser(ctx context.Context, subjectID, email, name string) (uuid.UUID, error) { - var id uuid.UUID - - // ON CONFLICT (subject_id) tells Postgres: "If this OIDC user already exists, - // just update their info and last_login instead of throwing an error." - query := ` - INSERT INTO users (subject_id, email, display_name, last_login) - VALUES ($1, $2, $3, NOW()) - ON CONFLICT (subject_id) DO UPDATE - SET last_login = NOW(), - display_name = EXCLUDED.display_name, - email = EXCLUDED.email - RETURNING id; - ` - - err := s.db.QueryRowContext(ctx, query, subjectID, email, name).Scan(&id) - if err != nil { - return uuid.Nil, fmt.Errorf("failed to upsert user: %w", err) - } - - return id, nil -} diff --git a/backend/internal/state/middleware.go b/backend/internal/state/middleware.go new file mode 100644 index 0000000..e1d44c9 --- /dev/null +++ b/backend/internal/state/middleware.go @@ -0,0 +1,87 @@ +package state + +import ( + "context" + "log/slog" + "net/http" + "strings" + + "echsen.club/radio/internal/oauth" +) + +func (state *State) AuthMiddleware(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/auth/login" || r.URL.Path == "/auth/callback" { + next.ServeHTTP(w, r) + return + } + unauthorized := func() { + // If the request is asking for JSON or sending JSON, it's a "Data" request + // This covers Connect-over-HTTP, REST, and standard AJAX + isDataRequest := r.Header.Get("Accept") == "application/json" || + strings.Contains(r.Header.Get("Content-Type"), "application/json") || + strings.Contains(r.Header.Get("Content-Type"), "application/connect") + + if isDataRequest { + w.WriteHeader(http.StatusUnauthorized) // 401 + return + } + + http.Redirect(w, r, "/auth/login", http.StatusTemporaryRedirect) // 307 + } + + session, err := state.oauth.Store.Get(r, "auth-session") + if err != nil { + slog.Error("Failed to get session", "error", err) + unauthorized() + return + } + + if session == nil { + slog.Error("Session is nil") + unauthorized() + return + } + + userID, ok := session.Values["user_id"].(string) + + if !ok || userID == "" { + unauthorized() + return + } + + user, err := state.GetUserByID(r.Context(), userID) + if err != nil { + slog.Warn("Session user not found in DB", "user_id", userID) + unauthorized() + return + } + + ctx := context.WithValue(r.Context(), userKey, user) + + next.ServeHTTP(w, r.WithContext(ctx)) + }) +} + +func (state *State) handleAuthSuccess(w http.ResponseWriter, r *http.Request, claims oauth.UserClaims) { + userID, err := state.UpsertUser(r.Context(), claims.Subject, claims.Email, claims.Name) + // upserting on every endpoint may be a little wasteful but this app is made for a versy specific and tiny + // audience so it shouldnt matter. If it does end up mattering at some point you know what to change :) + if err != nil { + slog.Error("Database sync failed", "error", err) + http.Error(w, "Failed to sync user", http.StatusInternalServerError) + return + } + + session, _ := state.oauth.Store.Get(r, "auth-session") + session.Values["user_id"] = userID.String() + session.Values["authenticated"] = true + + if err := session.Save(r, w); err != nil { + slog.Error("Could not save session", "error", err) + http.Error(w, "Internal error", http.StatusInternalServerError) + return + } + + http.Redirect(w, r, "/", http.StatusFound) +} diff --git a/backend/internal/state/state.go b/backend/internal/state/state.go new file mode 100644 index 0000000..f474d94 --- /dev/null +++ b/backend/internal/state/state.go @@ -0,0 +1,58 @@ +package state + +import ( + "database/sql" + "log" + "net/http" + + "connectrpc.com/connect" + "connectrpc.com/validate" + "echsen.club/radio/gen/radio/v1/radiov1connect" + "echsen.club/radio/internal/channel" + "echsen.club/radio/internal/config" + "echsen.club/radio/internal/oauth" +) + +type State struct { + db *sql.DB + oauth *oauth.Oauth + mux *http.ServeMux + channelHandler *channel.ChannelHandler +} + +func New(db *sql.DB, config config.Config) *State { + oauth, err := oauth.New(db, config.AuthConfig) + if err != nil { + log.Fatalf("Failed to initialize OAuth: %v", err) + } + + state := &State{ + db: db, + oauth: oauth, + mux: http.NewServeMux(), + } + + state.channelHandler = channel.NewHandler(state) + + state.RegisterHandlers() + + return state +} + +func (state *State) Handler() http.Handler { + return state.AuthMiddleware(state.mux) +} + +func (state *State) RegisterHandlers() { + state.mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/" { + http.NotFound(w, r) + return + } + w.Write([]byte(`

Welcome to Radio

Login with OIDC`)) + }) + state.mux.HandleFunc("/auth/login", state.oauth.LoginHandler) + state.mux.HandleFunc("/auth/callback", state.oauth.CallbackHandler(state.handleAuthSuccess)) + + state.mux.Handle(radiov1connect.NewChannelServiceHandler(state.channelHandler, connect.WithInterceptors(validate.NewInterceptor()))) +} diff --git a/backend/internal/state/user.go b/backend/internal/state/user.go new file mode 100644 index 0000000..84de725 --- /dev/null +++ b/backend/internal/state/user.go @@ -0,0 +1,52 @@ +package state + +import ( + "context" + "fmt" + + "echsen.club/radio/internal/db" + "github.com/google/uuid" +) + +type contextKey string + +const userKey contextKey = "user" + +func (state *State) GetUserByID(ctx context.Context, id string) (*db.User, error) { + var u db.User + query := `SELECT id, subject_id, email, display_name FROM users WHERE id = $1` + + err := state.db.QueryRowContext(ctx, query, id).Scan(&u.ID, &u.SubjectID, &u.Email, &u.DisplayName) + if err != nil { + return nil, err + } + return &u, nil +} + +func (state *State) UserFromContext(ctx context.Context) (*db.User, bool) { + u, ok := ctx.Value(userKey).(*db.User) + return u, ok +} + +func (state *State) UpsertUser(ctx context.Context, subjectID, email, name string) (uuid.UUID, error) { + var id uuid.UUID + + // ON CONFLICT (subject_id) tells Postgres: "If this OIDC user already exists, + // just update their info and last_login instead of throwing an error." + query := ` + INSERT INTO users (subject_id, email, display_name, last_login) + VALUES ($1, $2, $3, NOW()) + ON CONFLICT (subject_id) DO UPDATE + SET last_login = NOW(), + display_name = EXCLUDED.display_name, + email = EXCLUDED.email + RETURNING id; + ` + + err := state.db.QueryRowContext(ctx, query, subjectID, email, name).Scan(&id) + if err != nil { + return uuid.Nil, fmt.Errorf("failed to upsert user: %w", err) + } + + return id, nil +} diff --git a/proto/radio/v1/radio.proto b/proto/radio/v1/radio.proto index 831567f..fe890d0 100644 --- a/proto/radio/v1/radio.proto +++ b/proto/radio/v1/radio.proto @@ -29,18 +29,3 @@ service ChannelService { rpc CreateChannel(CreateChannelRequest) returns (CreateChannelResponse); rpc DeleteChannel(DeleteChannelRequest) returns (DeleteChannelResponse); } - -message GreetRequest { - string name = 1 [(buf.validate.field).string = { - min_len: 1, - max_len: 50, - }]; -} - -message GreetResponse { - string greeting = 1; -} - -service GreetService { - rpc Greet(GreetRequest) returns (GreetResponse) {} -}