diff --git a/apps/console/src/layouts/MainLayout.tsx b/apps/console/src/layouts/MainLayout.tsx index 3aed59a27..edefafadc 100644 --- a/apps/console/src/layouts/MainLayout.tsx +++ b/apps/console/src/layouts/MainLayout.tsx @@ -53,9 +53,6 @@ const MainLayoutQuery = graphql` fullName email } - invitations(first: 1, filter: {statuses: [PENDING]}) { - totalCount - } } organization: node(id: $organizationId) { ... on Organization { @@ -258,49 +255,79 @@ interface OrganizationsResponse { organizations: Organization[]; } +interface Invitation { + id: string; + email: string; + fullName: string; + role: string; + expiresAt: string; + acceptedAt?: string | null; + createdAt: string; + organization: { + id: string; + name: string; + }; +} + +interface InvitationsResponse { + invitations: Invitation[]; +} + function OrganizationSelectorWrapper({ organizationId }: { organizationId: string }) { const data = useLazyLoadQuery(MainLayoutQuery, { organizationId }); - return ; + return ; } function OrganizationSelector({ - viewer, currentOrganization }: { - viewer: MainLayoutQueryType["response"]["viewer"]; currentOrganization: MainLayoutQueryType["response"]["organization"]; }) { const [organizations, setOrganizations] = useState([]); + const [pendingInvitationsCount, setPendingInvitationsCount] = useState(0); const [isLoading, setIsLoading] = useState(true); const [error, setError] = useState(null); const { __ } = useTranslate(); - const pendingInvitationsCount = viewer.invitations.totalCount; - useEffect(() => { - const fetchOrganizations = async () => { + const fetchData = async () => { try { setIsLoading(true); - const response = await fetch('/auth/organizations', { - credentials: 'include', - }); - if (!response.ok) { + // Fetch organizations and invitations in parallel + const [orgsResponse, invitationsResponse] = await Promise.all([ + fetch('/auth/organizations', { credentials: 'include' }), + fetch('/auth/invitations', { credentials: 'include' }) + ]); + + if (!orgsResponse.ok) { throw new Error('Failed to fetch organizations'); } - const data: OrganizationsResponse = await response.json(); - setOrganizations(data.organizations); + if (!invitationsResponse.ok) { + throw new Error('Failed to fetch invitations'); + } + + const orgsData: OrganizationsResponse = await orgsResponse.json(); + const invitationsData: InvitationsResponse = await invitationsResponse.json(); + + // Count pending invitations (those without acceptedAt) + const pendingCount = invitationsData.invitations.filter( + inv => !inv.acceptedAt + ).length; + + setOrganizations(orgsData.organizations); + setPendingInvitationsCount(pendingCount); setError(null); } catch (err) { setError(err instanceof Error ? err.message : 'Unknown error'); - console.error('Failed to fetch organizations:', err); + console.error('Failed to fetch data:', err); } finally { setIsLoading(false); } }; - fetchOrganizations(); + fetchData(); }, []); if (error) { diff --git a/apps/console/src/layouts/__generated__/MainLayoutQuery.graphql.ts b/apps/console/src/layouts/__generated__/MainLayoutQuery.graphql.ts index d23b0081f..0aee147b9 100644 --- a/apps/console/src/layouts/__generated__/MainLayoutQuery.graphql.ts +++ b/apps/console/src/layouts/__generated__/MainLayoutQuery.graphql.ts @@ -1,5 +1,5 @@ /** - * @generated SignedSource<<26b3620e6aed7f97ffb1710be1eb267a>> + * @generated SignedSource<> * @lightSyntaxTransform * @nogrep */ @@ -20,9 +20,6 @@ export type MainLayoutQuery$data = { }; readonly viewer: { readonly id: string; - readonly invitations: { - readonly totalCount: number; - }; readonly user: { readonly email: string; readonly fullName: string; @@ -63,54 +60,21 @@ v3 = { "name": "email", "storageKey": null }, -v4 = { - "alias": null, - "args": [ - { - "kind": "Literal", - "name": "filter", - "value": { - "statuses": [ - "PENDING" - ] - } - }, - { - "kind": "Literal", - "name": "first", - "value": 1 - } - ], - "concreteType": "InvitationConnection", - "kind": "LinkedField", - "name": "invitations", - "plural": false, - "selections": [ - { - "alias": null, - "args": null, - "kind": "ScalarField", - "name": "totalCount", - "storageKey": null - } - ], - "storageKey": "invitations(filter:{\"statuses\":[\"PENDING\"]},first:1)" -}, -v5 = [ +v4 = [ { "kind": "Variable", "name": "id", "variableName": "organizationId" } ], -v6 = { +v5 = { "alias": null, "args": null, "kind": "ScalarField", "name": "name", "storageKey": null }, -v7 = { +v6 = { "alias": null, "args": null, "kind": "ScalarField", @@ -145,14 +109,13 @@ return { (v3/*: any*/) ], "storageKey": null - }, - (v4/*: any*/) + } ], "storageKey": null }, { "alias": "organization", - "args": (v5/*: any*/), + "args": (v4/*: any*/), "concreteType": null, "kind": "LinkedField", "name": "node", @@ -162,8 +125,8 @@ return { "kind": "InlineFragment", "selections": [ (v1/*: any*/), - (v6/*: any*/), - (v7/*: any*/) + (v5/*: any*/), + (v6/*: any*/) ], "type": "Organization", "abstractKey": null @@ -203,14 +166,13 @@ return { (v1/*: any*/) ], "storageKey": null - }, - (v4/*: any*/) + } ], "storageKey": null }, { "alias": "organization", - "args": (v5/*: any*/), + "args": (v4/*: any*/), "concreteType": null, "kind": "LinkedField", "name": "node", @@ -227,8 +189,8 @@ return { { "kind": "InlineFragment", "selections": [ - (v6/*: any*/), - (v7/*: any*/) + (v5/*: any*/), + (v6/*: any*/) ], "type": "Organization", "abstractKey": null @@ -239,16 +201,16 @@ return { ] }, "params": { - "cacheID": "a8f9f58d27677c55b5a217617db83e27", + "cacheID": "ee5a60e709dee856df7d2fef13974c9f", "id": null, "metadata": {}, "name": "MainLayoutQuery", "operationKind": "query", - "text": "query MainLayoutQuery(\n $organizationId: ID!\n) {\n viewer {\n id\n user {\n fullName\n email\n id\n }\n invitations(first: 1, filter: {statuses: [PENDING]}) {\n totalCount\n }\n }\n organization: node(id: $organizationId) {\n __typename\n ... on Organization {\n id\n name\n logoUrl\n }\n id\n }\n}\n" + "text": "query MainLayoutQuery(\n $organizationId: ID!\n) {\n viewer {\n id\n user {\n fullName\n email\n id\n }\n }\n organization: node(id: $organizationId) {\n __typename\n ... on Organization {\n id\n name\n logoUrl\n }\n id\n }\n}\n" } }; })(); -(node as any).hash = "17986fcea321c4567d86584d1a9f89c1"; +(node as any).hash = "9ea3e5a91a2d2be0993e7deebafa11b0"; export default node; diff --git a/pkg/auth/access.go b/pkg/auth/access.go new file mode 100644 index 000000000..586769bac --- /dev/null +++ b/pkg/auth/access.go @@ -0,0 +1,98 @@ +package auth + +import ( + "fmt" + + "github.com/getprobo/probo/pkg/coredata" + "github.com/getprobo/probo/pkg/gid" +) + +// AuthMethod represents a method of authentication +type AuthMethod int + +const ( + AuthMethodPassword AuthMethod = iota + AuthMethodSAML + AuthMethodAny // Used when either password or SAML would work +) + +// OrgAuthRequirement encapsulates the authentication requirements for accessing an organization +type OrgAuthRequirement struct { + OrganizationID gid.GID + EmailDomain string + SAMLConfig *coredata.SAMLConfiguration // nil if no SAML config applies to this org+domain +} + +// AccessResult represents the result of an organization access check +type AccessResult struct { + OrganizationID gid.GID + Allowed bool + MissingAuth AuthMethod // Which auth method is missing (if not allowed) + SAMLConfig *coredata.SAMLConfiguration // The SAML config involved (if any) +} + +// Check performs the access control decision based on session state +// This is pure business logic with no side effects - easily testable +func (r OrgAuthRequirement) Check(session coredata.SessionData) AccessResult { + // No SAML config or disabled → requires password authentication + if r.SAMLConfig == nil || !r.SAMLConfig.Enabled || !r.SAMLConfig.DomainVerified { + return AccessResult{ + OrganizationID: r.OrganizationID, + Allowed: session.PasswordAuthenticated, + MissingAuth: AuthMethodPassword, + SAMLConfig: nil, + } + } + + // Check if user has SAML authentication for this specific organization + orgKey := r.OrganizationID.String() + _, hasSAML := session.SAMLAuthenticatedOrgs[orgKey] + + // SAML enforcement: REQUIRED → must have SAML auth for this org + if r.SAMLConfig.EnforcementPolicy == coredata.SAMLEnforcementPolicyRequired { + return AccessResult{ + OrganizationID: r.OrganizationID, + Allowed: hasSAML, + MissingAuth: AuthMethodSAML, + SAMLConfig: r.SAMLConfig, + } + } + + // SAML enforcement: OPTIONAL → needs either password (global) OR SAML (for this org) + hasAnyAuth := session.PasswordAuthenticated || hasSAML + missingAuth := AuthMethodAny + if hasAnyAuth { + missingAuth = AuthMethodPassword // Not actually missing, but need a value + } + + return AccessResult{ + OrganizationID: r.OrganizationID, + Allowed: hasAnyAuth, + MissingAuth: missingAuth, + SAMLConfig: r.SAMLConfig, + } +} + +// ToError converts an AccessResult to an error if access is denied +// This handles the presentation layer concern of generating appropriate errors and redirect URLs +func (r AccessResult) ToError(baseURL string) error { + if r.Allowed { + return nil + } + + switch r.MissingAuth { + case AuthMethodPassword: + return ErrPasswordAuthRequired{ + OrganizationID: r.OrganizationID, + RedirectURL: fmt.Sprintf("%s/authentication/login?method=password", baseURL), + } + case AuthMethodSAML, AuthMethodAny: + return ErrSAMLAuthRequired{ + ConfigID: r.SAMLConfig.ID, + OrganizationID: r.OrganizationID, + RedirectURL: fmt.Sprintf("%s/auth/saml/login/%s", baseURL, r.SAMLConfig.ID), + } + default: + return fmt.Errorf("access denied to organization %s", r.OrganizationID) + } +} diff --git a/pkg/auth/service.go b/pkg/auth/service.go index abba58165..8e184d436 100644 --- a/pkg/auth/service.go +++ b/pkg/auth/service.go @@ -36,8 +36,6 @@ import ( ) type ( - // Service handles ONLY user authentication and management - // No organization-related logic - that belongs to authz service Service struct { pg *pg.Client encryptionKey cipher.EncryptionKey @@ -49,7 +47,6 @@ type ( invitationTokenValidity time.Duration } - // TenantAuthService handles tenant-scoped authentication operations TenantAuthService struct { pg *pg.Client encryptionKey cipher.EncryptionKey @@ -60,6 +57,13 @@ type ( scope coredata.Scoper } + OrganizationAccessResponse struct { + OrganizationID gid.GID + HasAccess bool + Error error + SAMLConfig *coredata.SAMLConfiguration + } + ErrInvalidCredentials struct { message string } @@ -95,6 +99,8 @@ type ( ErrSignupDisabled struct{} + ErrSAMLAutoSignupDisabled struct{} + EmailConfirmationData struct { UserID gid.GID `json:"uid"` Email string `json:"email"` @@ -109,6 +115,17 @@ type ( PasswordResetData struct { Email string `json:"email"` } + + ErrSAMLAuthRequired struct { + ConfigID gid.GID + OrganizationID gid.GID + RedirectURL string + } + + ErrPasswordAuthRequired struct { + OrganizationID gid.GID + RedirectURL string + } ) const ( @@ -152,6 +169,18 @@ func (e ErrSignupDisabled) Error() string { return "signup is disabled, contact the owner of the Probo instance" } +func (e ErrSAMLAutoSignupDisabled) Error() string { + return "SAML auto-signup is disabled for this organization" +} + +func (e ErrSAMLAuthRequired) Error() string { + return "SAML authentication required for this organization" +} + +func (e ErrPasswordAuthRequired) Error() string { + return "password authentication required for this organization" +} + func NewService( ctx context.Context, pgClient *pg.Client, @@ -221,6 +250,7 @@ func (s Service) ForgetPassword( if errors.As(err, &errUserNotFound) { return nil // Don't leak information about non-existent users } + return fmt.Errorf("cannot load user: %w", err) } @@ -259,7 +289,8 @@ func (s Service) SignUp( return nil, nil, &ErrSignupDisabled{} } - if _, err := mail.ParseAddress(emailAddress); err != nil { + emailAddress2, err := mail.ParseAddress(emailAddress) + if err != nil { return nil, nil, &ErrInvalidEmail{emailAddress} } @@ -279,7 +310,7 @@ func (s Service) SignUp( now := time.Now() user := &coredata.User{ ID: gid.New(gid.NilTenant, coredata.UserEntityType), - EmailAddress: emailAddress, + EmailAddress: emailAddress2.Address, HashedPassword: hashedPassword, EmailAddressVerified: false, FullName: fullName, @@ -307,6 +338,7 @@ func (s Service) SignUp( if errors.As(err, &errUserAlreadyExists) { return &ErrUserAlreadyExists{errUserAlreadyExists.Error()} } + return fmt.Errorf("cannot insert user: %w", err) } @@ -365,25 +397,30 @@ func (s Service) SignUp( return user, session, nil } -func (s Service) CreateOrGetSAMLUser( +func (s Service) ProvisionSAMLUser( ctx context.Context, + samlConfigID gid.GID, + organizationID gid.GID, emailAddress string, fullName string, samlSubject string, -) (*coredata.User, error) { + existingSession *coredata.Session, + sessionDuration time.Duration, +) (*coredata.Session, *coredata.User, error) { if _, err := mail.ParseAddress(emailAddress); err != nil { - return nil, &ErrInvalidEmail{emailAddress} + return nil, nil, &ErrInvalidEmail{emailAddress} } if fullName == "" { - return nil, &ErrInvalidFullName{fullName} + return nil, nil, &ErrInvalidFullName{fullName} } if samlSubject == "" { - return nil, fmt.Errorf("SAML subject cannot be empty") + return nil, nil, fmt.Errorf("SAML subject cannot be empty") } - var user coredata.User + user := &coredata.User{} + session := &coredata.Session{} now := time.Now() err := s.pg.WithTx( @@ -398,124 +435,76 @@ func (s Service) CreateOrGetSAMLUser( if err := user.Update(ctx, tx); err != nil { return fmt.Errorf("cannot update user: %w", err) } - - return nil - } - - user = coredata.User{ - ID: gid.New(gid.NilTenant, coredata.UserEntityType), - EmailAddress: emailAddress, - HashedPassword: nil, - EmailAddressVerified: true, - FullName: fullName, - SAMLSubject: &samlSubject, - CreatedAt: now, - UpdatedAt: now, - } - - if err := user.Insert(ctx, tx, coredata.NewNoScope()); err != nil { - return fmt.Errorf("cannot insert SAML user: %w", err) - } - - return nil - }, - ) - - if err != nil { - return nil, err - } - - return &user, nil -} - -func (s Service) CreateSessionForUser( - ctx context.Context, - userID gid.GID, - sessionDuration time.Duration, -) (*coredata.Session, error) { - now := time.Now() - - session := &coredata.Session{ - ID: gid.New(gid.NilTenant, coredata.SessionEntityType), - UserID: userID, - Data: coredata.SessionData{}, - ExpiredAt: now.Add(sessionDuration), - CreatedAt: now, - UpdatedAt: now, - } - - err := s.pg.WithTx( - ctx, - func(tx pg.Conn) error { - if err := session.Insert(ctx, tx); err != nil { - return fmt.Errorf("cannot insert session: %w", err) - } - return nil - }, - ) - - if err != nil { - return nil, err - } - - return session, nil -} - -func (s Service) SignIn( - ctx context.Context, - emailAddress string, - password string, -) (*coredata.Session, *coredata.User, error) { - if _, err := mail.ParseAddress(emailAddress); err != nil { - return nil, nil, &ErrInvalidCredentials{"invalid email or password"} - } - - user := &coredata.User{} - session := &coredata.Session{} - - err := s.pg.WithTx( - ctx, - func(tx pg.Conn) error { - // Load user by email (all users are global now) - if err := user.LoadByEmail(ctx, tx, emailAddress); err != nil { - var errUserNotFound *coredata.ErrUserNotFound - if errors.As(err, &errUserNotFound) { - return &ErrInvalidCredentials{"invalid email or password"} + } else { + var samlConfig coredata.SAMLConfiguration + if err := samlConfig.LoadByID(ctx, tx, coredata.NewNoScope(), samlConfigID); err != nil { + return fmt.Errorf("cannot load SAML configuration: %w", err) + } + + if !samlConfig.AutoSignupEnabled { + return &ErrSAMLAutoSignupDisabled{} + } + + *user = coredata.User{ + ID: gid.New(gid.NilTenant, coredata.UserEntityType), + EmailAddress: emailAddress, + HashedPassword: nil, + EmailAddressVerified: true, + FullName: fullName, + SAMLSubject: &samlSubject, + CreatedAt: now, + UpdatedAt: now, + } + + if err := user.Insert(ctx, tx, coredata.NewNoScope()); err != nil { + return fmt.Errorf("cannot insert SAML user: %w", err) } - return fmt.Errorf("cannot load user by email: %w", err) } - // Verify password - match, err := s.hp.ComparePasswordAndHash([]byte(password), user.HashedPassword) - if err != nil { - return fmt.Errorf("cannot verify password: %w", err) - } - if !match { - return &ErrInvalidCredentials{"invalid email or password"} - } + if existingSession != nil && existingSession.UserID == user.ID { + if err := session.LoadByID(ctx, tx, existingSession.ID); err != nil { + return fmt.Errorf("cannot load session: %w", err) + } - // Create new session with password authentication flag set - now := time.Now() - session = &coredata.Session{ - ID: gid.New(gid.NilTenant, coredata.SessionEntityType), - UserID: user.ID, - Data: coredata.SessionData{ - PasswordAuthenticated: true, - SAMLAuthenticatedOrgs: make(map[string]coredata.SAMLAuthInfo), - }, - ExpiredAt: now.Add(24 * time.Hour * 7), // 7 days - CreatedAt: now, - UpdatedAt: now, - } + if session.Data.SAMLAuthenticatedOrgs == nil { + session.Data.SAMLAuthenticatedOrgs = make(map[string]coredata.SAMLAuthInfo) + } + session.Data.SAMLAuthenticatedOrgs[organizationID.String()] = coredata.SAMLAuthInfo{ + AuthenticatedAt: now, + SAMLConfigID: samlConfigID, + SAMLSubject: samlSubject, + } + session.UpdatedAt = now - if err := session.Insert(ctx, tx); err != nil { - return fmt.Errorf("cannot insert session: %w", err) + if err := session.Update(ctx, tx); err != nil { + return fmt.Errorf("cannot update session: %w", err) + } + } else { + *session = coredata.Session{ + ID: gid.New(gid.NilTenant, coredata.SessionEntityType), + UserID: user.ID, + Data: coredata.SessionData{ + SAMLAuthenticatedOrgs: map[string]coredata.SAMLAuthInfo{ + organizationID.String(): { + AuthenticatedAt: now, + SAMLConfigID: samlConfigID, + SAMLSubject: samlSubject, + }, + }, + }, + ExpiredAt: now.Add(sessionDuration), + CreatedAt: now, + UpdatedAt: now, + } + + if err := session.Insert(ctx, tx); err != nil { + return fmt.Errorf("cannot insert session: %w", err) + } } return nil }, ) - if err != nil { return nil, nil, err } @@ -523,7 +512,7 @@ func (s Service) SignIn( return session, user, nil } -func (s Service) SignInWithExistingSession( +func (s Service) SignIn( ctx context.Context, emailAddress string, password string, @@ -579,7 +568,7 @@ func (s Service) SignInWithExistingSession( PasswordAuthenticated: true, SAMLAuthenticatedOrgs: make(map[string]coredata.SAMLAuthInfo), }, - ExpiredAt: now.Add(24 * time.Hour * 7), + ExpiredAt: now.Add(24 * time.Hour * 7), // 7 days CreatedAt: now, UpdatedAt: now, } @@ -629,7 +618,6 @@ func (s Service) GetSession(ctx context.Context, sessionID gid.GID) (*coredata.S } if time.Now().After(session.ExpiredAt) { - // Clean up expired session _ = coredata.DeleteSession(ctx, conn, sessionID) return &ErrSessionExpired{"session expired"} } @@ -654,26 +642,7 @@ func (s Service) GetUserByID(ctx context.Context, userID gid.GID) (*coredata.Use if err := user.LoadByID(ctx, conn, userID); err != nil { return fmt.Errorf("cannot load user by ID: %w", err) } - return nil - }, - ) - if err != nil { - return nil, err - } - - return user, nil -} - -func (s Service) GetUserByEmail(ctx context.Context, email string) (*coredata.User, error) { - user := &coredata.User{} - - err := s.pg.WithConn( - ctx, - func(conn pg.Conn) error { - if err := user.LoadByEmail(ctx, conn, email); err != nil { - return fmt.Errorf("cannot load user by email: %w", err) - } return nil }, ) @@ -709,7 +678,7 @@ func (s Service) UpdateSession(ctx context.Context, sessionID gid.GID) (*coredat } now := time.Now() - session.ExpiredAt = now.Add(24 * time.Hour * 7) // Extend by 7 days + session.ExpiredAt = now.Add(24 * time.Hour * 7) session.UpdatedAt = now if err := session.Update(ctx, tx); err != nil { return fmt.Errorf("cannot update session: %w", err) @@ -756,9 +725,11 @@ func (s Service) ConfirmEmail(ctx context.Context, tokenString string) error { TokenTypeEmailConfirmation, tokenString, ) + if err != nil { return &ErrInvalidTokenType{"invalid confirmation token"} } + emailConfirmationData := payload.Data return s.pg.WithTx( @@ -788,9 +759,11 @@ func (s Service) ResetPassword(ctx context.Context, tokenString string, newPassw TokenTypePasswordReset, tokenString, ) + if err != nil { return &ErrInvalidTokenType{"invalid reset token"} } + passwordResetData := payload.Data if len(newPassword) < 8 || len(newPassword) > 128 { @@ -811,6 +784,7 @@ func (s Service) ResetPassword(ctx context.Context, tokenString string, newPassw if errors.As(err, &errUserNotFound) { return nil // Don't leak information about non-existent users } + return fmt.Errorf("cannot load user: %w", err) } @@ -837,18 +811,17 @@ func (s Service) SignupFromInvitation( if err != nil { return nil, nil, &ErrInvalidTokenType{"invalid invitation token"} } - invitationData := payload.Data if len(password) < 8 || len(password) > 128 { return nil, nil, &ErrInvalidPassword{minLength: 8, maxLength: 128} } - if _, err := mail.ParseAddress(invitationData.Email); err != nil { - return nil, nil, &ErrInvalidEmail{invitationData.Email} + if _, err := mail.ParseAddress(payload.Data.Email); err != nil { + return nil, nil, &ErrInvalidEmail{payload.Data.Email} } if fullName == "" { - fullName = invitationData.FullName + fullName = payload.Data.FullName } if fullName == "" { @@ -863,16 +836,17 @@ func (s Service) SignupFromInvitation( var user *coredata.User var session *coredata.Session - scope := coredata.NewScope(invitationData.InvitationID.TenantID()) + scope := coredata.NewScope(payload.Data.InvitationID.TenantID()) err = s.pg.WithTx( ctx, func(tx pg.Conn) error { invitation := &coredata.Invitation{} - if err := invitation.LoadByID(ctx, tx, scope, invitationData.InvitationID); err != nil { + if err := invitation.LoadByID(ctx, tx, scope, payload.Data.InvitationID); err != nil { var errInvitationNotFound *coredata.ErrInvitationNotFound if errors.As(err, &errInvitationNotFound) { return fmt.Errorf("invitation was deleted or no longer exists") } + return fmt.Errorf("cannot load invitation: %w", err) } @@ -887,7 +861,7 @@ func (s Service) SignupFromInvitation( now := time.Now() user = &coredata.User{ ID: gid.New(gid.NilTenant, coredata.UserEntityType), - EmailAddress: invitationData.Email, + EmailAddress: payload.Data.Email, HashedPassword: hashedPassword, EmailAddressVerified: true, FullName: fullName, @@ -900,6 +874,7 @@ func (s Service) SignupFromInvitation( if errors.As(err, &errUserAlreadyExists) { return &ErrUserAlreadyExists{errUserAlreadyExists.Error()} } + return fmt.Errorf("cannot insert user: %w", err) } @@ -930,234 +905,246 @@ func (s Service) SignupFromInvitation( return user, session, nil } -// IsTenantUser removed - all users are now global (no tenant distinction) - -func (s Service) GetUserAuthMethod( +func (s *TenantAuthService) GetUserAuthMethod( ctx context.Context, - scope coredata.Scoper, userID gid.GID, organizationID gid.GID, session *coredata.Session, ) (coredata.UserAuthMethod, error) { - // Load the user to check their email and SAML subject - user := &coredata.User{} - err := s.pg.WithConn(ctx, func(conn pg.Conn) error { - return user.LoadByID(ctx, conn, userID) - }) - if err != nil { - return "", fmt.Errorf("cannot load user: %w", err) - } + var authMethod coredata.UserAuthMethod - // If user doesn't have a SAML subject, they only use password auth - if user.SAMLSubject == nil || *user.SAMLSubject == "" { - return coredata.UserAuthMethodPassword, nil - } + err := s.pg.WithConn( + ctx, + func(conn pg.Conn) error { + user := &coredata.User{} + if err := user.LoadByID(ctx, conn, userID); err != nil { + return fmt.Errorf("cannot load user: %w", err) + } - // User has SAML subject - check if there's SAML config for this org + user's domain - // Extract domain from user email - emailParts := []byte(user.EmailAddress) + // If user doesn't have a SAML subject, they only use password auth + if user.SAMLSubject == nil || *user.SAMLSubject == "" { + authMethod = coredata.UserAuthMethodPassword + return nil + } + + // User has SAML subject - check if there's SAML config for this org + user's domain + domain := extractDomain(user.EmailAddress) + if domain == "" { + authMethod = coredata.UserAuthMethodPassword + return nil + } + + // Check if SAML is configured for this org + domain + var samlConfig coredata.SAMLConfiguration + err := samlConfig.LoadByOrganizationIDAndEmailDomain(ctx, conn, s.scope, organizationID, domain) + if err != nil { + if errors.Is(err, pgx.ErrNoRows) { + authMethod = coredata.UserAuthMethodPassword + return nil + } + return fmt.Errorf("cannot check SAML configuration: %w", err) + } + + authMethod = coredata.UserAuthMethodSAML + return nil + }, + ) + + return authMethod, err +} + +func extractDomain(email string) string { atIndex := -1 - for i, b := range emailParts { - if b == '@' { + for i := 0; i < len(email); i++ { + if email[i] == '@' { atIndex = i break } } - if atIndex == -1 { - return coredata.UserAuthMethodPassword, nil + if atIndex == -1 || atIndex == len(email)-1 { + return "" } - domain := string(emailParts[atIndex+1:]) + return email[atIndex+1:] +} - // Check if SAML is configured for this org + domain - var samlConfig coredata.SAMLConfiguration - orgScope := coredata.NewScope(organizationID.TenantID()) - err = s.pg.WithConn(ctx, func(conn pg.Conn) error { - err := samlConfig.LoadByOrganizationIDAndEmailDomain(ctx, conn, orgScope, organizationID, domain) - if err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return nil // No SAML config for this org+domain +// CheckOrganizationAccess checks access to multiple organizations in a single database query +// Always uses batch processing to avoid N+1 query problems +func (s Service) CheckOrganizationAccess( + ctx context.Context, + user *coredata.User, + organizationIDs []gid.GID, + session *coredata.Session, +) (map[gid.GID]AccessResult, error) { + results := make(map[gid.GID]AccessResult, len(organizationIDs)) + + domain := extractDomain(user.EmailAddress) + if domain == "" { + // All organizations fail with invalid email + for _, orgID := range organizationIDs { + results[orgID] = AccessResult{ + OrganizationID: orgID, + Allowed: false, + MissingAuth: AuthMethodPassword, + SAMLConfig: nil, } - return err } - return nil + return results, nil + } + + // Fetch all SAML configs for these organizations in a single query + var samlConfigs map[gid.GID]*coredata.SAMLConfiguration + err := s.pg.WithConn(ctx, func(conn pg.Conn) error { + var err error + samlConfigs, err = coredata.LoadSAMLConfigurationsByOrganizationIDsAndEmailDomain( + ctx, + conn, + organizationIDs, + domain, + ) + return err }) if err != nil { - return "", fmt.Errorf("cannot check SAML configuration: %w", err) + return nil, fmt.Errorf("cannot load SAML configurations: %w", err) } - // If SAML config exists for this org+domain, user enrolled via SAML - if samlConfig.ID != (gid.GID{}) { - return coredata.UserAuthMethodSAML, nil + // Apply business logic for each organization using pure function + for _, orgID := range organizationIDs { + requirement := OrgAuthRequirement{ + OrganizationID: orgID, + EmailDomain: domain, + SAMLConfig: samlConfigs[orgID], // May be nil + } + results[orgID] = requirement.Check(session.Data) } - // No SAML config for this org, user uses password - return coredata.UserAuthMethodPassword, nil + return results, nil } -// Organization Access Control - -type ( - // ErrSAMLAuthRequired indicates user must authenticate via SAML to access org - ErrSAMLAuthRequired struct { - ConfigID gid.GID - OrganizationID gid.GID - RedirectURL string // SAML IdP login URL - } - - // ErrPasswordAuthRequired indicates user must authenticate with password to access org - ErrPasswordAuthRequired struct { - OrganizationID gid.GID - RedirectURL string // Password login page URL - } -) - -func (e ErrSAMLAuthRequired) Error() string { - return "SAML authentication required for this organization" -} - -func (e ErrPasswordAuthRequired) Error() string { - return "password authentication required for this organization" -} - -// CheckOrganizationAccess determines if a user can access an organization -// based on SAML configuration and session authentication state -func (s Service) CheckOrganizationAccess( +// CheckSingleOrganizationAccess is a convenience wrapper for checking access to a single organization +// It uses the batch method internally to maintain a single code path +func (s Service) CheckSingleOrganizationAccess( ctx context.Context, user *coredata.User, organizationID gid.GID, session *coredata.Session, ) error { - // Extract domain from user email - emailParts := []byte(user.EmailAddress) - atIndex := -1 - for i, b := range emailParts { - if b == '@' { - atIndex = i - break - } - } - if atIndex == -1 { - return fmt.Errorf("invalid email address format") - } - domain := string(emailParts[atIndex+1:]) - - // Find SAML configuration for this organization and domain - var samlConfig coredata.SAMLConfiguration - scope := coredata.NewScope(organizationID.TenantID()) - err := s.pg.WithConn(ctx, func(conn pg.Conn) error { - err := samlConfig.LoadByOrganizationIDAndEmailDomain(ctx, conn, scope, organizationID, domain) - if err != nil { - // If no SAML config found for this organization and domain, that's okay - not an error - // Just means this organization doesn't have SAML configured for this domain - if errors.Is(err, pgx.ErrNoRows) { - return nil - } - return err - } - return nil - }) + results, err := s.CheckOrganizationAccess(ctx, user, []gid.GID{organizationID}, session) if err != nil { - return fmt.Errorf("cannot check SAML configuration: %w", err) + return err } - // Check if SAML is configured and enabled for this domain and organization - if samlConfig.ID != (gid.GID{}) && samlConfig.Enabled && samlConfig.DomainVerified { - // SAML config exists for this org - check enforcement policy - if samlConfig.EnforcementPolicy == coredata.SAMLEnforcementPolicyRequired { - // SAML is REQUIRED - check if user has SAML-authenticated for this org - authInfo, hasSAMLAuth := session.Data.SAMLAuthenticatedOrgs[organizationID.String()] - if !hasSAMLAuth { - // Build SAML login URL - samlLoginURL := fmt.Sprintf("%s/auth/saml/login/%s", s.baseURL, samlConfig.ID) - return ErrSAMLAuthRequired{ - ConfigID: samlConfig.ID, - OrganizationID: organizationID, - RedirectURL: samlLoginURL, - } - } - - // Optional: Check if SAML auth is still recent (not too old) - // For now, we trust the session lifetime - _ = authInfo - } else { - // SAML is OPTIONAL or OFF - allow either password OR SAML auth for this specific org - hasSAMLAuth := false - if _, ok := session.Data.SAMLAuthenticatedOrgs[organizationID.String()]; ok { - hasSAMLAuth = true - } - - if !session.Data.PasswordAuthenticated && !hasSAMLAuth { - // User needs to authenticate - offer SAML as option - samlLoginURL := fmt.Sprintf("%s/auth/saml/login/%s", s.baseURL, samlConfig.ID) - return ErrSAMLAuthRequired{ - ConfigID: samlConfig.ID, - OrganizationID: organizationID, - RedirectURL: samlLoginURL, - } - } - } - } else { - // No SAML configuration for this org+domain combination - // Require password authentication for password-only organizations - if !session.Data.PasswordAuthenticated { - // User hasn't authenticated with password - require password authentication - loginURL := fmt.Sprintf("%s/authentication/login?method=password", s.baseURL) - return ErrPasswordAuthRequired{ - OrganizationID: organizationID, - RedirectURL: loginURL, - } - } + result, ok := results[organizationID] + if !ok { + return fmt.Errorf("no access result for organization %s", organizationID) } - return nil // Access granted + return result.ToError(s.baseURL) } -// InitiateDomainVerification creates a SAML configuration with unverified domain and generates verification token -func (s Service) InitiateDomainVerification( +// BaseURL returns the base URL for the service +func (s Service) BaseURL() string { + return s.baseURL +} + +func (s Service) GetOrganizationLogoFile( + ctx context.Context, + user *coredata.User, + organizationID gid.GID, + session *coredata.Session, +) (*coredata.File, error) { + // Check authentication requirements before allowing access to logo + err := s.CheckSingleOrganizationAccess(ctx, user, organizationID, session) + if err != nil { + return nil, fmt.Errorf("access denied: %w", err) + } + + var logoFile *coredata.File + + err = s.pg.WithConn( + ctx, + func(conn pg.Conn) error { + scope := coredata.NewScope(organizationID.TenantID()) + + var membership coredata.Membership + err := membership.LoadByUserAndOrg(ctx, conn, scope, user.ID, organizationID) + if err != nil { + if _, ok := err.(coredata.ErrMembershipNotFound); ok { + return fmt.Errorf("user does not have access to this organization") + } + + return fmt.Errorf("cannot verify membership: %w", err) + } + + var organization coredata.Organization + if err := organization.LoadByID(ctx, conn, scope, organizationID); err != nil { + return fmt.Errorf("cannot load organization: %w", err) + } + + if organization.LogoFileID == nil { + return fmt.Errorf("organization has no logo") + } + + var file coredata.File + if err := file.LoadByID(ctx, conn, scope, *organization.LogoFileID); err != nil { + return fmt.Errorf("cannot load logo file: %w", err) + } + + logoFile = &file + return nil + }, + ) + + if err != nil { + return nil, err + } + + return logoFile, nil +} + +func (s *TenantAuthService) InitiateDomainVerification( ctx context.Context, - tenantID gid.TenantID, organizationID gid.GID, emailDomain string, ) (*coredata.SAMLConfiguration, error) { - token, err := GenerateDomainVerificationToken() + token, err := generateDomainVerificationToken() if err != nil { return nil, fmt.Errorf("cannot generate verification token: %w", err) } var config *coredata.SAMLConfiguration - err = s.pg.WithTx(ctx, func(tx pg.Conn) error { - now := time.Now() - scope := coredata.NewScope(tenantID) + err = s.pg.WithTx( + ctx, + func(tx pg.Conn) error { + now := time.Now() - config = &coredata.SAMLConfiguration{ - ID: gid.New(tenantID, coredata.SAMLConfigurationEntityType), - OrganizationID: organizationID, - EmailDomain: emailDomain, - Enabled: false, - EnforcementPolicy: coredata.SAMLEnforcementPolicyOff, - DomainVerified: false, - DomainVerificationToken: &token, - // Default IdP values (placeholders until configured) - IdPEntityID: "not-configured", - IdPSsoURL: "not-configured", - IdPCertificate: "not-configured", - // Default attribute mappings - AttributeEmail: "http://schemas.xmlsoap.org/ws/2005/05/identity/claims/emailaddress", - AttributeFirstname: "http://schemas.xmlsoap.org/ws/2005/05/identity/claims/givenname", - AttributeLastname: "http://schemas.xmlsoap.org/ws/2005/05/identity/claims/surname", - AttributeRole: "http://schemas.xmlsoap.org/ws/2005/05/identity/claims/role", - AutoSignupEnabled: false, - CreatedAt: now, - UpdatedAt: now, - } + config = &coredata.SAMLConfiguration{ + ID: gid.New(s.scope.GetTenantID(), coredata.SAMLConfigurationEntityType), + OrganizationID: organizationID, + EmailDomain: emailDomain, + Enabled: false, + EnforcementPolicy: coredata.SAMLEnforcementPolicyOff, + DomainVerified: false, + DomainVerificationToken: &token, + IdPEntityID: "", + IdPSsoURL: "", + IdPCertificate: "", + AttributeEmail: "http://schemas.xmlsoap.org/ws/2005/05/identity/claims/emailaddress", + AttributeFirstname: "http://schemas.xmlsoap.org/ws/2005/05/identity/claims/givenname", + AttributeLastname: "http://schemas.xmlsoap.org/ws/2005/05/identity/claims/surname", + AttributeRole: "http://schemas.xmlsoap.org/ws/2005/05/identity/claims/role", + AutoSignupEnabled: false, + CreatedAt: now, + UpdatedAt: now, + } - if err := config.Insert(ctx, tx, scope); err != nil { - return fmt.Errorf("cannot insert SAML configuration: %w", err) - } + if err := config.Insert(ctx, tx, s.scope); err != nil { + return fmt.Errorf("cannot insert SAML configuration: %w", err) + } - return nil - }) + return nil + }, + ) if err != nil { return nil, err @@ -1166,54 +1153,51 @@ func (s Service) InitiateDomainVerification( return config, nil } -// VerifyDomain checks DNS TXT record and marks domain as verified if found -func (s Service) VerifyDomain( +func (s *TenantAuthService) VerifyDomain( ctx context.Context, - tenantID gid.TenantID, configID gid.GID, ) (*coredata.SAMLConfiguration, bool, error) { var config *coredata.SAMLConfiguration var verified bool - err := s.pg.WithTx(ctx, func(tx pg.Conn) error { - scope := coredata.NewScope(tenantID) - - // Load config - config = &coredata.SAMLConfiguration{} - if err := config.LoadByID(ctx, tx, scope, configID); err != nil { - return fmt.Errorf("cannot load SAML configuration: %w", err) - } - - if config.DomainVerificationToken == nil { - return fmt.Errorf("no verification token found for this configuration") - } - - if config.DomainVerified { - verified = true - return nil // Already verified - } - - // Check DNS TXT record - isVerified, err := VerifyDomainOwnership(ctx, config.EmailDomain, *config.DomainVerificationToken) - if err != nil { - return fmt.Errorf("cannot verify domain ownership: %w", err) - } - - verified = isVerified - - if isVerified { - now := time.Now() - config.DomainVerified = true - config.DomainVerifiedAt = &now - config.UpdatedAt = now - - if err := config.Update(ctx, tx, scope); err != nil { - return fmt.Errorf("cannot update SAML configuration: %w", err) + err := s.pg.WithTx( + ctx, + func(tx pg.Conn) error { + config = &coredata.SAMLConfiguration{} + if err := config.LoadByID(ctx, tx, s.scope, configID); err != nil { + return fmt.Errorf("cannot load SAML configuration: %w", err) } - } - return nil - }) + if config.DomainVerificationToken == nil { + return fmt.Errorf("no verification token found for this configuration") + } + + if config.DomainVerified { + verified = true + return nil + } + + isVerified, err := verifyDomainOwnership(ctx, config.EmailDomain, *config.DomainVerificationToken) + if err != nil { + return fmt.Errorf("cannot verify domain ownership: %w", err) + } + + verified = isVerified + + if isVerified { + now := time.Now() + config.DomainVerified = true + config.DomainVerifiedAt = &now + config.UpdatedAt = now + + if err := config.Update(ctx, tx, s.scope); err != nil { + return fmt.Errorf("cannot update SAML configuration: %w", err) + } + } + + return nil + }, + ) if err != nil { return nil, false, err @@ -1222,10 +1206,7 @@ func (s Service) VerifyDomain( return config, verified, nil } -// Domain Verification Methods - -// GenerateDomainVerificationToken generates a random 32-character hex token for domain verification -func GenerateDomainVerificationToken() (string, error) { +func generateDomainVerificationToken() (string, error) { bytes := make([]byte, 16) // 16 bytes = 32 hex characters if _, err := rand.Read(bytes); err != nil { return "", fmt.Errorf("cannot generate domain verification token: %w", err) @@ -1233,30 +1214,21 @@ func GenerateDomainVerificationToken() (string, error) { return hex.EncodeToString(bytes), nil } -// GetDomainVerificationRecord returns the DNS TXT record string that should be added to the domain func GetDomainVerificationRecord(token string) string { return fmt.Sprintf("probo-verification=%s", token) } -// VerifyDomainOwnership performs DNS lookup to verify domain ownership via TXT record -func VerifyDomainOwnership(ctx context.Context, domain, expectedToken string) (bool, error) { - // Use net package for DNS TXT record lookup +func verifyDomainOwnership(ctx context.Context, domain, expectedToken string) (bool, error) { var txtRecords []string var err error - // Create a DNS resolver with timeout from context - resolver := &net.Resolver{ - PreferGo: true, - } + resolver := &net.Resolver{PreferGo: true} txtRecords, err = resolver.LookupTXT(ctx, domain) if err != nil { - // DNS lookup errors are expected if the domain doesn't exist or has no TXT records - // We return false (not verified) but not an error, as this is a normal case return false, nil } - // Check if any TXT record matches our verification token expectedRecord := GetDomainVerificationRecord(expectedToken) for _, record := range txtRecords { if record == expectedRecord { @@ -1264,6 +1236,5 @@ func VerifyDomainOwnership(ctx context.Context, domain, expectedToken string) (b } } - // Token not found in DNS records return false, nil } diff --git a/pkg/authz/service.go b/pkg/authz/service.go index 0be246a17..b1b530a8a 100644 --- a/pkg/authz/service.go +++ b/pkg/authz/service.go @@ -84,24 +84,22 @@ func (s *Service) WithTenant(tenantID gid.TenantID) *TenantAuthzService { } } -// This method is on Service (not TenantAuthzService) because it operates across tenants -// and doesn't require tenant-scoped access. func (s *Service) GetAllUserOrganizations( ctx context.Context, userID gid.GID, -) ([]*coredata.Organization, error) { - var organizations []*coredata.Organization +) (coredata.Organizations, error) { + organizations := coredata.Organizations{} - err := s.pg.WithConn(ctx, func(conn pg.Conn) error { - var organizationList coredata.Organizations - if err := organizationList.LoadAllByUserID(ctx, conn, userID); err != nil { - return fmt.Errorf("cannot load user organizations: %w", err) - } + err := s.pg.WithConn( + ctx, + func(conn pg.Conn) error { + if err := organizations.LoadAllByUserID(ctx, conn, userID); err != nil { + return fmt.Errorf("cannot load user organizations: %w", err) + } - organizations = organizationList - - return nil - }) + return nil + }, + ) return organizations, err } @@ -110,13 +108,13 @@ func (s *Service) GetUserOrganizations( ctx context.Context, userID gid.GID, cursor *page.Cursor[coredata.OrganizationOrderField], -) ([]*coredata.Organization, error) { - var organizations coredata.Organizations +) (coredata.Organizations, error) { + organizations := coredata.Organizations{} err := s.pg.WithConn( ctx, func(conn pg.Conn) error { - if err := organizations.LoadByUserID(ctx, conn, userID, cursor); err != nil { + if err := organizations.LoadByUserID(ctx, conn, coredata.NewNoScope(), userID, cursor); err != nil { return fmt.Errorf("cannot load user organizations: %w", err) } return nil @@ -268,7 +266,7 @@ func (s *Service) GetUserInvitations( err := s.pg.WithConn( ctx, func(conn pg.Conn) error { - if err := invitations.LoadByEmail(ctx, conn, email, cursor, filter); err != nil { + if err := invitations.LoadByEmail(ctx, conn, coredata.NewNoScope(), email, cursor, filter); err != nil { return fmt.Errorf("cannot load invitations: %w", err) } @@ -282,44 +280,107 @@ func (s *Service) GetUserInvitations( return page.NewPage(invitations, cursor), nil } -func (s *Service) CountUserInvitations( +type UserInvitation struct { + ID gid.GID + Email string + FullName string + Role string + ExpiresAt time.Time + AcceptedAt *time.Time + CreatedAt time.Time + OrganizationID gid.GID + Organization OrganizationSummary +} + +type OrganizationSummary struct { + ID gid.GID + Name string +} + +func (s *Service) GetUserPendingInvitations( ctx context.Context, email string, - filter *coredata.InvitationFilter, -) (int, error) { - var count int +) ([]*UserInvitation, error) { + userInvitations := []*UserInvitation{} err := s.pg.WithConn( ctx, func(conn pg.Conn) error { - var invitations coredata.Invitations - var err error - count, err = invitations.CountByEmail(ctx, conn, email, filter) - return err + cursor := page.NewCursor( + 1000, + nil, + page.Head, + page.OrderBy[coredata.InvitationOrderField]{ + Field: coredata.InvitationOrderFieldCreatedAt, + Direction: page.OrderDirectionDesc, + }, + ) + filter := coredata.NewInvitationFilter([]coredata.InvitationStatus{coredata.InvitationStatusPending}) + invitations := coredata.Invitations{} + + if err := invitations.LoadByEmail(ctx, conn, coredata.NewNoScope(), email, cursor, filter); err != nil { + return fmt.Errorf("cannot load invitations: %w", err) + } + + organizationIDs := []gid.GID{} + for _, invitation := range invitations { + organizationIDs = append(organizationIDs, invitation.OrganizationID) + } + + organizations := coredata.Organizations{} + if err := organizations.BatchLoadByID(ctx, conn, coredata.NewNoScope(), organizationIDs); err != nil { + return fmt.Errorf("cannot load organizations: %w", err) + } + + for _, invitation := range invitations { + userInvitation := &UserInvitation{ + ID: invitation.ID, + Email: invitation.Email, + FullName: invitation.FullName, + Role: invitation.Role, + ExpiresAt: invitation.ExpiresAt, + AcceptedAt: invitation.AcceptedAt, + CreatedAt: invitation.CreatedAt, + OrganizationID: invitation.OrganizationID, + } + + for _, org := range organizations { + if org.ID == invitation.OrganizationID { + userInvitation.Organization = OrganizationSummary{ + ID: org.ID, + Name: org.Name, + } + } + } + + userInvitations = append(userInvitations, userInvitation) + } + + return nil }, ) - return count, err + if err != nil { + return nil, err + } + + return userInvitations, nil } -// This method is on Service (not TenantAuthzService) because the user viewing -// the invitation organization doesn't have tenant access yet. -func (s *Service) GetOrganizationByInvitationID( +func (s *TenantAuthzService) GetOrganizationByInvitationID( ctx context.Context, invitationID gid.GID, ) (*coredata.Organization, error) { - scope := coredata.NewScope(invitationID.TenantID()) - var organization coredata.Organization err := s.pg.WithConn( ctx, func(conn pg.Conn) error { var invitation coredata.Invitation - if err := invitation.LoadByID(ctx, conn, scope, invitationID); err != nil { + if err := invitation.LoadByID(ctx, conn, s.scope, invitationID); err != nil { return fmt.Errorf("cannot load invitation: %w", err) } - if err := organization.LoadByID(ctx, conn, scope, invitation.OrganizationID); err != nil { + if err := organization.LoadByID(ctx, conn, s.scope, invitation.OrganizationID); err != nil { return fmt.Errorf("cannot load organization: %w", err) } @@ -333,18 +394,14 @@ func (s *Service) GetOrganizationByInvitationID( return &organization, nil } -// This method is on Service (not TenantAuthzService) because the user added to the organization -// doesn't have tenant access yet -func (s *Service) AddUserToOrganization( +func (s *TenantAuthzService) AddUserToOrganization( ctx context.Context, userID gid.GID, orgID gid.GID, role string, ) error { now := time.Now() - tenantID := orgID.TenantID() - membershipID := gid.New(tenantID, coredata.MembershipEntityType) - scope := coredata.NewScope(tenantID) + membershipID := gid.New(s.scope.GetTenantID(), coredata.MembershipEntityType) membership := &coredata.Membership{ ID: membershipID, @@ -358,7 +415,7 @@ func (s *Service) AddUserToOrganization( return s.pg.WithConn( ctx, func(conn pg.Conn) error { - if err := membership.Create(ctx, conn, scope); err != nil { + if err := membership.Create(ctx, conn, s.scope); err != nil { return fmt.Errorf("cannot add user to organization: %w", err) } return nil @@ -384,6 +441,7 @@ func (s *TenantAuthzService) GetInvitationsByOrganizationID( return nil }, ) + if err != nil { return nil, err } @@ -397,17 +455,23 @@ func (s *TenantAuthzService) CountOrganizationInvitations( filter *coredata.InvitationFilter, ) (int, error) { var count int + err := s.pg.WithConn( ctx, - func(conn pg.Conn) error { + func(conn pg.Conn) (err error) { var invitations coredata.Invitations - var err error + count, err = invitations.CountByOrganizationID(ctx, conn, s.scope, orgID, filter) - return err + if err != nil { + return fmt.Errorf("cannot count organization invitations: %w", err) + } + + return nil }, ) + if err != nil { - return 0, fmt.Errorf("cannot count invitations: %w", err) + return 0, err } return count, nil @@ -430,6 +494,7 @@ func (s *TenantAuthzService) GetInvitationByID( if err != nil { return nil, err } + return invitation, nil } @@ -504,13 +569,18 @@ func (s *TenantAuthzService) CountOrganizationUsers( orgID gid.GID, ) (int, error) { var count int + err := s.pg.WithConn( ctx, - func(conn pg.Conn) error { + func(conn pg.Conn) (err error) { var users coredata.Users - var err error + count, err = users.CountByOrganizationID(ctx, conn, s.scope, orgID) - return err + if err != nil { + return fmt.Errorf("cannot count organization users: %w", err) + } + + return nil }, ) if err != nil { @@ -636,93 +706,96 @@ func (s *TenantAuthzService) InviteUserToOrganization( ) (*coredata.Invitation, error) { var invitation *coredata.Invitation - err := s.pg.WithTx(ctx, func(tx pg.Conn) error { - user := &coredata.User{} - userExists := true - if err := user.LoadByEmail(ctx, tx, emailAddress); err != nil { - var userNotFound *coredata.ErrUserNotFound - if errors.As(err, &userNotFound) { - userExists = false - } else { - return fmt.Errorf("cannot check if user exists: %w", err) + err := s.pg.WithTx( + ctx, + func(tx pg.Conn) error { + user := &coredata.User{} + userExists := true + if err := user.LoadByEmail(ctx, tx, emailAddress); err != nil { + var userNotFound *coredata.ErrUserNotFound + if errors.As(err, &userNotFound) { + userExists = false + } else { + return fmt.Errorf("cannot check if user exists: %w", err) + } } - } - organization := &coredata.Organization{} - if err := organization.LoadByID(ctx, tx, s.scope, organizationID); err != nil { - return fmt.Errorf("cannot load organization: %w", err) - } + organization := &coredata.Organization{} + if err := organization.LoadByID(ctx, tx, s.scope, organizationID); err != nil { + return fmt.Errorf("cannot load organization: %w", err) + } - invitationID := gid.New(s.scope.GetTenantID(), coredata.InvitationEntityType) - now := time.Now() - invitation = &coredata.Invitation{ - ID: invitationID, - OrganizationID: organizationID, - Email: emailAddress, - FullName: fullName, - Role: role, - ExpiresAt: now.Add(s.invitationTokenValidity), - CreatedAt: now, - } - - var err error - var invitationURL string - var recipientName string - - if userExists { - recipientName = user.FullName - invitationURL = fmt.Sprintf("https://%s/", s.hostname) - } else { - recipientName = fullName - invitationData := coredata.InvitationData{ - InvitationID: invitationID, + invitationID := gid.New(s.scope.GetTenantID(), coredata.InvitationEntityType) + now := time.Now() + invitation = &coredata.Invitation{ + ID: invitationID, OrganizationID: organizationID, Email: emailAddress, FullName: fullName, Role: role, + ExpiresAt: now.Add(s.invitationTokenValidity), + CreatedAt: now, } - invitationToken, err := statelesstoken.NewToken( - s.tokenSecret, - TokenTypeOrganizationInvitation, - s.invitationTokenValidity, - invitationData, + var err error + var invitationURL string + var recipientName string + + if userExists { + recipientName = user.FullName + invitationURL = fmt.Sprintf("https://%s/", s.hostname) + } else { + recipientName = fullName + invitationData := coredata.InvitationData{ + InvitationID: invitationID, + OrganizationID: organizationID, + Email: emailAddress, + FullName: fullName, + Role: role, + } + + invitationToken, err := statelesstoken.NewToken( + s.tokenSecret, + TokenTypeOrganizationInvitation, + s.invitationTokenValidity, + invitationData, + ) + if err != nil { + return fmt.Errorf("cannot generate invitation token: %w", err) + } + + invitationURL = fmt.Sprintf("https://%s/auth/signup-from-invitation?token=%s&fullName=%s", s.hostname, invitationToken, url.QueryEscape(fullName)) + } + + subject, textBody, htmlBody, err := emails.RenderInvitation( + s.hostname, + recipientName, + organization.Name, + invitationURL, ) if err != nil { - return fmt.Errorf("cannot generate invitation token: %w", err) + return fmt.Errorf("cannot render invitation email: %w", err) } - invitationURL = fmt.Sprintf("https://%s/auth/signup-from-invitation?token=%s&fullName=%s", s.hostname, invitationToken, url.QueryEscape(fullName)) - } + email := coredata.NewEmail( + fullName, + emailAddress, + subject, + textBody, + htmlBody, + ) - subject, textBody, htmlBody, err := emails.RenderInvitation( - s.hostname, - recipientName, - organization.Name, - invitationURL, - ) - if err != nil { - return fmt.Errorf("cannot render invitation email: %w", err) - } + if err := email.Insert(ctx, tx); err != nil { + return fmt.Errorf("cannot insert email: %w", err) + } - email := coredata.NewEmail( - fullName, - emailAddress, - subject, - textBody, - htmlBody, - ) + if err := invitation.Create(ctx, tx, s.scope); err != nil { + return fmt.Errorf("cannot create invitation: %w", err) + } - if err := email.Insert(ctx, tx); err != nil { - return fmt.Errorf("cannot insert email: %w", err) - } - - if err := invitation.Create(ctx, tx, s.scope); err != nil { - return fmt.Errorf("cannot create invitation: %w", err) - } - - return nil - }) + return nil + }, + ) if err != nil { return nil, err @@ -731,8 +804,6 @@ func (s *TenantAuthzService) InviteUserToOrganization( return invitation, nil } -// EnsureSAMLMembership creates or updates a user's membership in an organization. -// This is used during SAML authentication to ensure the user has the correct role. func (s *TenantAuthzService) EnsureSAMLMembership( ctx context.Context, userID gid.GID, diff --git a/pkg/coredata/file.go b/pkg/coredata/file.go index b8a519b25..ac926fdf5 100644 --- a/pkg/coredata/file.go +++ b/pkg/coredata/file.go @@ -153,3 +153,52 @@ WHERE %s return err } + +// LoadFilesByIDs loads multiple files by their IDs in a single query +// Returns a map of file ID to File for efficient lookup +func LoadFilesByIDs( + ctx context.Context, + conn pg.Conn, + fileIDs []gid.GID, +) (map[gid.GID]*File, error) { + if len(fileIDs) == 0 { + return make(map[gid.GID]*File), nil + } + + q := ` +SELECT + id, + bucket_name, + mime_type, + file_name, + file_key, + file_size, + created_at, + updated_at, + deleted_at +FROM + files +WHERE + id = ANY(@file_ids) + AND deleted_at IS NULL +` + + args := pgx.StrictNamedArgs{"file_ids": fileIDs} + + rows, err := conn.Query(ctx, q, args) + if err != nil { + return nil, fmt.Errorf("cannot query files: %w", err) + } + + files, err := pgx.CollectRows(rows, pgx.RowToStructByName[File]) + if err != nil { + return nil, fmt.Errorf("cannot collect files: %w", err) + } + + result := make(map[gid.GID]*File, len(files)) + for i := range files { + result[files[i].ID] = &files[i] + } + + return result, nil +} diff --git a/pkg/coredata/invitation.go b/pkg/coredata/invitation.go index 1f7bcbc78..7636eba48 100644 --- a/pkg/coredata/invitation.go +++ b/pkg/coredata/invitation.go @@ -237,11 +237,10 @@ WHERE return nil } -// Tenant scope is not applied because this is used to query invitations across all tenants -// for a user who doesn't have tenant access yet (before accepting an invitation). func (i *Invitations) LoadByEmail( ctx context.Context, conn pg.Conn, + scope Scoper, email string, cursor *page.Cursor[InvitationOrderField], filter *InvitationFilter, diff --git a/pkg/coredata/organization.go b/pkg/coredata/organization.go index 47f7d215a..a96eb910f 100644 --- a/pkg/coredata/organization.go +++ b/pkg/coredata/organization.go @@ -106,10 +106,10 @@ LIMIT 1; return nil } -// Tenant id scope is not applied in this functions because we want to access all user's organizations. func (o *Organizations) LoadByUserID( ctx context.Context, conn pg.Conn, + scope Scoper, userID gid.GID, cursor *page.Cursor[OrganizationOrderField], ) error { @@ -140,10 +140,11 @@ FROM INNER JOIN user_org ON organizations.id = user_org.organization_id WHERE - %s + %S + AND %s ` - q = fmt.Sprintf(q, cursor.SQLFragment()) + q = fmt.Sprintf(q, scope.SQLFragment(), cursor.SQLFragment()) args := pgx.StrictNamedArgs{"user_id": userID} maps.Copy(args, cursor.SQLArguments()) @@ -163,7 +164,6 @@ WHERE return nil } -// Tenant id scope is not applied in this function because we want to access all user's organizations. func (o *Organizations) LoadAllByUserID( ctx context.Context, conn pg.Conn, @@ -379,3 +379,50 @@ LIMIT 1 return nil } + +func (o *Organizations) BatchLoadByID( + ctx context.Context, + conn pg.Conn, + scope Scoper, + organizationIDs []gid.GID, +) error { + q := ` +SELECT + tenant_id, + id, + name, + logo_file_id, + horizontal_logo_file_id, + description, + website_url, + email, + headquarter_address, + custom_domain_id, + created_at, + updated_at +FROM + organizations +WHERE + %s + AND id = ANY(@organization_ids) +` + + q = fmt.Sprintf(q, scope.SQLFragment()) + + args := pgx.StrictNamedArgs{"organization_ids": organizationIDs} + maps.Copy(args, scope.SQLArguments()) + + rows, err := conn.Query(ctx, q, args) + if err != nil { + return fmt.Errorf("cannot query organizations: %w", err) + } + + organizations, err := pgx.CollectRows(rows, pgx.RowToAddrOfStructByName[Organization]) + if err != nil { + return fmt.Errorf("cannot collect organizations: %w", err) + } + + *o = organizations + + return nil +} diff --git a/pkg/coredata/saml_configuration.go b/pkg/coredata/saml_configuration.go index 68e435730..aa2bf23f1 100644 --- a/pkg/coredata/saml_configuration.go +++ b/pkg/coredata/saml_configuration.go @@ -441,3 +441,66 @@ ORDER BY created_at ASC; return result, nil } +// LoadSAMLConfigurationsByOrganizationIDsAndEmailDomain loads SAML configurations for multiple organizations +// and a given email domain in a single query. This is used to avoid N+1 queries. +func LoadSAMLConfigurationsByOrganizationIDsAndEmailDomain( + ctx context.Context, + conn pg.Conn, + organizationIDs []gid.GID, + emailDomain string, +) (map[gid.GID]*SAMLConfiguration, error) { + if len(organizationIDs) == 0 { + return make(map[gid.GID]*SAMLConfiguration), nil + } + + q := ` +SELECT + id, + organization_id, + email_domain, + enabled, + enforcement_policy, + idp_entity_id, + idp_sso_url, + idp_certificate, + idp_metadata_url, + attribute_email, + attribute_firstname, + attribute_lastname, + attribute_role, + auto_signup_enabled, + domain_verified, + domain_verification_token, + domain_verified_at, + created_at, + updated_at +FROM + auth_saml_configurations +WHERE + organization_id = ANY(@organization_ids) + AND email_domain = @email_domain +` + + args := pgx.StrictNamedArgs{ + "organization_ids": organizationIDs, + "email_domain": emailDomain, + } + + rows, err := conn.Query(ctx, q, args) + if err != nil { + return nil, fmt.Errorf("cannot query auth_saml_configurations: %w", err) + } + + configs, err := pgx.CollectRows(rows, pgx.RowToStructByName[SAMLConfiguration]) + if err != nil { + return nil, fmt.Errorf("cannot collect saml_configurations: %w", err) + } + + result := make(map[gid.GID]*SAMLConfiguration, len(configs)) + for i := range configs { + result[configs[i].OrganizationID] = &configs[i] + } + + return result, nil +} + diff --git a/pkg/server/api/console/v1/resolver.go b/pkg/server/api/console/v1/resolver.go index 37f0b76b8..b3fc8e5f9 100644 --- a/pkg/server/api/console/v1/resolver.go +++ b/pkg/server/api/console/v1/resolver.go @@ -410,6 +410,10 @@ func (r *Resolver) AuthzService(ctx context.Context, tenantID gid.TenantID) *aut return GetTenantAuthzService(ctx, r.authzSvc, tenantID) } +func (r *Resolver) AuthService(ctx context.Context, tenantID gid.TenantID) *auth.TenantAuthService { + return GetTenantAuthService(ctx, r.authSvc, tenantID) +} + func UnwrapOmittable[T any](field graphql.Omittable[T]) *T { if !field.IsSet() { return nil @@ -428,6 +432,11 @@ func GetTenantAuthzService(ctx context.Context, authzSvc *authz.Service, tenantI return authzSvc.WithTenant(tenantID) } +func GetTenantAuthService(ctx context.Context, authSvc *auth.Service, tenantID gid.TenantID) *auth.TenantAuthService { + validateTenantAccess(ctx, tenantID) + return authSvc.WithTenant(tenantID) +} + func validateTenantAccess(ctx context.Context, tenantID gid.TenantID) { access, _ := ctx.Value(userTenantContextKey).(*userTenantAccess) diff --git a/pkg/server/api/console/v1/schema.graphql b/pkg/server/api/console/v1/schema.graphql index 6569341bd..16838acce 100644 --- a/pkg/server/api/console/v1/schema.graphql +++ b/pkg/server/api/console/v1/schema.graphql @@ -2402,15 +2402,6 @@ type Viewer { before: CursorKey orderBy: OrganizationOrder ): OrganizationConnection! @goField(forceResolver: true) - - invitations( - first: Int - after: CursorKey - last: Int - before: CursorKey - orderBy: InvitationOrder - filter: InvitationFilter - ): InvitationConnection! @goField(forceResolver: true) } # Connection Types diff --git a/pkg/server/api/console/v1/schema/schema.go b/pkg/server/api/console/v1/schema/schema.go index a5e259b28..ee82a7d1e 100644 --- a/pkg/server/api/console/v1/schema/schema.go +++ b/pkg/server/api/console/v1/schema/schema.go @@ -1749,7 +1749,6 @@ type ComplexityRoot struct { Viewer struct { ID func(childComplexity int) int - Invitations func(childComplexity int, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.InvitationOrder, filter *types.InvitationFilter) int Organizations func(childComplexity int, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.OrganizationOrder) int User func(childComplexity int) int } @@ -2172,7 +2171,6 @@ type VendorServiceResolver interface { } type ViewerResolver interface { Organizations(ctx context.Context, obj *types.Viewer, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.OrganizationOrder) (*types.OrganizationConnection, error) - Invitations(ctx context.Context, obj *types.Viewer, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.InvitationOrder, filter *types.InvitationFilter) (*types.InvitationConnection, error) } type executableSchema struct { @@ -9422,18 +9420,6 @@ func (e *executableSchema) Complexity(ctx context.Context, typeName, field strin return e.complexity.Viewer.ID(childComplexity), true - case "Viewer.invitations": - if e.complexity.Viewer.Invitations == nil { - break - } - - args, err := ec.field_Viewer_invitations_args(ctx, rawArgs) - if err != nil { - return 0, false - } - - return e.complexity.Viewer.Invitations(childComplexity, args["first"].(*int), args["after"].(*page.CursorKey), args["last"].(*int), args["before"].(*page.CursorKey), args["orderBy"].(*types.InvitationOrder), args["filter"].(*types.InvitationFilter)), true - case "Viewer.organizations": if e.complexity.Viewer.Organizations == nil { break @@ -12142,15 +12128,6 @@ type Viewer { before: CursorKey orderBy: OrganizationOrder ): OrganizationConnection! @goField(forceResolver: true) - - invitations( - first: Int - after: CursorKey - last: Int - before: CursorKey - orderBy: InvitationOrder - filter: InvitationFilter - ): InvitationConnection! @goField(forceResolver: true) } # Connection Types @@ -22840,119 +22817,6 @@ func (ec *executionContext) field_Vendor_services_argsOrderBy( return zeroVal, nil } -func (ec *executionContext) field_Viewer_invitations_args(ctx context.Context, rawArgs map[string]any) (map[string]any, error) { - var err error - args := map[string]any{} - arg0, err := ec.field_Viewer_invitations_argsFirst(ctx, rawArgs) - if err != nil { - return nil, err - } - args["first"] = arg0 - arg1, err := ec.field_Viewer_invitations_argsAfter(ctx, rawArgs) - if err != nil { - return nil, err - } - args["after"] = arg1 - arg2, err := ec.field_Viewer_invitations_argsLast(ctx, rawArgs) - if err != nil { - return nil, err - } - args["last"] = arg2 - arg3, err := ec.field_Viewer_invitations_argsBefore(ctx, rawArgs) - if err != nil { - return nil, err - } - args["before"] = arg3 - arg4, err := ec.field_Viewer_invitations_argsOrderBy(ctx, rawArgs) - if err != nil { - return nil, err - } - args["orderBy"] = arg4 - arg5, err := ec.field_Viewer_invitations_argsFilter(ctx, rawArgs) - if err != nil { - return nil, err - } - args["filter"] = arg5 - return args, nil -} -func (ec *executionContext) field_Viewer_invitations_argsFirst( - ctx context.Context, - rawArgs map[string]any, -) (*int, error) { - ctx = graphql.WithPathContext(ctx, graphql.NewPathWithField("first")) - if tmp, ok := rawArgs["first"]; ok { - return ec.unmarshalOInt2ᚖint(ctx, tmp) - } - - var zeroVal *int - return zeroVal, nil -} - -func (ec *executionContext) field_Viewer_invitations_argsAfter( - ctx context.Context, - rawArgs map[string]any, -) (*page.CursorKey, error) { - ctx = graphql.WithPathContext(ctx, graphql.NewPathWithField("after")) - if tmp, ok := rawArgs["after"]; ok { - return ec.unmarshalOCursorKey2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋpageᚐCursorKey(ctx, tmp) - } - - var zeroVal *page.CursorKey - return zeroVal, nil -} - -func (ec *executionContext) field_Viewer_invitations_argsLast( - ctx context.Context, - rawArgs map[string]any, -) (*int, error) { - ctx = graphql.WithPathContext(ctx, graphql.NewPathWithField("last")) - if tmp, ok := rawArgs["last"]; ok { - return ec.unmarshalOInt2ᚖint(ctx, tmp) - } - - var zeroVal *int - return zeroVal, nil -} - -func (ec *executionContext) field_Viewer_invitations_argsBefore( - ctx context.Context, - rawArgs map[string]any, -) (*page.CursorKey, error) { - ctx = graphql.WithPathContext(ctx, graphql.NewPathWithField("before")) - if tmp, ok := rawArgs["before"]; ok { - return ec.unmarshalOCursorKey2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋpageᚐCursorKey(ctx, tmp) - } - - var zeroVal *page.CursorKey - return zeroVal, nil -} - -func (ec *executionContext) field_Viewer_invitations_argsOrderBy( - ctx context.Context, - rawArgs map[string]any, -) (*types.InvitationOrder, error) { - ctx = graphql.WithPathContext(ctx, graphql.NewPathWithField("orderBy")) - if tmp, ok := rawArgs["orderBy"]; ok { - return ec.unmarshalOInvitationOrder2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐInvitationOrder(ctx, tmp) - } - - var zeroVal *types.InvitationOrder - return zeroVal, nil -} - -func (ec *executionContext) field_Viewer_invitations_argsFilter( - ctx context.Context, - rawArgs map[string]any, -) (*types.InvitationFilter, error) { - ctx = graphql.WithPathContext(ctx, graphql.NewPathWithField("filter")) - if tmp, ok := rawArgs["filter"]; ok { - return ec.unmarshalOInvitationFilter2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐInvitationFilter(ctx, tmp) - } - - var zeroVal *types.InvitationFilter - return zeroVal, nil -} - func (ec *executionContext) field_Viewer_organizations_args(ctx context.Context, rawArgs map[string]any) (map[string]any, error) { var err error args := map[string]any{} @@ -54498,8 +54362,6 @@ func (ec *executionContext) fieldContext_Query_viewer(_ context.Context, field g return ec.fieldContext_Viewer_user(ctx, field) case "organizations": return ec.fieldContext_Viewer_organizations(ctx, field) - case "invitations": - return ec.fieldContext_Viewer_invitations(ctx, field) } return nil, fmt.Errorf("no field named %q was found under type Viewer", field.Name) }, @@ -71204,69 +71066,6 @@ func (ec *executionContext) fieldContext_Viewer_organizations(ctx context.Contex return fc, nil } -func (ec *executionContext) _Viewer_invitations(ctx context.Context, field graphql.CollectedField, obj *types.Viewer) (ret graphql.Marshaler) { - fc, err := ec.fieldContext_Viewer_invitations(ctx, field) - if err != nil { - return graphql.Null - } - ctx = graphql.WithFieldContext(ctx, fc) - defer func() { - if r := recover(); r != nil { - ec.Error(ctx, ec.Recover(ctx, r)) - ret = graphql.Null - } - }() - resTmp, err := ec.ResolverMiddleware(ctx, func(rctx context.Context) (any, error) { - ctx = rctx // use context from middleware stack in children - return ec.resolvers.Viewer().Invitations(rctx, obj, fc.Args["first"].(*int), fc.Args["after"].(*page.CursorKey), fc.Args["last"].(*int), fc.Args["before"].(*page.CursorKey), fc.Args["orderBy"].(*types.InvitationOrder), fc.Args["filter"].(*types.InvitationFilter)) - }) - if err != nil { - ec.Error(ctx, err) - return graphql.Null - } - if resTmp == nil { - if !graphql.HasFieldError(ctx, fc) { - ec.Errorf(ctx, "must not be null") - } - return graphql.Null - } - res := resTmp.(*types.InvitationConnection) - fc.Result = res - return ec.marshalNInvitationConnection2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐInvitationConnection(ctx, field.Selections, res) -} - -func (ec *executionContext) fieldContext_Viewer_invitations(ctx context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { - fc = &graphql.FieldContext{ - Object: "Viewer", - Field: field, - IsMethod: true, - IsResolver: true, - Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { - switch field.Name { - case "totalCount": - return ec.fieldContext_InvitationConnection_totalCount(ctx, field) - case "edges": - return ec.fieldContext_InvitationConnection_edges(ctx, field) - case "pageInfo": - return ec.fieldContext_InvitationConnection_pageInfo(ctx, field) - } - return nil, fmt.Errorf("no field named %q was found under type InvitationConnection", field.Name) - }, - } - defer func() { - if r := recover(); r != nil { - err = ec.Recover(ctx, r) - ec.Error(ctx, err) - } - }() - ctx = graphql.WithFieldContext(ctx, fc) - if fc.Args, err = ec.field_Viewer_invitations_args(ctx, field.ArgumentMap(ec.Variables)); err != nil { - ec.Error(ctx, err) - return fc, err - } - return fc, nil -} - func (ec *executionContext) ___Directive_name(ctx context.Context, field graphql.CollectedField, obj *introspection.Directive) (ret graphql.Marshaler) { fc, err := ec.fieldContext___Directive_name(ctx, field) if err != nil { @@ -98483,42 +98282,6 @@ func (ec *executionContext) _Viewer(ctx context.Context, sel ast.SelectionSet, o continue } - out.Concurrently(i, func(ctx context.Context) graphql.Marshaler { return innerFunc(ctx, out) }) - case "invitations": - field := field - - innerFunc := func(ctx context.Context, fs *graphql.FieldSet) (res graphql.Marshaler) { - defer func() { - if r := recover(); r != nil { - ec.Error(ctx, ec.Recover(ctx, r)) - } - }() - res = ec._Viewer_invitations(ctx, field, obj) - if res == graphql.Null { - atomic.AddUint32(&fs.Invalids, 1) - } - return res - } - - if field.Deferrable != nil { - dfs, ok := deferred[field.Deferrable.Label] - di := 0 - if ok { - dfs.AddField(field) - di = len(dfs.Values) - 1 - } else { - dfs = graphql.NewFieldSet([]graphql.CollectedField{field}) - deferred[field.Deferrable.Label] = dfs - } - dfs.Concurrently(di, func(ctx context.Context) graphql.Marshaler { - return innerFunc(ctx, dfs) - }) - - // don't run the out.Concurrently() call below - out.Values[i] = graphql.Null - continue - } - out.Concurrently(i, func(ctx context.Context) graphql.Marshaler { return innerFunc(ctx, out) }) default: panic("unknown field " + strconv.Quote(field.Name)) diff --git a/pkg/server/api/console/v1/types/types.go b/pkg/server/api/console/v1/types/types.go index 94340d7a6..b2bd7cda7 100644 --- a/pkg/server/api/console/v1/types/types.go +++ b/pkg/server/api/console/v1/types/types.go @@ -2480,5 +2480,4 @@ type Viewer struct { ID gid.GID `json:"id"` User *User `json:"user"` Organizations *OrganizationConnection `json:"organizations"` - Invitations *InvitationConnection `json:"invitations"` } diff --git a/pkg/server/api/console/v1/v1_resolver.go b/pkg/server/api/console/v1/v1_resolver.go index 4bf4b0105..12bcc6888 100644 --- a/pkg/server/api/console/v1/v1_resolver.go +++ b/pkg/server/api/console/v1/v1_resolver.go @@ -901,7 +901,9 @@ func (r *frameworkConnectionResolver) TotalCount(ctx context.Context, obj *types // Organization is the resolver for the organization field. func (r *invitationResolver) Organization(ctx context.Context, obj *types.Invitation) (*types.Organization, error) { - organization, err := r.authzSvc.GetOrganizationByInvitationID(ctx, obj.ID) + authz := r.AuthzService(ctx, obj.ID.TenantID()) + + organization, err := authz.GetOrganizationByInvitationID(ctx, obj.ID) if err != nil { panic(fmt.Errorf("cannot load organization: %w", err)) } @@ -918,28 +920,12 @@ func (r *invitationConnectionResolver) TotalCount(ctx context.Context, obj *type invitationFilter = coredata.NewInvitationFilter(obj.Filter.Statuses) } - authzSvc := r.AuthzService(ctx, obj.ParentID.TenantID()) - count, err := authzSvc.CountOrganizationInvitations(ctx, obj.ParentID, invitationFilter) + authz := r.AuthzService(ctx, obj.ParentID.TenantID()) + count, err := authz.CountOrganizationInvitations(ctx, obj.ParentID, invitationFilter) if err != nil { panic(fmt.Errorf("cannot count organization invitations: %w", err)) } return count, nil - case *viewerResolver: - user := UserFromContext(ctx) - if user == nil { - panic(fmt.Errorf("no authenticated user")) - } - - invitationFilter := coredata.NewInvitationFilter(nil) - if obj.Filter != nil { - invitationFilter = coredata.NewInvitationFilter(obj.Filter.Statuses) - } - - count, err := r.authzSvc.CountUserInvitations(ctx, user.EmailAddress, invitationFilter) - if err != nil { - panic(fmt.Errorf("cannot count user invitations: %w", err)) - } - return count, nil } panic(fmt.Errorf("unsupported resolver: %T", obj.Resolver)) @@ -1090,7 +1076,9 @@ func (r *membershipResolver) AuthMethod(ctx context.Context, obj *types.Membersh return coredata.UserAuthMethodPassword, nil } - authMethod, err := r.authSvc.GetUserAuthMethod(ctx, coredata.NewScope(obj.UserID.TenantID()), obj.UserID, obj.OrganizationID, session) + auth := r.AuthService(ctx, obj.UserID.TenantID()) + + authMethod, err := auth.GetUserAuthMethod(ctx, obj.UserID, obj.OrganizationID, session) if err != nil { return "", fmt.Errorf("cannot get user auth method: %w", err) } @@ -1101,22 +1089,26 @@ func (r *membershipResolver) AuthMethod(ctx context.Context, obj *types.Membersh func (r *membershipConnectionResolver) TotalCount(ctx context.Context, obj *types.MembershipConnection) (int, error) { switch obj.Resolver.(type) { case *organizationResolver: - authzSvc := r.AuthzService(ctx, obj.ParentID.TenantID()) - count, err := authzSvc.CountOrganizationMemberships(ctx, obj.ParentID) + authz := r.AuthzService(ctx, obj.ParentID.TenantID()) + count, err := authz.CountOrganizationMemberships(ctx, obj.ParentID) if err != nil { panic(fmt.Errorf("cannot count organization memberships: %w", err)) } + return count, nil - default: - panic(fmt.Errorf("unknown resolver type for membership connection")) } + + panic(fmt.Errorf("unknown resolver type for membership connection")) } // CreateOrganization is the resolver for the createOrganization field. func (r *mutationResolver) CreateOrganization(ctx context.Context, input types.CreateOrganizationInput) (*types.CreateOrganizationPayload, error) { currentUser := UserFromContext(ctx) - prb := r.proboSvc.WithTenant(gid.NewTenantID()) + tenantID := gid.NewTenantID() + + prb := r.proboSvc.WithTenant(tenantID) + authz := r.authzSvc.WithTenant(tenantID) organization, err := prb.Organizations.Create( ctx, @@ -1128,21 +1120,16 @@ func (r *mutationResolver) CreateOrganization(ctx context.Context, input types.C return nil, fmt.Errorf("cannot create organization: %w", err) } - err = r.authzSvc.AddUserToOrganization( + authz.AddUserToOrganization( ctx, currentUser.ID, organization.ID, - string(authz.RoleMember), + "MEMBER", ) if err != nil { return nil, fmt.Errorf("cannot add user to organization: %w", err) } - tenantIDs, _ := ctx.Value(userTenantContextKey).(*[]gid.TenantID) - *tenantIDs = append(*tenantIDs, organization.ID.TenantID()) - - prb = r.ProboService(ctx, organization.ID.TenantID()) - _, err = prb.Peoples.Create( ctx, probo.CreatePeopleRequest{ @@ -1158,6 +1145,10 @@ func (r *mutationResolver) CreateOrganization(ctx context.Context, input types.C return nil, fmt.Errorf("cannot create people: %w", err) } + // Append tenant to allowed one + tenantIDs, _ := ctx.Value(userTenantContextKey).(*[]gid.TenantID) + *tenantIDs = append(*tenantIDs, organization.ID.TenantID()) + return &types.CreateOrganizationPayload{ OrganizationEdge: types.NewOrganizationEdge(organization, coredata.OrganizationOrderFieldCreatedAt), }, nil @@ -1917,11 +1908,7 @@ func (r *mutationResolver) GenerateFrameworkStateOfApplicability(ctx context.Con // ExportFramework is the resolver for the exportFramework field. func (r *mutationResolver) ExportFramework(ctx context.Context, input types.ExportFrameworkInput) (*types.ExportFrameworkPayload, error) { prb := r.ProboService(ctx, input.FrameworkID.TenantID()) - user := UserFromContext(ctx) - if user == nil { - panic(fmt.Errorf("user not found")) - } err, exportJobID := prb.Frameworks.RequestExport( ctx, @@ -2773,11 +2760,7 @@ func (r *mutationResolver) BulkExportDocuments(ctx context.Context, input types. } prb := r.ProboService(ctx, input.DocumentIds[0].TenantID()) - user := UserFromContext(ctx) - if user == nil { - panic(fmt.Errorf("user not found")) - } options := probo.BulkExportOptions{ WithWatermark: input.WithWatermark, @@ -3557,15 +3540,11 @@ func (r *mutationResolver) DeleteCustomDomain(ctx context.Context, input types.D // InitiateDomainVerification is the resolver for the initiateDomainVerification field. func (r *mutationResolver) InitiateDomainVerification(ctx context.Context, input types.InitiateDomainVerificationInput) (*types.InitiateDomainVerificationPayload, error) { - user := UserFromContext(ctx) - if user == nil { - return nil, fmt.Errorf("user not authenticated") - } - organizationID := input.OrganizationID tenantID := organizationID.TenantID() - config, err := r.authSvc.InitiateDomainVerification(ctx, tenantID, organizationID, input.EmailDomain) + authSvc := r.AuthService(ctx, tenantID) + config, err := authSvc.InitiateDomainVerification(ctx, organizationID, input.EmailDomain) if err != nil { return nil, fmt.Errorf("cannot initiate domain verification: %w", err) } @@ -3584,15 +3563,11 @@ func (r *mutationResolver) InitiateDomainVerification(ctx context.Context, input // VerifyDomain is the resolver for the verifyDomain field. func (r *mutationResolver) VerifyDomain(ctx context.Context, input types.VerifyDomainInput) (*types.VerifyDomainPayload, error) { - user := UserFromContext(ctx) - if user == nil { - return nil, fmt.Errorf("user not authenticated") - } - configID := input.ID tenantID := configID.TenantID() - config, verified, err := r.authSvc.VerifyDomain(ctx, tenantID, configID) + authSvc := r.AuthService(ctx, tenantID) + config, verified, err := authSvc.VerifyDomain(ctx, configID) if err != nil { return nil, fmt.Errorf("cannot verify domain: %w", err) } @@ -3609,11 +3584,6 @@ func (r *mutationResolver) VerifyDomain(ctx context.Context, input types.VerifyD // CreateSAMLConfiguration is the resolver for the createSAMLConfiguration field. func (r *mutationResolver) CreateSAMLConfiguration(ctx context.Context, input types.CreateSAMLConfigurationInput) (*types.CreateSAMLConfigurationPayload, error) { - user := UserFromContext(ctx) - if user == nil { - return nil, fmt.Errorf("user not authenticated") - } - organizationID := input.OrganizationID tenantID := organizationID.TenantID() @@ -3670,7 +3640,8 @@ func (r *mutationResolver) CreateSAMLConfiguration(ctx context.Context, input ty autoSignupEnabled = *input.AutoSignupEnabled } - config, err := r.authSvc.WithTenant(tenantID).CreateSAMLConfiguration(ctx, auth.CreateSAMLConfigurationRequest{ + authSvc := r.AuthService(ctx, tenantID) + config, err := authSvc.CreateSAMLConfiguration(ctx, auth.CreateSAMLConfigurationRequest{ OrganizationID: organizationID, EmailDomain: input.EmailDomain, EnforcementPolicy: input.EnforcementPolicy, @@ -3699,15 +3670,11 @@ func (r *mutationResolver) CreateSAMLConfiguration(ctx context.Context, input ty // UpdateSAMLConfiguration is the resolver for the updateSAMLConfiguration field. func (r *mutationResolver) UpdateSAMLConfiguration(ctx context.Context, input types.UpdateSAMLConfigurationInput) (*types.UpdateSAMLConfigurationPayload, error) { - user := UserFromContext(ctx) - if user == nil { - return nil, fmt.Errorf("user not authenticated") - } - configID := input.ID tenantID := configID.TenantID() - updatedConfig, err := r.authSvc.WithTenant(tenantID).UpdateSAMLConfiguration(ctx, auth.UpdateSAMLConfigurationRequest{ + authSvc := r.AuthService(ctx, tenantID) + updatedConfig, err := authSvc.UpdateSAMLConfiguration(ctx, auth.UpdateSAMLConfigurationRequest{ ID: configID, Enabled: input.Enabled, EnforcementPolicy: input.EnforcementPolicy, @@ -3736,15 +3703,11 @@ func (r *mutationResolver) UpdateSAMLConfiguration(ctx context.Context, input ty // DeleteSAMLConfiguration is the resolver for the deleteSAMLConfiguration field. func (r *mutationResolver) DeleteSAMLConfiguration(ctx context.Context, input types.DeleteSAMLConfigurationInput) (*types.DeleteSAMLConfigurationPayload, error) { - user := UserFromContext(ctx) - if user == nil { - return nil, fmt.Errorf("user not authenticated") - } - configID := input.ID tenantID := configID.TenantID() - err := r.authSvc.WithTenant(tenantID).DeleteSAMLConfiguration(ctx, configID) + authSvc := r.AuthService(ctx, tenantID) + err := authSvc.DeleteSAMLConfiguration(ctx, configID) if err != nil { return nil, fmt.Errorf("cannot delete SAML configuration: %w", err) } @@ -3756,15 +3719,11 @@ func (r *mutationResolver) DeleteSAMLConfiguration(ctx context.Context, input ty // EnableSaml is the resolver for the enableSAML field. func (r *mutationResolver) EnableSaml(ctx context.Context, input types.EnableSAMLInput) (*types.EnableSAMLPayload, error) { - user := UserFromContext(ctx) - if user == nil { - return nil, fmt.Errorf("user not authenticated") - } - configID := input.ID tenantID := configID.TenantID() - enabledConfig, err := r.authSvc.WithTenant(tenantID).EnableSAMLConfiguration(ctx, configID) + authSvc := r.AuthService(ctx, tenantID) + enabledConfig, err := authSvc.EnableSAMLConfiguration(ctx, configID) if err != nil { return nil, fmt.Errorf("cannot enable SAML: %w", err) } @@ -3780,15 +3739,11 @@ func (r *mutationResolver) EnableSaml(ctx context.Context, input types.EnableSAM // DisableSaml is the resolver for the disableSAML field. func (r *mutationResolver) DisableSaml(ctx context.Context, input types.DisableSAMLInput) (*types.DisableSAMLPayload, error) { - user := UserFromContext(ctx) - if user == nil { - return nil, fmt.Errorf("user not authenticated") - } - configID := input.ID tenantID := configID.TenantID() - disabledConfig, err := r.authSvc.WithTenant(tenantID).DisableSAMLConfiguration(ctx, configID) + authSvc := r.AuthService(ctx, tenantID) + disabledConfig, err := authSvc.DisableSAMLConfiguration(ctx, configID) if err != nil { return nil, fmt.Errorf("cannot disable SAML: %w", err) } @@ -4551,7 +4506,8 @@ func (r *organizationResolver) CustomDomain(ctx context.Context, obj *types.Orga func (r *organizationResolver) SamlConfigurations(ctx context.Context, obj *types.Organization) ([]*types.SAMLConfiguration, error) { tenantID := obj.ID.TenantID() - configs, err := r.authSvc.WithTenant(tenantID).GetSAMLConfigurationsByOrganizationID(ctx, obj.ID) + authSvc := r.AuthService(ctx, tenantID) + configs, err := authSvc.GetSAMLConfigurationsByOrganizationID(ctx, obj.ID) if err != nil { return nil, fmt.Errorf("cannot load SAML configurations: %w", err) } @@ -5018,7 +4974,8 @@ func (r *sAMLConfigurationResolver) Organization(ctx context.Context, obj *types tenantID := obj.ID.TenantID() prb := r.ProboService(ctx, tenantID) - config, err := r.authSvc.WithTenant(tenantID).GetSAMLConfigurationByID(ctx, obj.ID) + authSvc := r.AuthService(ctx, tenantID) + config, err := authSvc.GetSAMLConfigurationByID(ctx, obj.ID) if err != nil { return nil, fmt.Errorf("cannot load SAML configuration: %w", err) } @@ -5848,35 +5805,6 @@ func (r *viewerResolver) Organizations(ctx context.Context, obj *types.Viewer, f return types.NewOrganizationConnection(page), nil } -// Invitations is the resolver for the invitations field. -func (r *viewerResolver) Invitations(ctx context.Context, obj *types.Viewer, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.InvitationOrder, filter *types.InvitationFilter) (*types.InvitationConnection, error) { - user := UserFromContext(ctx) - - 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) - - invitationFilter := coredata.NewInvitationFilter(nil) - if filter != nil { - invitationFilter = coredata.NewInvitationFilter(filter.Statuses) - } - - invitations, err := r.authzSvc.GetUserInvitations(ctx, user.EmailAddress, cursor, invitationFilter) - if err != nil { - panic(fmt.Errorf("cannot list invitations for user: %w", err)) - } - - return types.NewInvitationConnection(invitations, r, gid.GID{}, filter), nil -} - // Asset returns schema.AssetResolver implementation. func (r *Resolver) Asset() schema.AssetResolver { return &assetResolver{r} } diff --git a/pkg/server/auth/accept_invitation_handler.go b/pkg/server/auth/accept_invitation_handler.go index 047ba09c5..b863b7a39 100644 --- a/pkg/server/auth/accept_invitation_handler.go +++ b/pkg/server/auth/accept_invitation_handler.go @@ -36,13 +36,13 @@ type ( } ) -func AcceptInvitationHandler(authSvc *authsvc.Service, authzSvc *authz.Service, authCfg RoutesConfig) http.HandlerFunc { +func AcceptInvitationHandler(authSvc *authsvc.Service, authzSvc *authz.Service, cookieName string, cookieSecret string) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { ctx := r.Context() sessionAuthCfg := session.AuthConfig{ - CookieName: authCfg.CookieName, - CookieSecret: authCfg.CookieSecret, + CookieName: cookieName, + CookieSecret: cookieSecret, } errorHandler := session.ErrorHandler{ diff --git a/pkg/server/auth/auth.go b/pkg/server/auth/auth.go index c4a5c0ca0..c1d8178c1 100644 --- a/pkg/server/auth/auth.go +++ b/pkg/server/auth/auth.go @@ -23,20 +23,18 @@ import ( "github.com/getprobo/probo/pkg/filemanager" "github.com/go-chi/chi/v5" "go.gearno.de/kit/log" - "go.gearno.de/kit/pg" ) type Config struct { - Auth *authsvc.Service - Authz *authz.Service - SAML *authsvc.SAMLService - CookieName string - CookieDomain string - SessionDuration time.Duration - CookieSecret string - FileManager *filemanager.Service - PGClient *pg.Client - Logger *log.Logger + Auth *authsvc.Service + Authz *authz.Service + SAML *authsvc.SAMLService + CookieName string + CookieDomain string + SessionDuration time.Duration + CookieSecret string + FileManager *filemanager.Service + Logger *log.Logger } type Server struct { @@ -46,21 +44,21 @@ type Server struct { func NewServer(cfg Config) (*Server, error) { router := chi.NewRouter() - MountRoutes( - router, - cfg.Auth, - cfg.Authz, - cfg.SAML, - RoutesConfig{ - CookieName: cfg.CookieName, - CookieDomain: cfg.CookieDomain, - SessionDuration: cfg.SessionDuration, - CookieSecret: cfg.CookieSecret, - FileManager: cfg.FileManager, - PGClient: cfg.PGClient, - }, - cfg.Logger, - ) + router.Post("/register", SignUpHandler(cfg.Auth, cfg.CookieName, cfg.CookieSecret)) + router.Post("/login", SignInHandler(cfg.Auth, cfg.CookieName, cfg.CookieSecret)) + router.Delete("/logout", SignOutHandler(cfg.Auth, cfg.CookieName, cfg.CookieSecret)) + router.Post("/signup-from-invitation", SignupFromInvitationHandler(cfg.Auth, cfg.CookieName, cfg.CookieSecret)) + router.Post("/forget-password", ForgetPasswordHandler(cfg.Auth)) + router.Post("/reset-password", ResetPasswordHandler(cfg.Auth)) + router.Post("/check-sso", SAMLCheckSSOHandler(cfg.Auth, cfg.Logger)) + router.Get("/organizations", RequireAuth(cfg.Auth, cfg.Authz, cfg.CookieName, cfg.CookieSecret, ListOrganizationsHandler(cfg.Auth, cfg.Authz))) + router.Get("/organizations/{organizationID}/logo", RequireAuth(cfg.Auth, cfg.Authz, cfg.CookieName, cfg.CookieSecret, OrganizationLogoHandler(cfg.Auth, cfg.FileManager))) + router.Get("/invitations", RequireAuth(cfg.Auth, cfg.Authz, cfg.CookieName, cfg.CookieSecret, ListInvitationsHandler(cfg.Authz))) + router.Post("/invitations/accept", AcceptInvitationHandler(cfg.Auth, cfg.Authz, cfg.CookieName, cfg.CookieSecret)) + + router.Get("/saml/login/{samlConfigID}", SAMLLoginHandler(cfg.SAML, cfg.Auth, cfg.Logger)) + router.Post("/saml/consume", SAMLACSHandler(cfg.SAML, cfg.Auth, cfg.Authz, cfg.CookieName, cfg.CookieSecret, cfg.SessionDuration, cfg.Logger)) + router.Get("/saml/metadata", SAMLMetadataHandler(cfg.SAML)) return &Server{ router: router, diff --git a/pkg/server/auth/auth_middleware.go b/pkg/server/auth/auth_middleware.go new file mode 100644 index 000000000..741f310d5 --- /dev/null +++ b/pkg/server/auth/auth_middleware.go @@ -0,0 +1,93 @@ +// Copyright (c) 2025 Probo Inc . +// +// 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 auth + +import ( + "context" + "fmt" + "net/http" + + authsvc "github.com/getprobo/probo/pkg/auth" + "github.com/getprobo/probo/pkg/authz" + "github.com/getprobo/probo/pkg/coredata" + "github.com/getprobo/probo/pkg/server/session" + "go.gearno.de/kit/httpserver" +) + +type ctxKey struct{ name string } + +var ( + sessionContextKey = &ctxKey{name: "session"} + userContextKey = &ctxKey{name: "user"} +) + +func RequireAuth( + authSvc *authsvc.Service, + authzSvc *authz.Service, + cookieName string, + cookieSecret string, + next http.HandlerFunc, +) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + ctx := r.Context() + + sessionAuthCfg := session.AuthConfig{ + CookieName: cookieName, + CookieSecret: cookieSecret, + } + + errorHandler := session.ErrorHandler{ + OnCookieError: func(err error) { + httpserver.RenderError(w, http.StatusUnauthorized, fmt.Errorf("invalid session")) + }, + OnParseError: func(w http.ResponseWriter, authCfg session.AuthConfig) { + session.ClearCookie(w, authCfg) + httpserver.RenderError(w, http.StatusUnauthorized, fmt.Errorf("invalid session")) + }, + OnSessionError: func(w http.ResponseWriter, authCfg session.AuthConfig) { + session.ClearCookie(w, authCfg) + httpserver.RenderError(w, http.StatusUnauthorized, fmt.Errorf("session expired")) + }, + OnUserError: func(w http.ResponseWriter, authCfg session.AuthConfig) { + session.ClearCookie(w, authCfg) + httpserver.RenderError(w, http.StatusUnauthorized, fmt.Errorf("user not found")) + }, + OnTenantError: func(err error) { + panic(fmt.Errorf("cannot list tenants for user: %w", err)) + }, + } + + authResult := session.TryAuth(ctx, w, r, authSvc, authzSvc, sessionAuthCfg, errorHandler) + if authResult == nil { + httpserver.RenderError(w, http.StatusUnauthorized, fmt.Errorf("authentication required")) + return + } + + ctx = context.WithValue(ctx, sessionContextKey, authResult.Session) + ctx = context.WithValue(ctx, userContextKey, authResult.User) + + next(w, r.WithContext(ctx)) + } +} + +func SessionFromContext(ctx context.Context) *coredata.Session { + session, _ := ctx.Value(sessionContextKey).(*coredata.Session) + return session +} + +func UserFromContext(ctx context.Context) *coredata.User { + user, _ := ctx.Value(userContextKey).(*coredata.User) + return user +} diff --git a/pkg/server/auth/forget_password_handler.go b/pkg/server/auth/forget_password_handler.go index 70ca8cef9..a7b3fd636 100644 --- a/pkg/server/auth/forget_password_handler.go +++ b/pkg/server/auth/forget_password_handler.go @@ -33,7 +33,7 @@ type ( } ) -func ForgetPasswordHandler(authSvc *authsvc.Service, authCfg RoutesConfig) http.HandlerFunc { +func ForgetPasswordHandler(authSvc *authsvc.Service) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { var req ForgetPasswordRequest if err := json.NewDecoder(r.Body).Decode(&req); err != nil { diff --git a/pkg/server/auth/list_invitations_handler.go b/pkg/server/auth/list_invitations_handler.go index 80e7c2df7..4c0f8e5a4 100644 --- a/pkg/server/auth/list_invitations_handler.go +++ b/pkg/server/auth/list_invitations_handler.go @@ -15,18 +15,12 @@ package auth import ( - "context" "fmt" "net/http" - authsvc "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/page" - "github.com/getprobo/probo/pkg/server/session" "go.gearno.de/kit/httpserver" - "go.gearno.de/kit/pg" ) type ( @@ -35,167 +29,56 @@ type ( } InvitationResponse struct { - ID gid.GID `json:"id"` - Email string `json:"email"` - FullName string `json:"fullName"` - Role string `json:"role"` - ExpiresAt string `json:"expiresAt"` - AcceptedAt *string `json:"acceptedAt,omitempty"` - CreatedAt string `json:"createdAt"` - Organization OrganizationSummary `json:"organization"` + ID gid.GID `json:"id"` + Email string `json:"email"` + FullName string `json:"fullName"` + Role string `json:"role"` + ExpiresAt string `json:"expiresAt"` + AcceptedAt *string `json:"acceptedAt,omitempty"` + CreatedAt string `json:"createdAt"` + Organization OrganizationResponseSummary `json:"organization"` } - OrganizationSummary struct { + OrganizationResponseSummary struct { ID gid.GID `json:"id"` Name string `json:"name"` } ) -// loadOrganizationByID loads an organization by ID without tenant scope -func loadOrganizationByID( - ctx context.Context, - conn pg.Conn, - orgID gid.GID, -) (*coredata.Organization, error) { - query := ` -SELECT - id, - tenant_id, - name, - logo_file_id, - horizontal_logo_file_id, - description, - website_url, - email, - headquarter_address, - custom_domain_id, - created_at, - updated_at -FROM - authz_organizations -WHERE - id = $1 -` - - row := conn.QueryRow(ctx, query, orgID) - - var org coredata.Organization - err := row.Scan( - &org.ID, - &org.TenantID, - &org.Name, - &org.LogoFileID, - &org.HorizontalLogoFileID, - &org.Description, - &org.WebsiteURL, - &org.Email, - &org.HeadquarterAddress, - &org.CustomDomainID, - &org.CreatedAt, - &org.UpdatedAt, - ) - if err != nil { - return nil, fmt.Errorf("cannot load organization: %w", err) - } - - return &org, nil -} - -func ListInvitationsHandler(authSvc *authsvc.Service, authzSvc *authz.Service, authCfg RoutesConfig) http.HandlerFunc { +func ListInvitationsHandler(authzSvc *authz.Service) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { ctx := r.Context() + user := UserFromContext(ctx) - sessionAuthCfg := session.AuthConfig{ - CookieName: authCfg.CookieName, - CookieSecret: authCfg.CookieSecret, - } - - errorHandler := session.ErrorHandler{ - OnCookieError: func(err error) { - httpserver.RenderError(w, http.StatusUnauthorized, fmt.Errorf("invalid session")) - }, - OnParseError: func(w http.ResponseWriter, authCfg session.AuthConfig) { - session.ClearCookie(w, authCfg) - httpserver.RenderError(w, http.StatusUnauthorized, fmt.Errorf("invalid session")) - }, - OnSessionError: func(w http.ResponseWriter, authCfg session.AuthConfig) { - session.ClearCookie(w, authCfg) - httpserver.RenderError(w, http.StatusUnauthorized, fmt.Errorf("session expired")) - }, - OnUserError: func(w http.ResponseWriter, authCfg session.AuthConfig) { - session.ClearCookie(w, authCfg) - httpserver.RenderError(w, http.StatusUnauthorized, fmt.Errorf("user not found")) - }, - OnTenantError: func(err error) { - panic(fmt.Errorf("cannot list tenants for user: %w", err)) - }, - } - - authResult := session.TryAuth(ctx, w, r, authSvc, authzSvc, sessionAuthCfg, errorHandler) - if authResult == nil { - httpserver.RenderError(w, http.StatusUnauthorized, fmt.Errorf("authentication required")) - return - } - - // Get pending invitations for the user - cursor := page.NewCursor( - 1000, - nil, - page.Head, - page.OrderBy[coredata.InvitationOrderField]{ - Field: coredata.InvitationOrderFieldCreatedAt, - Direction: page.OrderDirectionDesc, - }, - ) - - invitationFilter := coredata.NewInvitationFilter([]coredata.InvitationStatus{coredata.InvitationStatusPending}) - - invitationsPage, err := authzSvc.GetUserInvitations(ctx, authResult.User.EmailAddress, cursor, invitationFilter) + invitations, err := authzSvc.GetUserPendingInvitations(ctx, user.EmailAddress) if err != nil { panic(fmt.Errorf("cannot list invitations for user: %w", err)) } - // Build response response := ListInvitationsResponse{ - Invitations: make([]InvitationResponse, 0, len(invitationsPage.Data)), + Invitations: make([]InvitationResponse, 0, len(invitations)), } - // Load organization data for each invitation - err = authCfg.PGClient.WithConn(ctx, func(conn pg.Conn) error { - for _, invitation := range invitationsPage.Data { - invitationResp := InvitationResponse{ - ID: invitation.ID, - Email: invitation.Email, - FullName: invitation.FullName, - Role: invitation.Role, - ExpiresAt: invitation.ExpiresAt.Format("2006-01-02T15:04:05Z07:00"), - CreatedAt: invitation.CreatedAt.Format("2006-01-02T15:04:05Z07:00"), - } - - if invitation.AcceptedAt != nil { - acceptedAtStr := invitation.AcceptedAt.Format("2006-01-02T15:04:05Z07:00") - invitationResp.AcceptedAt = &acceptedAtStr - } - - // Load organization details - org, err := loadOrganizationByID(ctx, conn, invitation.OrganizationID) - if err != nil { - // Log error but continue - organization might have been deleted - return nil - } - - invitationResp.Organization = OrganizationSummary{ - ID: org.ID, - Name: org.Name, - } - - response.Invitations = append(response.Invitations, invitationResp) + for _, invitation := range invitations { + invitationResp := InvitationResponse{ + ID: invitation.ID, + Email: invitation.Email, + FullName: invitation.FullName, + Role: invitation.Role, + ExpiresAt: invitation.ExpiresAt.Format("2006-01-02T15:04:05Z07:00"), + CreatedAt: invitation.CreatedAt.Format("2006-01-02T15:04:05Z07:00"), + Organization: OrganizationResponseSummary{ + ID: invitation.Organization.ID, + Name: invitation.Organization.Name, + }, } - return nil - }) - if err != nil { - panic(fmt.Errorf("cannot load organization details: %w", err)) + if invitation.AcceptedAt != nil { + acceptedAtStr := invitation.AcceptedAt.Format("2006-01-02T15:04:05Z07:00") + invitationResp.AcceptedAt = &acceptedAtStr + } + + response.Invitations = append(response.Invitations, invitationResp) } httpserver.RenderJSON(w, http.StatusOK, response) diff --git a/pkg/server/auth/list_organizations_handler.go b/pkg/server/auth/list_organizations_handler.go index 2515377b8..e5caff35a 100644 --- a/pkg/server/auth/list_organizations_handler.go +++ b/pkg/server/auth/list_organizations_handler.go @@ -15,20 +15,14 @@ package auth import ( - "context" - "errors" "fmt" "net/http" - "time" authsvc "github.com/getprobo/probo/pkg/auth" "github.com/getprobo/probo/pkg/authz" "github.com/getprobo/probo/pkg/coredata" - "github.com/getprobo/probo/pkg/filemanager" "github.com/getprobo/probo/pkg/gid" - "github.com/getprobo/probo/pkg/server/session" "go.gearno.de/kit/httpserver" - "go.gearno.de/kit/pg" ) type ( @@ -54,145 +48,84 @@ const ( AuthStatusExpired AuthenticationStatus = "expired" ) -// generateLogoURL generates a presigned URL for an organization's logo -func generateLogoURL( - ctx context.Context, - fileManager *filemanager.Service, - conn pg.Conn, - logoFileID *gid.GID, -) (*string, error) { - if logoFileID == nil { - return nil, nil +func buildOrganizationResponse( + org *coredata.Organization, + accessResult authsvc.AccessResult, + sessionData coredata.SessionData, +) OrganizationResponse { + // Generate logo URL path if organization has a logo + var logoURL *string + if org.LogoFileID != nil { + url := fmt.Sprintf("/auth/organizations/%s/logo", org.ID) + logoURL = &url } - var file coredata.File - // Load file without scope since we're in auth context (cross-tenant) - q := `SELECT bucket_name, file_key, file_name, mime_type, file_size FROM files WHERE id = $1` - err := conn.QueryRow(ctx, q, logoFileID).Scan( - &file.BucketName, - &file.FileKey, - &file.FileName, - &file.MimeType, - &file.FileSize, - ) - if err != nil { - return nil, fmt.Errorf("cannot load file: %w", err) + orgResponse := OrganizationResponse{ + ID: org.ID, + Name: org.Name, + LogoURL: logoURL, } - presignedURL, err := fileManager.GenerateFileUrl(ctx, &file, 1*time.Hour) - if err != nil { - return nil, fmt.Errorf("cannot generate file URL: %w", err) + // User does not have required authentication + if !accessResult.Allowed { + orgResponse.AuthStatus = AuthStatusUnauthenticated + + switch accessResult.MissingAuth { + case authsvc.AuthMethodSAML, authsvc.AuthMethodAny: + orgResponse.AuthenticationMethod = "saml" + if accessResult.SAMLConfig != nil { + orgResponse.LoginURL = fmt.Sprintf("/auth/saml/login/%s", accessResult.SAMLConfig.ID) + } + case authsvc.AuthMethodPassword: + orgResponse.AuthenticationMethod = "password" + orgResponse.LoginURL = "/authentication/login?method=password" + } + return orgResponse } - return &presignedURL, nil + // User has required authentication + orgResponse.AuthStatus = AuthStatusAuthenticated + + if sessionData.PasswordAuthenticated { + orgResponse.AuthenticationMethod = "password" + orgResponse.LoginURL = "/authentication/login?method=password" + } else if samlInfo, ok := sessionData.SAMLAuthenticatedOrgs[org.ID.String()]; ok { + orgResponse.AuthenticationMethod = "saml" + orgResponse.LoginURL = fmt.Sprintf("/auth/saml/login/%s", samlInfo.SAMLConfigID) + } else { + orgResponse.AuthenticationMethod = "any" + orgResponse.LoginURL = "/authentication/login?method=password" + } + + return orgResponse } -func ListOrganizationsHandler(authSvc *authsvc.Service, authzSvc *authz.Service, authCfg RoutesConfig) http.HandlerFunc { +func ListOrganizationsHandler(authSvc *authsvc.Service, authzSvc *authz.Service) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { ctx := r.Context() + user := UserFromContext(ctx) + sess := SessionFromContext(ctx) - sessionAuthCfg := session.AuthConfig{ - CookieName: authCfg.CookieName, - CookieSecret: authCfg.CookieSecret, - } - - errorHandler := session.ErrorHandler{ - OnCookieError: func(err error) { - httpserver.RenderError(w, http.StatusUnauthorized, fmt.Errorf("invalid session")) - }, - OnParseError: func(w http.ResponseWriter, authCfg session.AuthConfig) { - session.ClearCookie(w, authCfg) - httpserver.RenderError(w, http.StatusUnauthorized, fmt.Errorf("invalid session")) - }, - OnSessionError: func(w http.ResponseWriter, authCfg session.AuthConfig) { - session.ClearCookie(w, authCfg) - httpserver.RenderError(w, http.StatusUnauthorized, fmt.Errorf("session expired")) - }, - OnUserError: func(w http.ResponseWriter, authCfg session.AuthConfig) { - session.ClearCookie(w, authCfg) - httpserver.RenderError(w, http.StatusUnauthorized, fmt.Errorf("user not found")) - }, - OnTenantError: func(err error) { - panic(fmt.Errorf("cannot list tenants for user: %w", err)) - }, - } - - authResult := session.TryAuth(ctx, w, r, authSvc, authzSvc, sessionAuthCfg, errorHandler) - if authResult == nil { - httpserver.RenderError(w, http.StatusUnauthorized, fmt.Errorf("authentication required")) - return - } - - // Get all organizations for the user (without filtering by authentication state) - organizations, err := authzSvc.GetAllUserOrganizations(ctx, authResult.User.ID) + organizations, err := authzSvc.GetAllUserOrganizations(ctx, user.ID) if err != nil { panic(fmt.Errorf("cannot list organizations for user: %w", err)) } - // Build response with authentication requirements for each organization + orgIDs := make([]gid.GID, len(organizations)) + for i, org := range organizations { + orgIDs[i] = org.ID + } + accessResults, err := authSvc.CheckOrganizationAccess(ctx, user, orgIDs, sess) + if err != nil { + panic(fmt.Errorf("cannot check organization access: %w", err)) + } + response := ListOrganizationsResponse{ Organizations: make([]OrganizationResponse, 0, len(organizations)), } - for _, org := range organizations { - orgResponse := OrganizationResponse{ - ID: org.ID, - Name: org.Name, - } - - // Generate logo URL if available - if authCfg.FileManager != nil && authCfg.PGClient != nil { - err := authCfg.PGClient.WithConn(ctx, func(conn pg.Conn) error { - logoURL, err := generateLogoURL(ctx, authCfg.FileManager, conn, org.LogoFileID) - if err != nil { - // Log error but don't fail the request - return nil - } - orgResponse.LogoURL = logoURL - return nil - }) - if err != nil { - // Log error but continue - } - } - - // Check authentication requirements for this organization - err := authSvc.CheckOrganizationAccess(ctx, authResult.User, org.ID, authResult.Session) - if err != nil { - // User needs additional authentication - var errSAMLRequired authsvc.ErrSAMLAuthRequired - if errors.As(err, &errSAMLRequired) { - orgResponse.AuthenticationMethod = "saml" - orgResponse.AuthStatus = AuthStatusUnauthenticated - orgResponse.LoginURL = fmt.Sprintf("/auth/saml/login/%s", errSAMLRequired.ConfigID) - } else { - orgResponse.AuthenticationMethod = "password" - orgResponse.AuthStatus = AuthStatusUnauthenticated - orgResponse.LoginURL = "/authentication/login?method=password" - } - } else { - // User has proper authentication - orgResponse.AuthStatus = AuthStatusAuthenticated - - // Determine which auth method they used - if authResult.Session.Data.PasswordAuthenticated { - orgResponse.AuthenticationMethod = "password" - orgResponse.LoginURL = "/authentication/login?method=password" - } else if len(authResult.Session.Data.SAMLAuthenticatedOrgs) > 0 { - // Find SAML config for this org - orgResponse.AuthenticationMethod = "saml" - // Try to find the SAML config ID for login URL - if samlInfo, ok := authResult.Session.Data.SAMLAuthenticatedOrgs[org.ID.String()]; ok { - orgResponse.LoginURL = fmt.Sprintf("/auth/saml/login/%s", samlInfo.SAMLConfigID) - } else { - orgResponse.LoginURL = "/authentication/login?method=password" - } - } else { - orgResponse.AuthenticationMethod = "any" - orgResponse.LoginURL = "/authentication/login?method=password" - } - } - + accessResult := accessResults[org.ID] + orgResponse := buildOrganizationResponse(org, accessResult, sess.Data) response.Organizations = append(response.Organizations, orgResponse) } diff --git a/pkg/server/auth/organization_logo_handler.go b/pkg/server/auth/organization_logo_handler.go new file mode 100644 index 000000000..fa34eeaae --- /dev/null +++ b/pkg/server/auth/organization_logo_handler.go @@ -0,0 +1,58 @@ +// Copyright (c) 2025 Probo Inc . +// +// 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 auth + +import ( + "context" + "fmt" + "net/http" + "time" + + authsvc "github.com/getprobo/probo/pkg/auth" + "github.com/getprobo/probo/pkg/coredata" + "github.com/getprobo/probo/pkg/gid" + "github.com/go-chi/chi/v5" +) + +func OrganizationLogoHandler(authSvc *authsvc.Service, fileManager interface { + GenerateFileUrl(ctx context.Context, file *coredata.File, duration time.Duration) (string, error) +}) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + ctx := r.Context() + user := UserFromContext(ctx) + session := SessionFromContext(ctx) + + organizationIDStr := chi.URLParam(r, "organizationID") + organizationID, err := gid.ParseGID(organizationIDStr) + if err != nil { + http.Error(w, "Invalid organization ID", http.StatusBadRequest) + return + } + + logoFile, err := authSvc.GetOrganizationLogoFile(ctx, user, organizationID, session) + if err != nil { + panic(fmt.Errorf("cannot get organization logo: %w", err)) + } + + presignedURL, err := fileManager.GenerateFileUrl(ctx, logoFile, 1*time.Hour) + if err != nil { + panic(fmt.Errorf("cannot generate presigned URL: %w", err)) + } + + w.Header().Set("Cache-Control", "public, max-age=3600") + + http.Redirect(w, r, presignedURL, http.StatusFound) + } +} diff --git a/pkg/server/auth/reset_password_handler.go b/pkg/server/auth/reset_password_handler.go index 32a03f14e..a2d71804f 100644 --- a/pkg/server/auth/reset_password_handler.go +++ b/pkg/server/auth/reset_password_handler.go @@ -36,7 +36,7 @@ type ( } ) -func ResetPasswordHandler(authSvc *authsvc.Service, authCfg RoutesConfig) http.HandlerFunc { +func ResetPasswordHandler(authSvc *authsvc.Service) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { var req ResetPasswordRequest if err := json.NewDecoder(r.Body).Decode(&req); err != nil { diff --git a/pkg/server/auth/router.go b/pkg/server/auth/router.go deleted file mode 100644 index f2981b047..000000000 --- a/pkg/server/auth/router.go +++ /dev/null @@ -1,60 +0,0 @@ -// Copyright (c) 2025 Probo Inc . -// -// 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 auth - -import ( - "time" - - authsvc "github.com/getprobo/probo/pkg/auth" - "github.com/getprobo/probo/pkg/authz" - "github.com/getprobo/probo/pkg/filemanager" - "github.com/go-chi/chi/v5" - "go.gearno.de/kit/log" - "go.gearno.de/kit/pg" -) - -type RoutesConfig struct { - CookieName string - CookieDomain string - SessionDuration time.Duration - CookieSecret string - FileManager *filemanager.Service - PGClient *pg.Client -} - -func MountRoutes( - r chi.Router, - authSvc *authsvc.Service, - authzSvc *authz.Service, - samlSvc *authsvc.SAMLService, - authCfg RoutesConfig, - logger *log.Logger, -) { - r.Post("/register", SignUpHandler(authSvc, authCfg)) - r.Post("/login", SignInHandler(authSvc, authCfg)) - r.Delete("/logout", SignOutHandler(authSvc, authCfg)) - r.Post("/signup-from-invitation", SignupFromInvitationHandler(authSvc, authCfg)) - r.Post("/forget-password", ForgetPasswordHandler(authSvc, authCfg)) - r.Post("/reset-password", ResetPasswordHandler(authSvc, authCfg)) - r.Post("/check-sso", SAMLCheckSSOHandler(authSvc, logger)) - r.Get("/organizations", ListOrganizationsHandler(authSvc, authzSvc, authCfg)) - r.Get("/invitations", ListInvitationsHandler(authSvc, authzSvc, authCfg)) - r.Post("/invitations/accept", AcceptInvitationHandler(authSvc, authzSvc, authCfg)) - - // SAML routes - r.Get("/saml/login/{samlConfigID}", SAMLLoginHandler(samlSvc, authSvc, logger)) - r.Post("/saml/consume", SAMLACSHandler(samlSvc, authSvc, authzSvc, authCfg, logger)) - r.Get("/saml/metadata", SAMLMetadataHandler(samlSvc)) -} diff --git a/pkg/server/auth/saml_acs_handler.go b/pkg/server/auth/saml_acs_handler.go index 3d6c2244b..d43e25704 100644 --- a/pkg/server/auth/saml_acs_handler.go +++ b/pkg/server/auth/saml_acs_handler.go @@ -15,6 +15,7 @@ package auth import ( + "errors" "fmt" "net/http" "time" @@ -27,11 +28,14 @@ import ( "go.gearno.de/kit/log" ) -func getSessionIDFromCookie(r *http.Request, authCfg RoutesConfig) (gid.GID, error) { - cookieValue, err := securecookie.Get(r, securecookie.DefaultConfig( - authCfg.CookieName, - authCfg.CookieSecret, - )) +func getSessionIDFromCookie(r *http.Request, cookieName string, cookieSecret string) (gid.GID, error) { + cookieValue, err := securecookie.Get( + r, + securecookie.DefaultConfig( + cookieName, + cookieSecret, + ), + ) if err != nil { return gid.GID{}, err } @@ -39,7 +43,7 @@ func getSessionIDFromCookie(r *http.Request, authCfg RoutesConfig) (gid.GID, err return gid.ParseGID(cookieValue) } -func SAMLACSHandler(samlSvc *authsvc.SAMLService, authSvc *authsvc.Service, authzSvc *authz.Service, authCfg RoutesConfig, logger *log.Logger) http.HandlerFunc { +func SAMLACSHandler(samlSvc *authsvc.SAMLService, authSvc *authsvc.Service, authzSvc *authz.Service, cookieName string, cookieSecret string, sessionDuration time.Duration, logger *log.Logger) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { ctx := r.Context() @@ -68,10 +72,32 @@ func SAMLACSHandler(samlSvc *authsvc.SAMLService, authSvc *authsvc.Service, auth return } - user, err := authSvc.CreateOrGetSAMLUser(ctx, userInfo.Email, userInfo.FullName, userInfo.SAMLSubject) + var existingSession *coredata.Session + if existingSessionID, err := getSessionIDFromCookie(r, cookieName, cookieSecret); err == nil { + if session, err := authSvc.GetSession(ctx, existingSessionID); err == nil { + existingSession = session + } + } + + session, user, err := authSvc.ProvisionSAMLUser( + ctx, + userInfo.SAMLConfigID, + userInfo.OrganizationID, + userInfo.Email, + userInfo.FullName, + userInfo.SAMLSubject, + existingSession, + sessionDuration, + ) if err != nil { - logger.ErrorCtx(ctx, "cannot create or get SAML user", log.Error(err)) - http.Error(w, "cannot create user", http.StatusInternalServerError) + var autoSignupDisabledErr *authsvc.ErrSAMLAutoSignupDisabled + if errors.As(err, &autoSignupDisabledErr) { + logger.WarnCtx(ctx, "SAML auto-signup is disabled") + http.Error(w, "User does not exist and auto-signup is disabled for this organization", http.StatusForbidden) + return + } + logger.ErrorCtx(ctx, "cannot provision SAML user", log.Error(err)) + http.Error(w, "cannot provision user", http.StatusInternalServerError) return } @@ -83,43 +109,11 @@ func SAMLACSHandler(samlSvc *authsvc.SAMLService, authSvc *authsvc.Service, auth return } - var session *coredata.Session - if existingSessionID, err := getSessionIDFromCookie(r, authCfg); err == nil { - if existingSession, err := authSvc.GetSession(ctx, existingSessionID); err == nil && existingSession.UserID == user.ID { - session = existingSession - } - } - - if session == nil { - session, err = authSvc.CreateSessionForUser(ctx, user.ID, authCfg.SessionDuration) - if err != nil { - logger.ErrorCtx(ctx, "cannot create session", log.Error(err), log.String("user_id", user.ID.String())) - http.Error(w, "cannot create session", http.StatusInternalServerError) - return - } - } - - if session.Data.SAMLAuthenticatedOrgs == nil { - session.Data.SAMLAuthenticatedOrgs = make(map[string]coredata.SAMLAuthInfo) - } - session.Data.SAMLAuthenticatedOrgs[userInfo.OrganizationID.String()] = coredata.SAMLAuthInfo{ - AuthenticatedAt: time.Now(), - SAMLConfigID: userInfo.SAMLConfigID, - SAMLSubject: userInfo.SAMLSubject, - } - - err = authSvc.UpdateSessionData(ctx, session.ID, session.Data) - if err != nil { - logger.ErrorCtx(ctx, "cannot update session data", log.Error(err), log.String("session_id", session.ID.String())) - http.Error(w, "cannot update session", http.StatusInternalServerError) - return - } - securecookie.Set( w, securecookie.DefaultConfig( - authCfg.CookieName, - authCfg.CookieSecret, + cookieName, + cookieSecret, ), session.ID.String(), ) diff --git a/pkg/server/auth/sign_in_handler.go b/pkg/server/auth/sign_in_handler.go index 586220eb9..f1fbb788e 100644 --- a/pkg/server/auth/sign_in_handler.go +++ b/pkg/server/auth/sign_in_handler.go @@ -47,7 +47,7 @@ type ( } ) -func SignInHandler(authSvc *authsvc.Service, authCfg RoutesConfig) http.HandlerFunc { +func SignInHandler(authSvc *authsvc.Service, cookieName string, cookieSecret string) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { var req SignInRequest @@ -57,13 +57,13 @@ func SignInHandler(authSvc *authsvc.Service, authCfg RoutesConfig) http.HandlerF } var existingSession *coredata.Session - if existingSessionID, err := getSessionIDFromCookie(r, authCfg); err == nil { + if existingSessionID, err := getSessionIDFromCookie(r, cookieName, cookieSecret); err == nil { if session, err := authSvc.GetSession(r.Context(), existingSessionID); err == nil { existingSession = session } } - session, user, err := authSvc.SignInWithExistingSession(r.Context(), req.Email, req.Password, existingSession) + session, user, err := authSvc.SignIn(r.Context(), req.Email, req.Password, existingSession) if err != nil { var ErrInvalidCredentials *authsvc.ErrInvalidCredentials if errors.As(err, &ErrInvalidCredentials) { @@ -77,8 +77,8 @@ func SignInHandler(authSvc *authsvc.Service, authCfg RoutesConfig) http.HandlerF securecookie.Set( w, securecookie.DefaultConfig( - authCfg.CookieName, - authCfg.CookieSecret, + cookieName, + cookieSecret, ), session.ID.String(), ) diff --git a/pkg/server/auth/sign_out_handler.go b/pkg/server/auth/sign_out_handler.go index 249f5f084..a8c5e71e1 100644 --- a/pkg/server/auth/sign_out_handler.go +++ b/pkg/server/auth/sign_out_handler.go @@ -24,12 +24,12 @@ import ( "go.gearno.de/kit/httpserver" ) -func SignOutHandler(authSvc *authsvc.Service, authCfg RoutesConfig) http.HandlerFunc { +func SignOutHandler(authSvc *authsvc.Service, cookieName string, cookieSecret string) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { sessionID, err := securecookie.Get(r, securecookie.DefaultConfig( - authCfg.CookieName, - authCfg.CookieSecret, + cookieName, + cookieSecret, )) if err != nil { httpserver.RenderError(w, http.StatusBadRequest, err) @@ -48,8 +48,8 @@ func SignOutHandler(authSvc *authsvc.Service, authCfg RoutesConfig) http.Handler } securecookie.Clear(w, securecookie.DefaultConfig( - authCfg.CookieName, - authCfg.CookieSecret, + cookieName, + cookieSecret, )) w.Header().Set("Clear-Site-Data", "*") diff --git a/pkg/server/auth/sign_up_handler.go b/pkg/server/auth/sign_up_handler.go index a83516caa..cbf8597e9 100644 --- a/pkg/server/auth/sign_up_handler.go +++ b/pkg/server/auth/sign_up_handler.go @@ -37,7 +37,7 @@ type ( } ) -func SignUpHandler(authSvc *authsvc.Service, authCfg RoutesConfig) http.HandlerFunc { +func SignUpHandler(authSvc *authsvc.Service, cookieName string, cookieSecret string) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { var req SignUpRequest if err := json.NewDecoder(r.Body).Decode(&req); err != nil { @@ -70,8 +70,8 @@ func SignUpHandler(authSvc *authsvc.Service, authCfg RoutesConfig) http.HandlerF securecookie.Set( w, securecookie.DefaultConfig( - authCfg.CookieName, - authCfg.CookieSecret, + cookieName, + cookieSecret, ), session.ID.String(), ) diff --git a/pkg/server/auth/signup_from_invitation_handler.go b/pkg/server/auth/signup_from_invitation_handler.go index 4adb7dab0..4bfbd7775 100644 --- a/pkg/server/auth/signup_from_invitation_handler.go +++ b/pkg/server/auth/signup_from_invitation_handler.go @@ -35,7 +35,7 @@ type ( } ) -func SignupFromInvitationHandler(authSvc *authsvc.Service, authCfg RoutesConfig) http.HandlerFunc { +func SignupFromInvitationHandler(authSvc *authsvc.Service, cookieName string, cookieSecret string) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { var req SignupFromInvitationRequest if err := json.NewDecoder(r.Body).Decode(&req); err != nil { @@ -52,8 +52,8 @@ func SignupFromInvitationHandler(authSvc *authsvc.Service, authCfg RoutesConfig) securecookie.Set( w, securecookie.DefaultConfig( - authCfg.CookieName, - authCfg.CookieSecret, + cookieName, + cookieSecret, ), session.ID.String(), ) diff --git a/pkg/server/server.go b/pkg/server/server.go index c6bc0bd80..4163746fb 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -109,7 +109,6 @@ func NewServer(cfg Config) (*Server, error) { SessionDuration: cfg.ConsoleAuth.SessionDuration, CookieSecret: cfg.ConsoleAuth.CookieSecret, FileManager: cfg.FileManager, - PGClient: cfg.PGClient, Logger: cfg.Logger.Named("auth"), }) if err != nil { diff --git a/pkg/server/session/session.go b/pkg/server/session/session.go index a3ba8be30..0de53fdac 100644 --- a/pkg/server/session/session.go +++ b/pkg/server/session/session.go @@ -103,15 +103,30 @@ func TryAuth( allowedTenantIDs := make([]gid.TenantID, 0, len(organizations)) authErrors := make(map[gid.TenantID]error) + // Extract organization IDs for batch check + orgIDs := make([]gid.GID, len(organizations)) + for i, org := range organizations { + orgIDs[i] = org.ID + } + + // Batch check access to all organizations in a single query + accessResults, err := authSvc.CheckOrganizationAccess(ctx, user, orgIDs, session) + if err != nil { + if errorHandler.OnTenantError != nil { + errorHandler.OnTenantError(err) + } + return nil + } + + // Process results for _, org := range organizations { - // Check if user has the required authentication for this organization - err := authSvc.CheckOrganizationAccess(ctx, user, org.ID, session) - if err == nil { + result := accessResults[org.ID] + if result.Allowed { // User has proper authentication for this org allowedTenantIDs = append(allowedTenantIDs, org.ID.TenantID()) } else { // Store the authentication error for later use - authErrors[org.ID.TenantID()] = err + authErrors[org.ID.TenantID()] = result.ToError(authSvc.BaseURL()) } }