Refactoring of authentification

Signed-off-by: Bryan Frimin <bryan@getprobo.com>
This commit is contained in:
Bryan Frimin
2025-09-08 10:02:52 +02:00
committed by Sacha Al Himdani
parent afa7e4fdd5
commit 0f96b8518f
53 changed files with 7997 additions and 1631 deletions

View File

@@ -20,13 +20,14 @@ import (
"time"
"github.com/getprobo/probo/pkg/auth"
"github.com/getprobo/probo/pkg/authz"
"github.com/getprobo/probo/pkg/connector"
"github.com/getprobo/probo/pkg/probo"
"github.com/getprobo/probo/pkg/saferedirect"
console_v1 "github.com/getprobo/probo/pkg/server/api/console/v1"
trust_v1 "github.com/getprobo/probo/pkg/server/api/trust/v1"
"github.com/getprobo/probo/pkg/trust"
"github.com/getprobo/probo/pkg/usrmgr"
"github.com/go-chi/chi/v5"
"github.com/go-chi/cors"
"go.gearno.de/kit/httpserver"
@@ -55,9 +56,10 @@ type (
Config struct {
AllowedOrigins []string
Probo *probo.Service
Usrmgr *usrmgr.Service
Auth *auth.Service
Authz *authz.Service
Trust *trust.Service
Auth ConsoleAuthConfig
ConsoleAuth ConsoleAuthConfig
TrustAuth TrustAuthConfig
ConnectorRegistry *connector.ConnectorRegistry
SafeRedirect *saferedirect.SafeRedirect
@@ -72,8 +74,9 @@ type (
)
var (
ErrMissingProboService = errors.New("server configuration requires a valid probo.Service instance")
ErrMissingUsrmgrService = errors.New("server configuration requires a valid usrmgr.Service instance")
ErrMissingProboService = errors.New("server configuration requires a valid probo.Service instance")
ErrMissingAuthService = errors.New("server configuration requires a valid auth.Service instance")
ErrMissingAuthzService = errors.New("server configuration requires a valid authz.Service instance")
)
func methodNotAllowed(w http.ResponseWriter, r *http.Request) {
@@ -105,20 +108,25 @@ func NewServer(cfg Config) (*Server, error) {
return nil, ErrMissingProboService
}
if cfg.Usrmgr == nil {
return nil, ErrMissingUsrmgrService
if cfg.Auth == nil {
return nil, ErrMissingAuthService
}
if cfg.Authz == nil {
return nil, ErrMissingAuthzService
}
// Create trust API handler once
trustAPIHandler := trust_v1.NewMux(
cfg.Logger.Named("trust.v1"),
cfg.Usrmgr,
cfg.Auth,
cfg.Authz,
cfg.Trust,
console_v1.AuthConfig{
CookieName: cfg.Auth.CookieName,
CookieDomain: cfg.Auth.CookieDomain,
SessionDuration: cfg.Auth.SessionDuration,
CookieSecret: cfg.Auth.CookieSecret,
CookieName: cfg.ConsoleAuth.CookieName,
CookieDomain: cfg.ConsoleAuth.CookieDomain,
SessionDuration: cfg.ConsoleAuth.SessionDuration,
CookieSecret: cfg.ConsoleAuth.CookieSecret,
},
trust_v1.TrustAuthConfig{
CookieName: cfg.TrustAuth.CookieName,
@@ -175,12 +183,13 @@ func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) {
console_v1.NewMux(
s.cfg.Logger.Named("console.v1"),
s.cfg.Probo,
s.cfg.Usrmgr,
s.cfg.Auth,
s.cfg.Authz,
console_v1.AuthConfig{
CookieName: s.cfg.Auth.CookieName,
CookieDomain: s.cfg.Auth.CookieDomain,
SessionDuration: s.cfg.Auth.SessionDuration,
CookieSecret: s.cfg.Auth.CookieSecret,
CookieName: s.cfg.ConsoleAuth.CookieName,
CookieDomain: s.cfg.ConsoleAuth.CookieDomain,
SessionDuration: s.cfg.ConsoleAuth.SessionDuration,
CookieSecret: s.cfg.ConsoleAuth.CookieSecret,
},
s.cfg.ConnectorRegistry,
s.cfg.SafeRedirect,

View File

@@ -19,7 +19,7 @@ import (
"fmt"
"net/http"
"github.com/getprobo/probo/pkg/usrmgr"
"github.com/getprobo/probo/pkg/auth"
"go.gearno.de/kit/httpserver"
)
@@ -33,7 +33,7 @@ type (
}
)
func ForgetPasswordHandler(usrmgrSvc *usrmgr.Service, authCfg AuthConfig) http.HandlerFunc {
func ForgetPasswordHandler(authSvc *auth.Service, authCfg AuthConfig) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
var req ForgetPasswordRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
@@ -41,7 +41,7 @@ func ForgetPasswordHandler(usrmgrSvc *usrmgr.Service, authCfg AuthConfig) http.H
return
}
err := usrmgrSvc.ForgetPassword(r.Context(), req.Email)
err := authSvc.ForgetPassword(r.Context(), req.Email)
if err != nil {
// For security reasons, we don't expose whether an email exists or not
httpserver.RenderError(w, http.StatusInternalServerError, fmt.Errorf("cannot process request: %w", err))

View File

@@ -16,11 +16,14 @@ package console_v1
import (
"encoding/json"
"errors"
"fmt"
"net/http"
"github.com/getprobo/probo/pkg/probo"
"github.com/getprobo/probo/pkg/usrmgr"
"github.com/getprobo/probo/pkg/auth"
"github.com/getprobo/probo/pkg/authz"
"github.com/getprobo/probo/pkg/coredata"
"github.com/getprobo/probo/pkg/statelesstoken"
"go.gearno.de/kit/httpserver"
)
@@ -34,7 +37,7 @@ type (
}
)
func InvitationConfirmationHandler(usrmgrSvc *usrmgr.Service, proboSvc *probo.Service, authCfg AuthConfig) http.HandlerFunc {
func InvitationConfirmationHandler(authSvc *auth.Service, authzSvc *authz.Service, authCfg AuthConfig) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
var req InvitationConfirmationRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
@@ -42,7 +45,32 @@ func InvitationConfirmationHandler(usrmgrSvc *usrmgr.Service, proboSvc *probo.Se
return
}
_, err := usrmgrSvc.ConfirmInvitation(r.Context(), req.Token, req.Password)
payload, err := statelesstoken.ValidateToken[coredata.InvitationData](
authCfg.CookieSecret,
authz.TokenTypeOrganizationInvitation,
req.Token,
)
if err != nil {
httpserver.RenderError(w, http.StatusBadRequest, fmt.Errorf("invalid invitation token: %w", err))
return
}
user, _, err := authSvc.SignUp(r.Context(), payload.Data.Email, req.Password, payload.Data.FullName)
if err != nil {
var errUserAlreadyExists *auth.ErrUserAlreadyExists
if errors.As(err, &errUserAlreadyExists) {
user, err = authSvc.GetUserByEmail(r.Context(), payload.Data.Email)
if err != nil {
httpserver.RenderError(w, http.StatusInternalServerError, fmt.Errorf("failed to load existing user: %w", err))
return
}
} else {
httpserver.RenderError(w, http.StatusInternalServerError, fmt.Errorf("failed to create user: %w", err))
return
}
}
err = authzSvc.AcceptInvitation(r.Context(), req.Token, user.ID)
if err != nil {
httpserver.RenderError(w, http.StatusInternalServerError, err)
return

View File

@@ -21,7 +21,7 @@ import (
"errors"
"github.com/getprobo/probo/pkg/usrmgr"
"github.com/getprobo/probo/pkg/auth"
"go.gearno.de/kit/httpserver"
)
@@ -36,7 +36,7 @@ type (
}
)
func ResetPasswordHandler(usrmgrSvc *usrmgr.Service, authCfg AuthConfig) http.HandlerFunc {
func ResetPasswordHandler(authSvc *auth.Service, authCfg AuthConfig) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
var req ResetPasswordRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
@@ -44,10 +44,10 @@ func ResetPasswordHandler(usrmgrSvc *usrmgr.Service, authCfg AuthConfig) http.Ha
return
}
err := usrmgrSvc.ResetPassword(r.Context(), req.Token, req.Password)
err := authSvc.ResetPassword(r.Context(), req.Token, req.Password)
if err != nil {
var invalidPasswordErr *usrmgr.ErrInvalidPassword
var invalidTokenErr *usrmgr.ErrInvalidTokenType
var invalidPasswordErr *auth.ErrInvalidPassword
var invalidTokenErr *auth.ErrInvalidTokenType
if errors.As(err, &invalidPasswordErr) {
httpserver.RenderError(w, http.StatusBadRequest, err)

View File

@@ -29,6 +29,8 @@ import (
"github.com/99designs/gqlgen/graphql/handler/extension"
"github.com/99designs/gqlgen/graphql/handler/transport"
"github.com/99designs/gqlgen/graphql/playground"
"github.com/getprobo/probo/pkg/auth"
"github.com/getprobo/probo/pkg/authz"
"github.com/getprobo/probo/pkg/connector"
"github.com/getprobo/probo/pkg/coredata"
"github.com/getprobo/probo/pkg/gid"
@@ -38,7 +40,6 @@ import (
gqlutils "github.com/getprobo/probo/pkg/server/graphql"
"github.com/getprobo/probo/pkg/server/session"
"github.com/getprobo/probo/pkg/statelesstoken"
"github.com/getprobo/probo/pkg/usrmgr"
"github.com/go-chi/chi/v5"
"github.com/vektah/gqlparser/v2/gqlerror"
"go.gearno.de/kit/log"
@@ -54,7 +55,8 @@ type (
Resolver struct {
proboSvc *probo.Service
usrmgrSvc *usrmgr.Service
authSvc *auth.Service
authzSvc *authz.Service
authCfg AuthConfig
customDomainCname string
}
@@ -81,7 +83,8 @@ func UserFromContext(ctx context.Context) *coredata.User {
func NewMux(
logger *log.Logger,
proboSvc *probo.Service,
usrmgrSvc *usrmgr.Service,
authSvc *auth.Service,
authzSvc *authz.Service,
authCfg AuthConfig,
connectorRegistry *connector.ConnectorRegistry,
safeRedirect *saferedirect.SafeRedirect,
@@ -151,14 +154,14 @@ func NewMux(
},
)
r.Post("/auth/register", SignUpHandler(usrmgrSvc, authCfg))
r.Post("/auth/login", SignInHandler(usrmgrSvc, authCfg))
r.Delete("/auth/logout", SignOutHandler(usrmgrSvc, authCfg))
r.Post("/auth/invitation", InvitationConfirmationHandler(usrmgrSvc, proboSvc, authCfg))
r.Post("/auth/forget-password", ForgetPasswordHandler(usrmgrSvc, authCfg))
r.Post("/auth/reset-password", ResetPasswordHandler(usrmgrSvc, authCfg))
r.Post("/auth/register", SignUpHandler(authSvc, authCfg))
r.Post("/auth/login", SignInHandler(authSvc, authCfg))
r.Delete("/auth/logout", SignOutHandler(authSvc, authCfg))
r.Post("/auth/invitation", InvitationConfirmationHandler(authSvc, authzSvc, authCfg))
r.Post("/auth/forget-password", ForgetPasswordHandler(authSvc, authCfg))
r.Post("/auth/reset-password", ResetPasswordHandler(authSvc, authCfg))
r.Get("/connectors/initiate", WithSession(usrmgrSvc, authCfg, func(w http.ResponseWriter, r *http.Request) {
r.Get("/connectors/initiate", WithSession(authSvc, authzSvc, authCfg, func(w http.ResponseWriter, r *http.Request) {
connectorID := r.URL.Query().Get("connector_id")
organizationID, err := gid.ParseGID(r.URL.Query().Get("organization_id"))
if err != nil {
@@ -175,7 +178,7 @@ func NewMux(
http.Redirect(w, r, redirectURL, http.StatusSeeOther)
}))
r.Get("/connectors/complete", WithSession(usrmgrSvc, authCfg, func(w http.ResponseWriter, r *http.Request) {
r.Get("/connectors/complete", WithSession(authSvc, authzSvc, authCfg, func(w http.ResponseWriter, r *http.Request) {
connectorID := r.URL.Query().Get("connector_id")
organizationID, err := gid.ParseGID(r.URL.Query().Get("organization_id"))
if err != nil {
@@ -206,19 +209,20 @@ func NewMux(
}))
r.Get("/", playground.Handler("GraphQL", "/api/console/v1/query"))
r.Post("/query", graphqlHandler(logger, proboSvc, usrmgrSvc, authCfg, customDomainCname))
r.Post("/query", graphqlHandler(logger, proboSvc, authSvc, authzSvc, authCfg, customDomainCname))
return r
}
func graphqlHandler(logger *log.Logger, proboSvc *probo.Service, usrmgrSvc *usrmgr.Service, authCfg AuthConfig, customDomainCname string) http.HandlerFunc {
func graphqlHandler(logger *log.Logger, proboSvc *probo.Service, authSvc *auth.Service, authzSvc *authz.Service, authCfg AuthConfig, customDomainCname string) http.HandlerFunc {
var mb int64 = 1 << 20
es := schema.NewExecutableSchema(
schema.Config{
Resolvers: &Resolver{
proboSvc: proboSvc,
usrmgrSvc: usrmgrSvc,
authSvc: authSvc,
authzSvc: authzSvc,
authCfg: authCfg,
customDomainCname: customDomainCname,
},
@@ -259,10 +263,10 @@ func graphqlHandler(logger *log.Logger, proboSvc *probo.Service, usrmgrSvc *usrm
},
)
return WithSession(usrmgrSvc, authCfg, srv.ServeHTTP)
return WithSession(authSvc, authzSvc, authCfg, srv.ServeHTTP)
}
func WithSession(usrmgrSvc *usrmgr.Service, authCfg AuthConfig, next http.HandlerFunc) http.HandlerFunc {
func WithSession(authSvc *auth.Service, authzSvc *authz.Service, authCfg AuthConfig, next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
@@ -289,7 +293,7 @@ func WithSession(usrmgrSvc *usrmgr.Service, authCfg AuthConfig, next http.Handle
},
}
authResult := session.TryAuth(ctx, w, r, usrmgrSvc, sessionAuthCfg, errorHandler)
authResult := session.TryAuth(ctx, w, r, authSvc, authzSvc, sessionAuthCfg, errorHandler)
if authResult == nil {
next(w, r)
return
@@ -302,7 +306,7 @@ func WithSession(usrmgrSvc *usrmgr.Service, authCfg AuthConfig, next http.Handle
next(w, r.WithContext(ctx))
// Update session after the handler completes
if err := usrmgrSvc.UpdateSession(ctx, authResult.Session); err != nil {
if _, err := authSvc.UpdateSession(ctx, authResult.Session.ID); err != nil {
panic(fmt.Errorf("failed to update session: %w", err))
}
}

View File

@@ -1209,6 +1209,54 @@ enum SnapshotOrderField
)
}
enum MembershipOrderField
@goModel(model: "github.com/getprobo/probo/pkg/coredata.MembershipOrderField") {
FULL_NAME
@goEnum(
value: "github.com/getprobo/probo/pkg/coredata.MembershipOrderFieldFullName"
)
EMAIL_ADDRESS
@goEnum(
value: "github.com/getprobo/probo/pkg/coredata.MembershipOrderFieldEmailAddress"
)
ROLE
@goEnum(
value: "github.com/getprobo/probo/pkg/coredata.MembershipOrderFieldRole"
)
CREATED_AT
@goEnum(
value: "github.com/getprobo/probo/pkg/coredata.MembershipOrderFieldCreatedAt"
)
}
enum InvitationOrderField
@goModel(model: "github.com/getprobo/probo/pkg/coredata.InvitationOrderField") {
FULL_NAME
@goEnum(
value: "github.com/getprobo/probo/pkg/coredata.InvitationOrderFieldFullName"
)
EMAIL
@goEnum(
value: "github.com/getprobo/probo/pkg/coredata.InvitationOrderFieldEmail"
)
ROLE
@goEnum(
value: "github.com/getprobo/probo/pkg/coredata.InvitationOrderFieldRole"
)
CREATED_AT
@goEnum(
value: "github.com/getprobo/probo/pkg/coredata.InvitationOrderFieldCreatedAt"
)
EXPIRES_AT
@goEnum(
value: "github.com/getprobo/probo/pkg/coredata.InvitationOrderFieldExpiresAt"
)
ACCEPTED_AT
@goEnum(
value: "github.com/getprobo/probo/pkg/coredata.InvitationOrderFieldAcceptedAt"
)
}
# Input Types
input UserOrder
@goModel(
@@ -1404,6 +1452,19 @@ input SnapshotOrder
field: SnapshotOrderField!
}
input MembershipOrder
@goModel(
model: "github.com/getprobo/probo/pkg/server/api/console/v1/types.MembershipOrderBy"
) {
direction: OrderDirection!
field: MembershipOrderField!
}
input InvitationOrder {
direction: OrderDirection!
field: InvitationOrderField!
}
input DocumentVersionFilter {
status: DocumentStatus
}
@@ -1497,13 +1558,21 @@ type Organization implements Node {
email: String
headquarterAddress: String
users(
memberships(
first: Int
after: CursorKey
last: Int
before: CursorKey
orderBy: UserOrder
): UserConnection! @goField(forceResolver: true)
orderBy: MembershipOrder
): MembershipConnection! @goField(forceResolver: true)
invitations(
first: Int
after: CursorKey
last: Int
before: CursorKey
orderBy: InvitationOrder
): InvitationConnection! @goField(forceResolver: true)
connectors(
first: Int
@@ -1671,6 +1740,27 @@ type User implements Node {
people(organizationId: ID!): People @goField(forceResolver: true)
}
type Membership implements Node {
id: ID!
userID: ID!
organizationID: ID!
role: String!
fullName: String!
emailAddress: String!
createdAt: Datetime!
updatedAt: Datetime!
}
type Invitation implements Node {
id: ID!
email: String!
fullName: String!
role: String!
expiresAt: Datetime!
acceptedAt: Datetime
createdAt: Datetime!
}
type Connector implements Node {
id: ID!
name: String!
@@ -2305,10 +2395,22 @@ type TrustCenterReferenceEdge {
}
type UserConnection {
totalCount: Int! @goField(forceResolver: true)
edges: [UserEdge!]!
pageInfo: PageInfo!
}
type MembershipConnection {
totalCount: Int! @goField(forceResolver: true)
edges: [MembershipEdge!]!
pageInfo: PageInfo!
}
type MembershipEdge {
cursor: CursorKey!
node: Membership!
}
type UserEdge {
cursor: CursorKey!
node: User!
@@ -2607,6 +2709,17 @@ type File {
updatedAt: Datetime!
}
type InvitationConnection {
totalCount: Int! @goField(forceResolver: true)
edges: [InvitationEdge!]!
pageInfo: PageInfo!
}
type InvitationEdge {
cursor: CursorKey!
node: Invitation!
}
# Root Types
type Query {
node(id: ID!): Node!
@@ -2667,7 +2780,8 @@ type Mutation {
# User mutations
confirmEmail(input: ConfirmEmailInput!): ConfirmEmailPayload!
inviteUser(input: InviteUserInput!): InviteUserPayload!
removeUser(input: RemoveUserInput!): RemoveUserPayload!
deleteInvitation(input: DeleteInvitationInput!): DeleteInvitationPayload!
removeMember(input: RemoveMemberInput!): RemoveMemberPayload!
# People mutations
createPeople(input: CreatePeopleInput!): CreatePeoplePayload!
@@ -3426,9 +3540,13 @@ input InviteUserInput {
createPeople: Boolean!
}
input RemoveUserInput {
input DeleteInvitationInput {
invitationId: ID!
}
input RemoveMemberInput {
organizationId: ID!
userId: ID!
memberId: ID!
}
input CreateControlInput {
@@ -3945,10 +4063,14 @@ type ConfirmEmailPayload {
}
type InviteUserPayload {
success: Boolean!
invitationEdge: InvitationEdge!
}
type RemoveUserPayload {
type DeleteInvitationPayload {
deletedInvitationId: ID!
}
type RemoveMemberPayload {
success: Boolean!
}

File diff suppressed because it is too large Load Diff

View File

@@ -23,7 +23,7 @@ import (
"github.com/getprobo/probo/pkg/gid"
"github.com/getprobo/probo/pkg/securecookie"
"github.com/getprobo/probo/pkg/usrmgr"
"github.com/getprobo/probo/pkg/auth"
"go.gearno.de/kit/httpserver"
)
@@ -46,7 +46,7 @@ type (
}
)
func SignInHandler(usrmgrSvc *usrmgr.Service, authCfg AuthConfig) http.HandlerFunc {
func SignInHandler(authSvc *auth.Service, authCfg AuthConfig) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
var req SignInRequest
@@ -55,9 +55,9 @@ func SignInHandler(usrmgrSvc *usrmgr.Service, authCfg AuthConfig) http.HandlerFu
return
}
user, session, err := usrmgrSvc.SignIn(r.Context(), req.Email, req.Password)
session, user, err := authSvc.SignIn(r.Context(), req.Email, req.Password)
if err != nil {
var ErrInvalidCredentials *usrmgr.ErrInvalidCredentials
var ErrInvalidCredentials *auth.ErrInvalidCredentials
if errors.As(err, &ErrInvalidCredentials) {
httpserver.RenderError(w, http.StatusUnauthorized, err)
return

View File

@@ -20,11 +20,11 @@ import (
"github.com/getprobo/probo/pkg/gid"
"github.com/getprobo/probo/pkg/securecookie"
"github.com/getprobo/probo/pkg/usrmgr"
"github.com/getprobo/probo/pkg/auth"
"go.gearno.de/kit/httpserver"
)
func SignOutHandler(usrmgrSvc *usrmgr.Service, authCfg AuthConfig) http.HandlerFunc {
func SignOutHandler(authSvc *auth.Service, authCfg AuthConfig) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
sessionID, err := securecookie.Get(r, securecookie.DefaultConfig(
@@ -42,7 +42,7 @@ func SignOutHandler(usrmgrSvc *usrmgr.Service, authCfg AuthConfig) http.HandlerF
return
}
err = usrmgrSvc.SignOut(r.Context(), gid)
err = authSvc.SignOut(r.Context(), gid)
if err != nil {
panic(fmt.Errorf("cannot sign out: %w", err))
}

View File

@@ -20,9 +20,8 @@ import (
"fmt"
"net/http"
"github.com/getprobo/probo/pkg/coredata"
"github.com/getprobo/probo/pkg/auth"
"github.com/getprobo/probo/pkg/securecookie"
"github.com/getprobo/probo/pkg/usrmgr"
"go.gearno.de/kit/httpserver"
)
@@ -38,7 +37,7 @@ type (
}
)
func SignUpHandler(usrmgrSvc *usrmgr.Service, authCfg AuthConfig) http.HandlerFunc {
func SignUpHandler(authSvc *auth.Service, authCfg AuthConfig) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
var req SignUpRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
@@ -46,20 +45,20 @@ func SignUpHandler(usrmgrSvc *usrmgr.Service, authCfg AuthConfig) http.HandlerFu
return
}
user, session, err := usrmgrSvc.SignUp(
user, session, err := authSvc.SignUp(
r.Context(),
req.Email,
req.Password,
req.FullName,
)
if err != nil {
var errUserAlreadyExists *coredata.ErrUserAlreadyExists
var errUserAlreadyExists *auth.ErrUserAlreadyExists
if errors.As(err, &errUserAlreadyExists) {
httpserver.RenderError(w, http.StatusBadRequest, fmt.Errorf("cannot register user: %w", err))
return
}
var errSignupDisabled *usrmgr.ErrSignupDisabled
var errSignupDisabled *auth.ErrSignupDisabled
if errors.As(err, &errSignupDisabled) {
httpserver.RenderError(w, http.StatusBadRequest, fmt.Errorf("cannot register user: %w", err))
return

View File

@@ -0,0 +1,52 @@
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package types
import (
"github.com/getprobo/probo/pkg/coredata"
"github.com/getprobo/probo/pkg/page"
)
func NewInvitationConnection(p *page.Page[*coredata.Invitation, coredata.InvitationOrderField]) *InvitationConnection {
var edges = make([]*InvitationEdge, len(p.Data))
for i := range edges {
edges[i] = NewInvitationEdge(p.Data[i], p.Cursor.OrderBy.Field)
}
return &InvitationConnection{
Edges: edges,
PageInfo: NewPageInfo(p),
}
}
func NewInvitationEdge(invitation *coredata.Invitation, orderBy coredata.InvitationOrderField) *InvitationEdge {
return &InvitationEdge{
Cursor: invitation.CursorKey(orderBy),
Node: NewInvitation(invitation),
}
}
func NewInvitation(i *coredata.Invitation) *Invitation {
return &Invitation{
ID: i.ID,
Email: i.Email,
FullName: i.FullName,
Role: i.Role,
ExpiresAt: i.ExpiresAt,
AcceptedAt: i.AcceptedAt,
CreatedAt: i.CreatedAt,
}
}

View File

@@ -0,0 +1,57 @@
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package types
import (
"github.com/getprobo/probo/pkg/coredata"
"github.com/getprobo/probo/pkg/page"
)
type (
MembershipOrderBy OrderBy[coredata.MembershipOrderField]
)
func NewMembershipConnection(p *page.Page[*coredata.Membership, coredata.MembershipOrderField]) *MembershipConnection {
var edges = make([]*MembershipEdge, len(p.Data))
for i := range edges {
edges[i] = NewMembershipEdge(p.Data[i], p.Cursor.OrderBy.Field)
}
return &MembershipConnection{
Edges: edges,
PageInfo: NewPageInfo(p),
}
}
func NewMembershipEdge(membership *coredata.Membership, orderBy coredata.MembershipOrderField) *MembershipEdge {
return &MembershipEdge{
Cursor: membership.CursorKey(orderBy),
Node: NewMembership(membership),
}
}
func NewMembership(m *coredata.Membership) *Membership {
return &Membership{
ID: m.ID,
UserID: m.UserID,
OrganizationID: m.OrganizationID,
Role: m.Role,
FullName: m.FullName,
EmailAddress: m.EmailAddress,
CreatedAt: m.CreatedAt,
UpdatedAt: m.UpdatedAt,
}
}

View File

@@ -807,6 +807,14 @@ type DeleteFrameworkPayload struct {
DeletedFrameworkID gid.GID `json:"deletedFrameworkId"`
}
type DeleteInvitationInput struct {
InvitationID gid.GID `json:"invitationId"`
}
type DeleteInvitationPayload struct {
DeletedInvitationID gid.GID `json:"deletedInvitationId"`
}
type DeleteMeasureInput struct {
MeasureID gid.GID `json:"measureId"`
}
@@ -1191,6 +1199,35 @@ type ImportMeasurePayload struct {
MeasureEdges []*MeasureEdge `json:"measureEdges"`
}
type Invitation struct {
ID gid.GID `json:"id"`
Email string `json:"email"`
FullName string `json:"fullName"`
Role string `json:"role"`
ExpiresAt time.Time `json:"expiresAt"`
AcceptedAt *time.Time `json:"acceptedAt,omitempty"`
CreatedAt time.Time `json:"createdAt"`
}
func (Invitation) IsNode() {}
func (this Invitation) GetID() gid.GID { return this.ID }
type InvitationConnection struct {
TotalCount int `json:"totalCount"`
Edges []*InvitationEdge `json:"edges"`
PageInfo *PageInfo `json:"pageInfo"`
}
type InvitationEdge struct {
Cursor page.CursorKey `json:"cursor"`
Node *Invitation `json:"node"`
}
type InvitationOrder struct {
Direction page.OrderDirection `json:"direction"`
Field coredata.InvitationOrderField `json:"field"`
}
type InviteUserInput struct {
OrganizationID gid.GID `json:"organizationId"`
Email string `json:"email"`
@@ -1199,7 +1236,7 @@ type InviteUserInput struct {
}
type InviteUserPayload struct {
Success bool `json:"success"`
InvitationEdge *InvitationEdge `json:"invitationEdge"`
}
type Measure struct {
@@ -1229,6 +1266,31 @@ type MeasureFilter struct {
State *coredata.MeasureState `json:"state,omitempty"`
}
type Membership struct {
ID gid.GID `json:"id"`
UserID gid.GID `json:"userID"`
OrganizationID gid.GID `json:"organizationID"`
Role string `json:"role"`
FullName string `json:"fullName"`
EmailAddress string `json:"emailAddress"`
CreatedAt time.Time `json:"createdAt"`
UpdatedAt time.Time `json:"updatedAt"`
}
func (Membership) IsNode() {}
func (this Membership) GetID() gid.GID { return this.ID }
type MembershipConnection struct {
TotalCount int `json:"totalCount"`
Edges []*MembershipEdge `json:"edges"`
PageInfo *PageInfo `json:"pageInfo"`
}
type MembershipEdge struct {
Cursor page.CursorKey `json:"cursor"`
Node *Membership `json:"node"`
}
type Mutation struct {
}
@@ -1301,7 +1363,8 @@ type Organization struct {
WebsiteURL *string `json:"websiteUrl,omitempty"`
Email *string `json:"email,omitempty"`
HeadquarterAddress *string `json:"headquarterAddress,omitempty"`
Users *UserConnection `json:"users"`
Memberships *MembershipConnection `json:"memberships"`
Invitations *InvitationConnection `json:"invitations"`
Connectors *ConnectorConnection `json:"connectors"`
Frameworks *FrameworkConnection `json:"frameworks"`
Controls *ControlConnection `json:"controls"`
@@ -1424,12 +1487,12 @@ type PublishDocumentVersionPayload struct {
type Query struct {
}
type RemoveUserInput struct {
type RemoveMemberInput struct {
OrganizationID gid.GID `json:"organizationId"`
UserID gid.GID `json:"userId"`
MemberID gid.GID `json:"memberId"`
}
type RemoveUserPayload struct {
type RemoveMemberPayload struct {
Success bool `json:"success"`
}
@@ -2064,8 +2127,9 @@ func (User) IsNode() {}
func (this User) GetID() gid.GID { return this.ID }
type UserConnection struct {
Edges []*UserEdge `json:"edges"`
PageInfo *PageInfo `json:"pageInfo"`
TotalCount int `json:"totalCount"`
Edges []*UserEdge `json:"edges"`
PageInfo *PageInfo `json:"pageInfo"`
}
type UserEdge struct {

View File

@@ -12,6 +12,7 @@ import (
"fmt"
"time"
"github.com/getprobo/probo/pkg/authz"
"github.com/getprobo/probo/pkg/coredata"
"github.com/getprobo/probo/pkg/gid"
"github.com/getprobo/probo/pkg/page"
@@ -890,6 +891,27 @@ func (r *frameworkConnectionResolver) TotalCount(ctx context.Context, obj *types
panic(fmt.Errorf("unsupported resolver: %T", obj.Resolver))
}
// TotalCount is the resolver for the totalCount field.
func (r *invitationConnectionResolver) TotalCount(ctx context.Context, obj *types.InvitationConnection) (int, error) {
currentUser := UserFromContext(ctx)
if currentUser == nil {
return 0, fmt.Errorf("no authenticated user")
}
memberships, err := r.authzSvc.GetAllUserOrganizations(ctx, currentUser.ID)
if err != nil || len(memberships) == 0 {
return 0, fmt.Errorf("user has no organization memberships")
}
orgID := memberships[0].ID
count, err := r.authzSvc.CountOrganizationInvitations(ctx, orgID)
if err != nil {
return 0, fmt.Errorf("failed to count invitations: %w", err)
}
return count, nil
}
// Evidences is the resolver for the evidences field.
func (r *measureResolver) Evidences(ctx context.Context, obj *types.Measure, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.EvidenceOrderBy) (*types.EvidenceConnection, error) {
prb := r.ProboService(ctx, obj.ID.TenantID())
@@ -1028,6 +1050,27 @@ func (r *measureConnectionResolver) TotalCount(ctx context.Context, obj *types.M
panic(fmt.Errorf("unsupported resolver: %T", obj.Resolver))
}
// TotalCount is the resolver for the totalCount field.
func (r *membershipConnectionResolver) TotalCount(ctx context.Context, obj *types.MembershipConnection) (int, error) {
currentUser := UserFromContext(ctx)
if currentUser == nil {
return 0, fmt.Errorf("no authenticated user")
}
memberships, err := r.authzSvc.GetAllUserOrganizations(ctx, currentUser.ID)
if err != nil || len(memberships) == 0 {
return 0, fmt.Errorf("user has no organization memberships")
}
orgID := memberships[0].ID
count, err := r.authzSvc.CountOrganizationMemberships(ctx, orgID)
if err != nil {
return 0, fmt.Errorf("failed to count memberships: %w", err)
}
return count, nil
}
// CreateOrganization is the resolver for the createOrganization field.
func (r *mutationResolver) CreateOrganization(ctx context.Context, input types.CreateOrganizationInput) (*types.CreateOrganizationPayload, error) {
prb := r.proboSvc.WithTenant(gid.NewTenantID())
@@ -1042,10 +1085,11 @@ func (r *mutationResolver) CreateOrganization(ctx context.Context, input types.C
return nil, fmt.Errorf("cannot create organization: %w", err)
}
err = r.usrmgrSvc.EnrollUserInOrganization(
err = r.authzSvc.AddUserToOrganization(
ctx,
UserFromContext(ctx).ID,
organization.ID,
string(authz.RoleMember),
)
if err != nil {
return nil, fmt.Errorf("cannot add user to organization: %w", err)
@@ -1324,7 +1368,7 @@ func (r *mutationResolver) DeleteTrustCenterReference(ctx context.Context, input
// ConfirmEmail is the resolver for the confirmEmail field.
func (r *mutationResolver) ConfirmEmail(ctx context.Context, input types.ConfirmEmailInput) (*types.ConfirmEmailPayload, error) {
err := r.usrmgrSvc.ConfirmEmail(ctx, input.Token)
err := r.authSvc.ConfirmEmail(ctx, input.Token)
if err != nil {
return nil, err
@@ -1337,44 +1381,70 @@ func (r *mutationResolver) ConfirmEmail(ctx context.Context, input types.Confirm
func (r *mutationResolver) InviteUser(ctx context.Context, input types.InviteUserInput) (*types.InviteUserPayload, error) {
user := UserFromContext(ctx)
organizations, err := r.usrmgrSvc.ListOrganizationsForUserID(ctx, user.ID)
organizations, err := r.authzSvc.GetAllUserOrganizations(ctx, user.ID)
if err != nil {
panic(fmt.Errorf("failed to list organizations for user: %w", err))
}
for _, organization := range organizations {
if organization.ID == input.OrganizationID {
createPeople := input.CreatePeople
err := r.usrmgrSvc.InviteUser(ctx, input.OrganizationID, input.FullName, input.Email, createPeople)
invitation, err := r.authzSvc.InviteUserToOrganization(ctx, input.OrganizationID, input.Email, input.FullName, string(authz.RoleMember))
if err != nil {
return nil, err
}
return &types.InviteUserPayload{Success: true}, nil
if input.CreatePeople {
prb := r.ProboService(ctx, input.OrganizationID.TenantID())
_, err := prb.Peoples.Create(ctx, probo.CreatePeopleRequest{
OrganizationID: input.OrganizationID,
FullName: input.FullName,
PrimaryEmailAddress: input.Email,
AdditionalEmailAddresses: []string{},
Kind: coredata.PeopleKindEmployee,
})
if err != nil {
return nil, fmt.Errorf("failed to create people record: %w", err)
}
}
return &types.InviteUserPayload{
InvitationEdge: types.NewInvitationEdge(invitation, coredata.InvitationOrderFieldCreatedAt),
}, nil
}
}
return nil, fmt.Errorf("organization not found")
}
// RemoveUser is the resolver for the removeUser field.
func (r *mutationResolver) RemoveUser(ctx context.Context, input types.RemoveUserInput) (*types.RemoveUserPayload, error) {
// DeleteInvitation is the resolver for the deleteInvitation field.
func (r *mutationResolver) DeleteInvitation(ctx context.Context, input types.DeleteInvitationInput) (*types.DeleteInvitationPayload, error) {
err := r.authzSvc.DeleteInvitation(ctx, input.InvitationID)
if err != nil {
return nil, err
}
return &types.DeleteInvitationPayload{
DeletedInvitationID: input.InvitationID,
}, nil
}
// RemoveMember is the resolver for the removeMember field.
func (r *mutationResolver) RemoveMember(ctx context.Context, input types.RemoveMemberInput) (*types.RemoveMemberPayload, error) {
user := UserFromContext(ctx)
organizations, err := r.usrmgrSvc.ListOrganizationsForUserID(ctx, user.ID)
organizations, err := r.authzSvc.GetAllUserOrganizations(ctx, user.ID)
if err != nil {
panic(fmt.Errorf("failed to list organizations for user: %w", err))
}
for _, organization := range organizations {
if organization.ID == input.OrganizationID {
err := r.usrmgrSvc.RemoveUser(ctx, input.OrganizationID, input.UserID)
err := r.authzSvc.RemoveMemberFromOrganization(ctx, input.OrganizationID, input.MemberID)
if err != nil {
return nil, err
}
return &types.RemoveUserPayload{Success: true}, nil
return &types.RemoveMemberPayload{Success: true}, nil
}
}
@@ -3526,14 +3596,14 @@ func (r *organizationResolver) HorizontalLogoURL(ctx context.Context, obj *types
return prb.Organizations.GenerateHorizontalLogoURL(ctx, obj.ID, 1*time.Hour)
}
// Users is the resolver for the users field.
func (r *organizationResolver) Users(ctx context.Context, obj *types.Organization, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.UserOrderBy) (*types.UserConnection, error) {
pageOrderBy := page.OrderBy[coredata.UserOrderField]{
Field: coredata.UserOrderFieldCreatedAt,
// Memberships is the resolver for the memberships field.
func (r *organizationResolver) Memberships(ctx context.Context, obj *types.Organization, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.MembershipOrderBy) (*types.MembershipConnection, error) {
pageOrderBy := page.OrderBy[coredata.MembershipOrderField]{
Field: coredata.MembershipOrderFieldCreatedAt,
Direction: page.OrderDirectionDesc,
}
if orderBy != nil {
pageOrderBy = page.OrderBy[coredata.UserOrderField]{
pageOrderBy = page.OrderBy[coredata.MembershipOrderField]{
Field: orderBy.Field,
Direction: orderBy.Direction,
}
@@ -3541,12 +3611,35 @@ func (r *organizationResolver) Users(ctx context.Context, obj *types.Organizatio
cursor := types.NewCursor(first, after, last, before, pageOrderBy)
page, err := r.usrmgrSvc.ListUsersForTenant(ctx, obj.ID, cursor)
page, err := r.authzSvc.GetAllOrganizationMemberships(ctx, obj.ID, cursor)
if err != nil {
panic(fmt.Errorf("cannot list users: %w", err))
panic(fmt.Errorf("cannot list memberships: %w", err))
}
return types.NewUserConnection(page), nil
return types.NewMembershipConnection(page), nil
}
// Invitations is the resolver for the invitations field.
func (r *organizationResolver) Invitations(ctx context.Context, obj *types.Organization, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.InvitationOrder) (*types.InvitationConnection, error) {
pageOrderBy := page.OrderBy[coredata.InvitationOrderField]{
Field: coredata.InvitationOrderFieldCreatedAt,
Direction: page.OrderDirectionDesc,
}
if orderBy != nil {
pageOrderBy = page.OrderBy[coredata.InvitationOrderField]{
Field: orderBy.Field,
Direction: orderBy.Direction,
}
}
cursor := types.NewCursor(first, after, last, before, pageOrderBy)
page, err := r.authzSvc.GetAllOrganizationInvitations(ctx, obj.ID, cursor)
if err != nil {
panic(fmt.Errorf("cannot list invitations: %w", err))
}
return types.NewInvitationConnection(page), nil
}
// Connectors is the resolver for the connectors field.
@@ -4866,6 +4959,27 @@ func (r *userResolver) People(ctx context.Context, obj *types.User, organization
return types.NewPeople(people), nil
}
// TotalCount is the resolver for the totalCount field.
func (r *userConnectionResolver) TotalCount(ctx context.Context, obj *types.UserConnection) (int, error) {
currentUser := UserFromContext(ctx)
if currentUser == nil {
return 0, fmt.Errorf("no authenticated user")
}
memberships, err := r.authzSvc.GetAllUserOrganizations(ctx, currentUser.ID)
if err != nil || len(memberships) == 0 {
return 0, fmt.Errorf("user has no organization memberships")
}
orgID := memberships[0].ID
count, err := r.authzSvc.CountOrganizationMemberships(ctx, orgID)
if err != nil {
return 0, fmt.Errorf("failed to count memberships: %w", err)
}
return count, nil
}
// Organization is the resolver for the organization field.
func (r *vendorResolver) Organization(ctx context.Context, obj *types.Vendor) (*types.Organization, error) {
prb := r.ProboService(ctx, obj.ID.TenantID())
@@ -5229,7 +5343,7 @@ func (r *viewerResolver) Organizations(ctx context.Context, obj *types.Viewer, f
}
cursor := types.NewCursor(first, after, last, before, pageOrderBy)
organizations, err := r.usrmgrSvc.ListOrganizationsForUserIDPaginated(ctx, user.ID, cursor)
organizations, err := r.authzSvc.GetUserOrganizations(ctx, user.ID, cursor)
if err != nil {
panic(fmt.Errorf("failed to list organizations for user: %w", err))
}
@@ -5318,6 +5432,11 @@ func (r *Resolver) FrameworkConnection() schema.FrameworkConnectionResolver {
return &frameworkConnectionResolver{r}
}
// InvitationConnection returns schema.InvitationConnectionResolver implementation.
func (r *Resolver) InvitationConnection() schema.InvitationConnectionResolver {
return &invitationConnectionResolver{r}
}
// Measure returns schema.MeasureResolver implementation.
func (r *Resolver) Measure() schema.MeasureResolver { return &measureResolver{r} }
@@ -5326,6 +5445,11 @@ func (r *Resolver) MeasureConnection() schema.MeasureConnectionResolver {
return &measureConnectionResolver{r}
}
// MembershipConnection returns schema.MembershipConnectionResolver implementation.
func (r *Resolver) MembershipConnection() schema.MembershipConnectionResolver {
return &membershipConnectionResolver{r}
}
// Mutation returns schema.MutationResolver implementation.
func (r *Resolver) Mutation() schema.MutationResolver { return &mutationResolver{r} }
@@ -5420,6 +5544,9 @@ func (r *Resolver) TrustCenterReferenceConnection() schema.TrustCenterReferenceC
// User returns schema.UserResolver implementation.
func (r *Resolver) User() schema.UserResolver { return &userResolver{r} }
// UserConnection returns schema.UserConnectionResolver implementation.
func (r *Resolver) UserConnection() schema.UserConnectionResolver { return &userConnectionResolver{r} }
// Vendor returns schema.VendorResolver implementation.
func (r *Resolver) Vendor() schema.VendorResolver { return &vendorResolver{r} }
@@ -5476,8 +5603,10 @@ type evidenceConnectionResolver struct{ *Resolver }
type fileResolver struct{ *Resolver }
type frameworkResolver struct{ *Resolver }
type frameworkConnectionResolver struct{ *Resolver }
type invitationConnectionResolver struct{ *Resolver }
type measureResolver struct{ *Resolver }
type measureConnectionResolver struct{ *Resolver }
type membershipConnectionResolver struct{ *Resolver }
type mutationResolver struct{ *Resolver }
type nonconformityResolver struct{ *Resolver }
type nonconformityConnectionResolver struct{ *Resolver }
@@ -5502,6 +5631,7 @@ type trustCenterDocumentAccessConnectionResolver struct{ *Resolver }
type trustCenterReferenceResolver struct{ *Resolver }
type trustCenterReferenceConnectionResolver struct{ *Resolver }
type userResolver struct{ *Resolver }
type userConnectionResolver struct{ *Resolver }
type vendorResolver struct{ *Resolver }
type vendorBusinessAssociateAgreementResolver struct{ *Resolver }
type vendorComplianceReportResolver struct{ *Resolver }

View File

@@ -26,17 +26,18 @@ import (
"github.com/99designs/gqlgen/graphql/handler"
"github.com/99designs/gqlgen/graphql/handler/extension"
"github.com/99designs/gqlgen/graphql/handler/transport"
"github.com/getprobo/probo/pkg/auth"
"github.com/getprobo/probo/pkg/authz"
"github.com/getprobo/probo/pkg/coredata"
"github.com/getprobo/probo/pkg/gid"
"github.com/getprobo/probo/pkg/probo"
console_v1 "github.com/getprobo/probo/pkg/server/api/console/v1"
"github.com/getprobo/probo/pkg/server/api/trust/v1/auth"
"github.com/getprobo/probo/pkg/server/api/trust/v1/schema"
"github.com/getprobo/probo/pkg/server/api/trust/v1/trustauth"
gqlutils "github.com/getprobo/probo/pkg/server/graphql"
"github.com/getprobo/probo/pkg/server/session"
"github.com/getprobo/probo/pkg/statelesstoken"
"github.com/getprobo/probo/pkg/trust"
"github.com/getprobo/probo/pkg/usrmgr"
"github.com/go-chi/chi/v5"
"go.gearno.de/kit/log"
)
@@ -79,31 +80,32 @@ func UserFromContext(ctx context.Context) *coredata.User {
return user
}
func TokenAccessFromContext(ctx context.Context) *auth.TokenAccessData {
tokenAccess, _ := ctx.Value(tokenAccessContextKey).(*auth.TokenAccessData)
func TokenAccessFromContext(ctx context.Context) *trustauth.TokenAccessData {
tokenAccess, _ := ctx.Value(tokenAccessContextKey).(*trustauth.TokenAccessData)
return tokenAccess
}
// UserFromContext implements auth.ContextAccessor interface
// UserFromContext implements trustauth.ContextAccessor interface
func (r *Resolver) UserFromContext(ctx context.Context) *coredata.User {
return UserFromContext(ctx)
}
// TokenAccessFromContext implements auth.ContextAccessor interface
func (r *Resolver) TokenAccessFromContext(ctx context.Context) *auth.TokenAccessData {
// TokenAccessFromContext implements trustauth.ContextAccessor interface
func (r *Resolver) TokenAccessFromContext(ctx context.Context) *trustauth.TokenAccessData {
return TokenAccessFromContext(ctx)
}
func NewMux(
logger *log.Logger,
usrmgrSvc *usrmgr.Service,
authSvc *auth.Service,
authzSvc *authz.Service,
trustSvc *trust.Service,
authCfg console_v1.AuthConfig,
trustAuthCfg TrustAuthConfig,
) *chi.Mux {
r := chi.NewMux()
r.Handle("/graphql", graphqlHandler(logger, usrmgrSvc, trustSvc, authCfg, trustAuthCfg))
r.Handle("/graphql", graphqlHandler(logger, authSvc, authzSvc, trustSvc, authCfg, trustAuthCfg))
r.Post("/auth/authenticate", authTokenHandler(trustSvc, trustAuthCfg))
r.Delete("/auth/logout", trustCenterLogoutHandler(authCfg, trustAuthCfg))
@@ -111,7 +113,7 @@ func NewMux(
return r
}
func graphqlHandler(logger *log.Logger, usrmgrSvc *usrmgr.Service, trustSvc *trust.Service, authCfg console_v1.AuthConfig, trustAuthCfg TrustAuthConfig) http.HandlerFunc {
func graphqlHandler(logger *log.Logger, authSvc *auth.Service, authzSvc *authz.Service, trustSvc *trust.Service, authCfg console_v1.AuthConfig, trustAuthCfg TrustAuthConfig) http.HandlerFunc {
resolver := &Resolver{
trustCenterSvc: trustSvc,
authCfg: authCfg,
@@ -122,7 +124,7 @@ func graphqlHandler(logger *log.Logger, usrmgrSvc *usrmgr.Service, trustSvc *tru
Resolvers: resolver,
}
c.Directives.MustBeAuthenticated = auth.MustBeAuthenticatedDirective(resolver)
c.Directives.MustBeAuthenticated = trustauth.MustBeAuthenticatedDirective(resolver)
es := schema.NewExecutableSchema(c)
@@ -137,7 +139,7 @@ func graphqlHandler(logger *log.Logger, usrmgrSvc *usrmgr.Service, trustSvc *tru
srv.SetRecoverFunc(gqlutils.RecoverFunc)
return WithSession(usrmgrSvc, trustSvc, authCfg, trustAuthCfg, srv.ServeHTTP)
return WithSession(authSvc, authzSvc, trustSvc, authCfg, trustAuthCfg, srv.ServeHTTP)
}
func (r *Resolver) RootTrustService(ctx context.Context) *trust.TenantService {
@@ -149,14 +151,14 @@ func (r *Resolver) PublicTrustService(ctx context.Context, tenantID gid.TenantID
}
func (r *Resolver) PrivateTrustService(ctx context.Context, tenantID gid.TenantID) (*trust.TenantService, error) {
if err := auth.ValidateTenantAccess(ctx, r, userTenantContextKey, tenantID); err != nil {
if err := trustauth.ValidateTenantAccess(ctx, r, userTenantContextKey, tenantID); err != nil {
return nil, fmt.Errorf("cannot access trust center: %w", err)
}
return r.trustCenterSvc.WithTenant(tenantID), nil
}
func WithSession(usrmgrSvc *usrmgr.Service, trustSvc *trust.Service, authCfg console_v1.AuthConfig, trustAuthCfg TrustAuthConfig, next http.HandlerFunc) http.HandlerFunc {
func WithSession(authSvc *auth.Service, authzSvc *authz.Service, trustSvc *trust.Service, authCfg console_v1.AuthConfig, trustAuthCfg TrustAuthConfig, next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
@@ -168,9 +170,9 @@ func WithSession(usrmgrSvc *usrmgr.Service, trustSvc *trust.Service, authCfg con
return
}
if authCtx := trySessionAuth(ctx, w, r, usrmgrSvc, authCfg); authCtx != nil {
if authCtx := trySessionAuth(ctx, w, r, authSvc, authzSvc, authCfg); authCtx != nil {
next(w, r.WithContext(authCtx))
updateSessionIfNeeded(authCtx, usrmgrSvc)
updateSessionIfNeeded(authCtx, authSvc)
return
}
@@ -178,7 +180,7 @@ func WithSession(usrmgrSvc *usrmgr.Service, trustSvc *trust.Service, authCfg con
}
}
func trySessionAuth(ctx context.Context, w http.ResponseWriter, r *http.Request, usrmgrSvc *usrmgr.Service, authCfg console_v1.AuthConfig) context.Context {
func trySessionAuth(ctx context.Context, w http.ResponseWriter, r *http.Request, authSvc *auth.Service, authzSvc *authz.Service, authCfg console_v1.AuthConfig) context.Context {
sessionAuthCfg := session.AuthConfig{
CookieName: authCfg.CookieName,
CookieSecret: authCfg.CookieSecret,
@@ -199,7 +201,7 @@ func trySessionAuth(ctx context.Context, w http.ResponseWriter, r *http.Request,
},
}
authResult := session.TryAuth(ctx, w, r, usrmgrSvc, sessionAuthCfg, errorHandler)
authResult := session.TryAuth(ctx, w, r, authSvc, authzSvc, sessionAuthCfg, errorHandler)
if authResult == nil {
return nil
}
@@ -235,7 +237,7 @@ func tryTokenAuth(ctx context.Context, w http.ResponseWriter, r *http.Request, t
return nil
}
tokenAccess := &auth.TokenAccessData{
tokenAccess := &trustauth.TokenAccessData{
TrustCenterID: basicPayload.Data.TrustCenterID,
Email: basicPayload.Data.Email,
TenantID: tenantID,
@@ -258,10 +260,10 @@ func clearTokenCookie(w http.ResponseWriter, trustAuthCfg TrustAuthConfig) {
})
}
func updateSessionIfNeeded(ctx context.Context, usrmgrSvc *usrmgr.Service) {
func updateSessionIfNeeded(ctx context.Context, authSvc *auth.Service) {
session := SessionFromContext(ctx)
if session != nil {
if err := usrmgrSvc.UpdateSession(ctx, session); err != nil {
if _, err := authSvc.UpdateSession(ctx, session.ID); err != nil {
panic(fmt.Errorf("failed to update session: %w", err))
}
}

View File

@@ -12,7 +12,7 @@
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package auth
package trustauth
import (
"context"

View File

@@ -21,6 +21,8 @@ import (
"strings"
"github.com/getprobo/probo/pkg/agents"
"github.com/getprobo/probo/pkg/auth"
"github.com/getprobo/probo/pkg/authz"
"github.com/getprobo/probo/pkg/connector"
"github.com/getprobo/probo/pkg/gid"
"github.com/getprobo/probo/pkg/probo"
@@ -30,7 +32,6 @@ import (
"github.com/getprobo/probo/pkg/server/trust"
"github.com/getprobo/probo/pkg/server/web"
trust_pkg "github.com/getprobo/probo/pkg/trust"
"github.com/getprobo/probo/pkg/usrmgr"
"github.com/go-chi/chi/v5"
"go.gearno.de/kit/log"
)
@@ -40,9 +41,10 @@ type Config struct {
AllowedOrigins []string
ExtraHeaderFields map[string]string
Probo *probo.Service
Usrmgr *usrmgr.Service
Auth *auth.Service
Authz *authz.Service
Trust *trust_pkg.Service
Auth api.ConsoleAuthConfig
ConsoleAuth api.ConsoleAuthConfig
TrustAuth api.TrustAuthConfig
ConnectorRegistry *connector.ConnectorRegistry
Agent *agents.Agent
@@ -68,9 +70,10 @@ func NewServer(cfg Config) (*Server, error) {
apiCfg := api.Config{
AllowedOrigins: cfg.AllowedOrigins,
Probo: cfg.Probo,
Usrmgr: cfg.Usrmgr,
Trust: cfg.Trust,
Auth: cfg.Auth,
Authz: cfg.Authz,
Trust: cfg.Trust,
ConsoleAuth: cfg.ConsoleAuth,
TrustAuth: cfg.TrustAuth,
ConnectorRegistry: cfg.ConnectorRegistry,
SafeRedirect: cfg.SafeRedirect,

View File

@@ -19,10 +19,11 @@ import (
"errors"
"net/http"
"github.com/getprobo/probo/pkg/authz"
"github.com/getprobo/probo/pkg/coredata"
"github.com/getprobo/probo/pkg/gid"
"github.com/getprobo/probo/pkg/securecookie"
"github.com/getprobo/probo/pkg/usrmgr"
"github.com/getprobo/probo/pkg/auth"
)
type AuthConfig struct {
@@ -48,7 +49,8 @@ func TryAuth(
ctx context.Context,
w http.ResponseWriter,
r *http.Request,
usrmgrSvc *usrmgr.Service,
authSvc *auth.Service,
authzSvc *authz.Service,
authCfg AuthConfig,
errorHandler ErrorHandler,
) *AuthResult {
@@ -71,7 +73,7 @@ func TryAuth(
return nil
}
session, err := usrmgrSvc.GetSession(ctx, sessionID)
session, err := authSvc.GetSession(ctx, sessionID)
if err != nil {
if errorHandler.OnSessionError != nil {
errorHandler.OnSessionError(w, authCfg)
@@ -79,7 +81,7 @@ func TryAuth(
return nil
}
user, err := usrmgrSvc.GetUserBySession(ctx, sessionID)
user, err := authSvc.GetUserBySession(ctx, sessionID)
if err != nil {
if errorHandler.OnUserError != nil {
errorHandler.OnUserError(w, authCfg)
@@ -87,7 +89,7 @@ func TryAuth(
return nil
}
tenantIDs, err := usrmgrSvc.ListTenantsForUserID(ctx, user.ID)
organizations, err := authzSvc.GetAllUserOrganizations(ctx, user.ID)
if err != nil {
if errorHandler.OnTenantError != nil {
errorHandler.OnTenantError(err)
@@ -95,6 +97,11 @@ func TryAuth(
return nil
}
tenantIDs := make([]gid.TenantID, len(organizations))
for i, org := range organizations {
tenantIDs[i] = org.ID.TenantID()
}
return &AuthResult{
Session: session,
User: user,