Rewire IAM, mailman, and probod for the portal

Wire the certificate manager and trust center base domain into IAM so
organization creation provisions a managed default domain and certificate
atomically. Email presenters in IAM and mailman resolve public URLs
through the compliance portal resolver and read profile fields from the
trust center. probod initializes the certmanager service and injects the
new management and visitor services.

Signed-off-by: Bryan Frimin <bryan@probo.com>
This commit is contained in:
Bryan Frimin
2026-07-10 15:13:46 +02:00
parent 44ad34561e
commit a9a126899e
7 changed files with 234 additions and 208 deletions

View File

@@ -22,14 +22,13 @@ package iam
import ( import (
"context" "context"
"errors"
"fmt" "fmt"
"net/url"
"path/filepath" "path/filepath"
"time" "time"
"go.gearno.de/kit/pg" "go.gearno.de/kit/pg"
"go.probo.inc/probo/packages/emails" "go.probo.inc/probo/packages/emails"
"go.probo.inc/probo/pkg/complianceportal/resolver"
"go.probo.inc/probo/pkg/coredata" "go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/gid"
) )
@@ -96,7 +95,7 @@ func (s *CompliancePageService) EmailPresenterConfig(ctx context.Context, compli
var ( var (
compliancePage = &coredata.TrustCenter{} compliancePage = &coredata.TrustCenter{}
organization = &coredata.Organization{} organization = &coredata.Organization{}
customDomain *coredata.CustomDomain compliancePageURL string
logoFile = &coredata.File{} logoFile = &coredata.File{}
emailPresenterCfg = emails.DefaultPresenterConfig(s.baseURL) emailPresenterCfg = emails.DefaultPresenterConfig(s.baseURL)
) )
@@ -120,13 +119,19 @@ func (s *CompliancePageService) EmailPresenterConfig(ctx context.Context, compli
return fmt.Errorf("cannot load organization: %w", err) return fmt.Errorf("cannot load organization: %w", err)
} }
customDomain = &coredata.CustomDomain{} publicURL, err := resolver.PublicURLForTrustCenter(
if err := customDomain.LoadByOrganizationID(ctx, conn, scope, organization.ID); err != nil { ctx,
if !errors.Is(err, coredata.ErrResourceNotFound) { conn,
return fmt.Errorf("cannot load custom domain: %w", err) scope,
} compliancePage,
s.trustCenterBaseDomain,
)
if err != nil {
return fmt.Errorf("cannot resolve compliance page URL: %w", err)
} }
compliancePageURL = publicURL
return nil return nil
}, },
) )
@@ -134,24 +139,7 @@ func (s *CompliancePageService) EmailPresenterConfig(ctx context.Context, compli
return emailPresenterCfg, err return emailPresenterCfg, err
} }
parsedBaseURL, err := url.Parse(s.baseURL) emailPresenterCfg.BaseURL = compliancePageURL
if err != nil {
return emailPresenterCfg, fmt.Errorf("cannot parse base URL: %w", err)
}
baseURL := url.URL{
Scheme: parsedBaseURL.Scheme,
Host: parsedBaseURL.Host,
Path: "/trust/" + compliancePage.Slug,
}
if customDomain != nil && customDomain.SSLStatus == coredata.CustomDomainSSLStatusActive {
baseURL.Host = customDomain.Domain
baseURL.Scheme = "https"
baseURL.Path = ""
}
emailPresenterCfg.BaseURL = baseURL.String()
if compliancePage.LogoFileID != nil { if compliancePage.LogoFileID != nil {
if logoFile.FileKey == "" { if logoFile.FileKey == "" {
@@ -162,12 +150,12 @@ func (s *CompliancePageService) EmailPresenterConfig(ctx context.Context, compli
emailPresenterCfg.SenderCompanyLogoPath = filepath.Join("/api/files/v1/public/", logoFile.ID.String()) emailPresenterCfg.SenderCompanyLogoPath = filepath.Join("/api/files/v1/public/", logoFile.ID.String())
emailPresenterCfg.SenderCompanyName = organization.Name emailPresenterCfg.SenderCompanyName = organization.Name
if organization.WebsiteURL != nil { if compliancePage.WebsiteURL != nil {
emailPresenterCfg.SenderCompanyWebsiteURL = *organization.WebsiteURL emailPresenterCfg.SenderCompanyWebsiteURL = *compliancePage.WebsiteURL
} }
if organization.HeadquarterAddress != nil { if compliancePage.HeadquarterAddress != nil {
emailPresenterCfg.SenderCompanyHeadquarterAddress = *organization.HeadquarterAddress emailPresenterCfg.SenderCompanyHeadquarterAddress = *compliancePage.HeadquarterAddress
} }
} }

View File

@@ -33,6 +33,7 @@ import (
"go.gearno.de/kit/pg" "go.gearno.de/kit/pg"
"go.opentelemetry.io/otel/trace" "go.opentelemetry.io/otel/trace"
"go.probo.inc/probo/pkg/baseurl" "go.probo.inc/probo/pkg/baseurl"
"go.probo.inc/probo/pkg/certmanager"
"go.probo.inc/probo/pkg/connector" "go.probo.inc/probo/pkg/connector"
"go.probo.inc/probo/pkg/coredata" "go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/crypto/cipher" "go.probo.inc/probo/pkg/crypto/cipher"
@@ -61,6 +62,9 @@ type (
magicLinkTokenValidity time.Duration magicLinkTokenValidity time.Duration
sessionDuration time.Duration sessionDuration time.Duration
bucket string bucket string
encryptionKey cipher.EncryptionKey
trustCenterBaseDomain string
certManager *certmanager.Service
certificate *x509.Certificate certificate *x509.Certificate
privateKey *rsa.PrivateKey privateKey *rsa.PrivateKey
logger *log.Logger logger *log.Logger
@@ -90,6 +94,8 @@ type (
Bucket string Bucket string
TokenSecret string TokenSecret string
BaseURL *baseurl.BaseURL BaseURL *baseurl.BaseURL
TrustCenterBaseDomain string
CertManager *certmanager.Service
EncryptionKey cipher.EncryptionKey EncryptionKey cipher.EncryptionKey
Certificate *x509.Certificate Certificate *x509.Certificate
PrivateKey *rsa.PrivateKey PrivateKey *rsa.PrivateKey
@@ -158,6 +164,9 @@ func NewService(
magicLinkTokenValidity: cfg.MagicLinkTokenValidity, magicLinkTokenValidity: cfg.MagicLinkTokenValidity,
sessionDuration: cfg.SessionDuration, sessionDuration: cfg.SessionDuration,
bucket: cfg.Bucket, bucket: cfg.Bucket,
encryptionKey: cfg.EncryptionKey,
trustCenterBaseDomain: cfg.TrustCenterBaseDomain,
certManager: cfg.CertManager,
certificate: cfg.Certificate, certificate: cfg.Certificate,
privateKey: cfg.PrivateKey, privateKey: cfg.PrivateKey,
logger: cfg.Logger, logger: cfg.Logger,

View File

@@ -29,6 +29,7 @@ import (
"go.gearno.de/kit/pg" "go.gearno.de/kit/pg"
"go.probo.inc/probo/packages/emails" "go.probo.inc/probo/packages/emails"
"go.probo.inc/probo/pkg/baseurl" "go.probo.inc/probo/pkg/baseurl"
"go.probo.inc/probo/pkg/complianceportal/resolver"
"go.probo.inc/probo/pkg/coredata" "go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/mail" "go.probo.inc/probo/pkg/mail"
@@ -62,12 +63,12 @@ func (s *Service) mailingListEmailConfig(
mailingListID gid.GID, mailingListID gid.GID,
) (emails.PresenterConfig, string, string, *mail.Addr, error) { ) (emails.PresenterConfig, string, string, *mail.Addr, error) {
var ( var (
mailingList = &coredata.MailingList{} mailingList = &coredata.MailingList{}
compliancePage = &coredata.TrustCenter{} compliancePage = &coredata.TrustCenter{}
organization = &coredata.Organization{} organization = &coredata.Organization{}
customDomain *coredata.CustomDomain compliancePageURL string
logoFile = &coredata.File{} logoFile = &coredata.File{}
defaultCfg = emails.DefaultPresenterConfig(s.apiBaseURL.String()) defaultCfg = emails.DefaultPresenterConfig(s.apiBaseURL.String())
) )
scope := coredata.NewScopeFromObjectID(mailingListID) scope := coredata.NewScopeFromObjectID(mailingListID)
@@ -97,13 +98,19 @@ func (s *Service) mailingListEmailConfig(
return fmt.Errorf("cannot load organization: %w", err) return fmt.Errorf("cannot load organization: %w", err)
} }
customDomain = &coredata.CustomDomain{} publicURL, err := resolver.PublicURLForTrustCenter(
if err := customDomain.LoadByOrganizationID(ctx, conn, scope, organization.ID); err != nil { ctx,
if !errors.Is(err, coredata.ErrResourceNotFound) { conn,
return fmt.Errorf("cannot load custom domain: %w", err) scope,
} compliancePage,
s.trustCenterBaseDomain,
)
if err != nil {
return fmt.Errorf("cannot resolve compliance page URL: %w", err)
} }
compliancePageURL = publicURL
return nil return nil
}, },
) )
@@ -111,7 +118,7 @@ func (s *Service) mailingListEmailConfig(
return defaultCfg, "", "", nil, err return defaultCfg, "", "", nil, err
} }
cfg, compliancePageURL, err := s.presenterConfigFromTrustCenter(compliancePage, organization, customDomain, logoFile) cfg, err := s.presenterConfigFromTrustCenter(compliancePage, organization, compliancePageURL, logoFile)
if err != nil { if err != nil {
return defaultCfg, "", "", nil, err return defaultCfg, "", "", nil, err
} }
@@ -132,41 +139,24 @@ func (s *Service) mailingListEmailConfig(
func (s *Service) presenterConfigFromTrustCenter( func (s *Service) presenterConfigFromTrustCenter(
compliancePage *coredata.TrustCenter, compliancePage *coredata.TrustCenter,
organization *coredata.Organization, organization *coredata.Organization,
customDomain *coredata.CustomDomain, compliancePageURL string,
logoFile *coredata.File, logoFile *coredata.File,
) (emails.PresenterConfig, string, error) { ) (emails.PresenterConfig, error) {
cfg := emails.DefaultPresenterConfig(s.apiBaseURL.String()) cfg := emails.DefaultPresenterConfig(s.apiBaseURL.String())
compliancePageBase := s.apiBaseURL.WithPath("/trust/" + compliancePage.ID.String())
if customDomain != nil && customDomain.SSLStatus == coredata.CustomDomainSSLStatusActive {
customBase, err := baseurl.Parse("https://" + customDomain.Domain)
if err != nil {
return cfg, "", fmt.Errorf("cannot parse custom domain URL: %w", err)
}
compliancePageBase = customBase.WithPath("")
}
compliancePageURL, err := compliancePageBase.String()
if err != nil {
return cfg, "", fmt.Errorf("cannot build compliance page URL: %w", err)
}
cfg.BaseURL = compliancePageURL cfg.BaseURL = compliancePageURL
if compliancePage.LogoFileID != nil && logoFile != nil && logoFile.FileKey != "" { if compliancePage.LogoFileID != nil && logoFile != nil && logoFile.FileKey != "" {
cfg.SenderCompanyLogoPath = filepath.Join("/api/files/v1/public/", logoFile.ID.String()) cfg.SenderCompanyLogoPath = filepath.Join("/api/files/v1/public/", logoFile.ID.String())
cfg.SenderCompanyName = organization.Name cfg.SenderCompanyName = organization.Name
if organization.WebsiteURL != nil { if compliancePage.WebsiteURL != nil {
cfg.SenderCompanyWebsiteURL = *organization.WebsiteURL cfg.SenderCompanyWebsiteURL = *compliancePage.WebsiteURL
} }
if organization.HeadquarterAddress != nil { if compliancePage.HeadquarterAddress != nil {
cfg.SenderCompanyHeadquarterAddress = *organization.HeadquarterAddress cfg.SenderCompanyHeadquarterAddress = *compliancePage.HeadquarterAddress
} }
} }
return cfg, compliancePageURL, nil return cfg, nil
} }

View File

@@ -51,17 +51,36 @@ const (
) )
type Service struct { type Service struct {
pg *pg.Client pg *pg.Client
fm *filemanager.Service fm *filemanager.Service
tokenSecret string tokenSecret string
apiBaseURL *baseurl.BaseURL apiBaseURL *baseurl.BaseURL
bucket string trustCenterBaseDomain string
encryptionKey cipher.EncryptionKey bucket string
logger *log.Logger encryptionKey cipher.EncryptionKey
logger *log.Logger
} }
func NewService(pgClient *pg.Client, fm *filemanager.Service, tokenSecret string, apiBaseURL *baseurl.BaseURL, bucket string, encryptionKey cipher.EncryptionKey, logger *log.Logger) *Service { func NewService(
return &Service{pg: pgClient, fm: fm, tokenSecret: tokenSecret, apiBaseURL: apiBaseURL, bucket: bucket, encryptionKey: encryptionKey, logger: logger} pgClient *pg.Client,
fm *filemanager.Service,
tokenSecret string,
apiBaseURL *baseurl.BaseURL,
trustCenterBaseDomain string,
bucket string,
encryptionKey cipher.EncryptionKey,
logger *log.Logger,
) *Service {
return &Service{
pg: pgClient,
fm: fm,
tokenSecret: tokenSecret,
apiBaseURL: apiBaseURL,
trustCenterBaseDomain: trustCenterBaseDomain,
bucket: bucket,
encryptionKey: encryptionKey,
logger: logger,
}
} }
type ( type (

View File

@@ -52,6 +52,9 @@ import (
"go.probo.inc/probo/pkg/awsconfig" "go.probo.inc/probo/pkg/awsconfig"
"go.probo.inc/probo/pkg/baseurl" "go.probo.inc/probo/pkg/baseurl"
"go.probo.inc/probo/pkg/certmanager" "go.probo.inc/probo/pkg/certmanager"
"go.probo.inc/probo/pkg/complianceportal"
"go.probo.inc/probo/pkg/complianceportal/management"
trust "go.probo.inc/probo/pkg/complianceportal/visitor"
"go.probo.inc/probo/pkg/connector" "go.probo.inc/probo/pkg/connector"
"go.probo.inc/probo/pkg/connector/provider" "go.probo.inc/probo/pkg/connector/provider"
"go.probo.inc/probo/pkg/cookiebanner" "go.probo.inc/probo/pkg/cookiebanner"
@@ -80,7 +83,6 @@ import (
"go.probo.inc/probo/pkg/server/trustedproxy" "go.probo.inc/probo/pkg/server/trustedproxy"
"go.probo.inc/probo/pkg/slack" "go.probo.inc/probo/pkg/slack"
"go.probo.inc/probo/pkg/thirdparty" "go.probo.inc/probo/pkg/thirdparty"
"go.probo.inc/probo/pkg/trust"
"go.probo.inc/probo/pkg/webhook" "go.probo.inc/probo/pkg/webhook"
"golang.org/x/sync/errgroup" "golang.org/x/sync/errgroup"
) )
@@ -144,8 +146,9 @@ func New() *Implm {
}, },
}, },
TrustCenter: TrustCenterConfig{ TrustCenter: TrustCenterConfig{
HTTPAddr: ":80", HTTPAddr: ":80",
HTTPSAddr: ":443", HTTPSAddr: ":443",
BaseDomain: "probopage.com",
}, },
AWS: AWSConfig{ AWS: AWSConfig{
Region: "us-east-1", Region: "us-east-1",
@@ -491,54 +494,11 @@ func (impl *Implm) Run(
oauth2ScopeRegistry := oauth2scope.NewRegistry(). oauth2ScopeRegistry := oauth2scope.NewRegistry().
Register(iam.IAMOAuth2ScopeMappings). Register(iam.IAMOAuth2ScopeMappings).
Register(probo.OAuth2ScopeMappings). Register(probo.OAuth2ScopeMappings).
Register(complianceportal.OAuth2ScopeMappings).
Register(agentrun.OAuth2ScopeMappings). Register(agentrun.OAuth2ScopeMappings).
Register(accessreview.OAuth2ScopeMappings). Register(accessreview.OAuth2ScopeMappings).
Register(resourcealias.OAuth2ScopeMappings) Register(resourcealias.OAuth2ScopeMappings)
iamService, err := iam.NewService(
ctx,
pgClient,
fileManagerService,
hp,
iam.Config{
DisableSignup: impl.cfg.Auth.DisableSignup,
InvitationTokenValidity: time.Duration(impl.cfg.Auth.InvitationConfirmationTokenValidity) * time.Second,
PasswordResetTokenValidity: time.Duration(impl.cfg.Auth.PasswordResetTokenValidity) * time.Second,
MagicLinkTokenValidity: time.Duration(impl.cfg.Auth.MagicLinkTokenValidity) * time.Second,
SessionDuration: time.Duration(impl.cfg.Auth.Cookie.Duration) * time.Hour,
Bucket: impl.cfg.AWS.Bucket,
TokenSecret: impl.cfg.Auth.Cookie.Secret,
BaseURL: baseURL,
EncryptionKey: encryptionKey,
Certificate: samlCert,
PrivateKey: samlKey,
Logger: l.Named("iam"),
TracerProvider: tp,
Registerer: r,
ConnectorRegistry: defaultConnectorRegistry,
DomainVerificationInterval: impl.cfg.Auth.SAML.DomainVerificationInterval(),
DomainVerificationResolverAddr: impl.cfg.Auth.SAML.DomainVerificationResolverAddr,
SCIMBridgeSyncInterval: time.Duration(impl.cfg.SCIMBridge.SyncInterval) * time.Second,
SCIMBridgePollInterval: time.Duration(impl.cfg.SCIMBridge.PollInterval) * time.Second,
GoogleOIDC: oidc.ProviderConfig{
ClientID: impl.cfg.Auth.Google.ClientID,
ClientSecret: impl.cfg.Auth.Google.ClientSecret,
Enabled: impl.cfg.Auth.Google.Enabled,
},
MicrosoftOIDC: oidc.ProviderConfig{
ClientID: impl.cfg.Auth.Microsoft.ClientID,
ClientSecret: impl.cfg.Auth.Microsoft.ClientSecret,
Enabled: impl.cfg.Auth.Microsoft.Enabled,
},
OAuth2ServerSigningKeys: oauth2SigningKeys,
OAuth2ServerOptions: oauth2ServerOptions(impl.cfg.Auth.OAuth2Server),
OAuth2ScopeRegistry: oauth2ScopeRegistry,
},
)
if err != nil {
return fmt.Errorf("cannot create iam service: %w", err)
}
var accountKey crypto.Signer var accountKey crypto.Signer
if impl.cfg.CustomDomains.ACME.AccountKey != "" { if impl.cfg.CustomDomains.ACME.AccountKey != "" {
accountKey, err = pemutil.DecodePrivateKey([]byte(impl.cfg.CustomDomains.ACME.AccountKey)) accountKey, err = pemutil.DecodePrivateKey([]byte(impl.cfg.CustomDomains.ACME.AccountKey))
@@ -569,6 +529,77 @@ func (impl *Implm) Run(
return fmt.Errorf("cannot initialize ACME service: %w", err) return fmt.Errorf("cannot initialize ACME service: %w", err)
} }
customDomainRenewalInterval := time.Duration(impl.cfg.CustomDomains.RenewalInterval) * time.Second
if customDomainRenewalInterval == 0 {
customDomainRenewalInterval = time.Hour
}
customDomainProvisionInterval := time.Duration(impl.cfg.CustomDomains.ProvisionInterval) * time.Second
if customDomainProvisionInterval == 0 {
customDomainProvisionInterval = 30 * time.Second
}
certManagerService := certmanager.NewService(
pgClient,
acmeService,
encryptionKey,
certmanager.Config{
CnameTarget: impl.cfg.CustomDomains.CnameTarget,
CAAIssuerDomain: impl.cfg.CustomDomains.CAAIssuerDomain,
ResolverAddr: impl.cfg.CustomDomains.ResolverAddr,
ManagedBaseDomain: impl.cfg.TrustCenter.BaseDomain,
RenewalInterval: customDomainRenewalInterval,
ProvisionInterval: customDomainProvisionInterval,
},
l.Named("certmanager"),
)
iamService, err := iam.NewService(
ctx,
pgClient,
fileManagerService,
hp,
iam.Config{
DisableSignup: impl.cfg.Auth.DisableSignup,
InvitationTokenValidity: time.Duration(impl.cfg.Auth.InvitationConfirmationTokenValidity) * time.Second,
PasswordResetTokenValidity: time.Duration(impl.cfg.Auth.PasswordResetTokenValidity) * time.Second,
MagicLinkTokenValidity: time.Duration(impl.cfg.Auth.MagicLinkTokenValidity) * time.Second,
SessionDuration: time.Duration(impl.cfg.Auth.Cookie.Duration) * time.Hour,
Bucket: impl.cfg.AWS.Bucket,
TokenSecret: impl.cfg.Auth.Cookie.Secret,
BaseURL: baseURL,
TrustCenterBaseDomain: impl.cfg.TrustCenter.BaseDomain,
EncryptionKey: encryptionKey,
Certificate: samlCert,
PrivateKey: samlKey,
Logger: l.Named("iam"),
TracerProvider: tp,
Registerer: r,
ConnectorRegistry: defaultConnectorRegistry,
DomainVerificationInterval: impl.cfg.Auth.SAML.DomainVerificationInterval(),
DomainVerificationResolverAddr: impl.cfg.Auth.SAML.DomainVerificationResolverAddr,
SCIMBridgeSyncInterval: time.Duration(impl.cfg.SCIMBridge.SyncInterval) * time.Second,
SCIMBridgePollInterval: time.Duration(impl.cfg.SCIMBridge.PollInterval) * time.Second,
GoogleOIDC: oidc.ProviderConfig{
ClientID: impl.cfg.Auth.Google.ClientID,
ClientSecret: impl.cfg.Auth.Google.ClientSecret,
Enabled: impl.cfg.Auth.Google.Enabled,
},
MicrosoftOIDC: oidc.ProviderConfig{
ClientID: impl.cfg.Auth.Microsoft.ClientID,
ClientSecret: impl.cfg.Auth.Microsoft.ClientSecret,
Enabled: impl.cfg.Auth.Microsoft.Enabled,
},
OAuth2ServerSigningKeys: oauth2SigningKeys,
OAuth2ServerOptions: oauth2ServerOptions(impl.cfg.Auth.OAuth2Server),
OAuth2ScopeRegistry: oauth2ScopeRegistry,
CertManager: certManagerService,
},
)
if err != nil {
return fmt.Errorf("cannot create iam service: %w", err)
}
slackService := slack.NewService( slackService := slack.NewService(
pgClient, pgClient,
impl.cfg.GetSlackSigningSecret(), impl.cfg.GetSlackSigningSecret(),
@@ -586,7 +617,16 @@ func (impl *Implm) Run(
l.Named("esign"), l.Named("esign"),
) )
mailmanService := mailman.NewService(pgClient, fileManagerService, impl.cfg.Auth.Cookie.Secret, baseURL, impl.cfg.AWS.Bucket, encryptionKey, l) mailmanService := mailman.NewService(
pgClient,
fileManagerService,
impl.cfg.Auth.Cookie.Secret,
baseURL,
impl.cfg.TrustCenter.BaseDomain,
impl.cfg.AWS.Bucket,
encryptionKey,
l,
)
cookieBannerService := cookiebanner.NewService(pgClient, impl.cfg.Branding) cookieBannerService := cookiebanner.NewService(pgClient, impl.cfg.Branding)
@@ -605,7 +645,6 @@ func (impl *Implm) Run(
MaxTokens: ref.UnrefOrZero(proboAgentCfg.MaxTokens), MaxTokens: ref.UnrefOrZero(proboAgentCfg.MaxTokens),
}, },
html2pdfConverter, html2pdfConverter,
acmeService,
fileManagerService, fileManagerService,
l.Named("probo"), l.Named("probo"),
slackService, slackService,
@@ -620,11 +659,24 @@ func (impl *Implm) Run(
resourceAliasService := resourcealias.NewService(pgClient) resourceAliasService := resourcealias.NewService(pgClient)
managementService := management.NewService(
pgClient,
s3Client,
impl.cfg.AWS.Bucket,
baseURL.String(),
impl.cfg.TrustCenter.BaseDomain,
fileManagerService,
certManagerService,
slackService,
l.Named("compliance-portal-management"),
)
trustService := trust.NewService( trustService := trust.NewService(
pgClient, pgClient,
s3Client, s3Client,
impl.cfg.AWS.Bucket, impl.cfg.AWS.Bucket,
baseURL.String(), baseURL.String(),
impl.cfg.TrustCenter.BaseDomain,
impl.cfg.GetSlackSigningSecret(), impl.cfg.GetSlackSigningSecret(),
iamService, iamService,
esignService, esignService,
@@ -648,6 +700,7 @@ func (impl *Implm) Run(
iamService.Authorizer.RegisterPolicySet(agentrun.PolicySet()) iamService.Authorizer.RegisterPolicySet(agentrun.PolicySet())
iamService.Authorizer.RegisterPolicySet(accessreview.PolicySet()) iamService.Authorizer.RegisterPolicySet(accessreview.PolicySet())
iamService.Authorizer.RegisterPolicySet(resourcealias.PolicySet()) iamService.Authorizer.RegisterPolicySet(resourcealias.PolicySet())
iamService.Authorizer.RegisterPolicySet(complianceportal.PolicySet())
thirdPartyService := thirdparty.NewService(pgClient, fileManagerService, thirdPartyVetter) thirdPartyService := thirdparty.NewService(pgClient, fileManagerService, thirdPartyVetter)
riskManagementService := riskmanagement.NewService(pgClient) riskManagementService := riskmanagement.NewService(pgClient)
@@ -662,6 +715,7 @@ func (impl *Implm) Run(
IAM: iamService, IAM: iamService,
Trust: trustService, Trust: trustService,
ESign: esignService, ESign: esignService,
CustomDomain: managementService,
AccessReview: accessReviewService, AccessReview: accessReviewService,
AgentRun: agentRunService, AgentRun: agentRunService,
Mailman: mailmanService, Mailman: mailmanService,
@@ -679,7 +733,6 @@ func (impl *Implm) Run(
QueryCacheSize: impl.cfg.Api.GraphQL.QueryCacheSize, QueryCacheSize: impl.cfg.Api.GraphQL.QueryCacheSize,
DisableSuggestion: impl.cfg.Api.GraphQL.DisableSuggestion, DisableSuggestion: impl.cfg.Api.GraphQL.DisableSuggestion,
}, },
CustomDomainCname: impl.cfg.CustomDomains.CnameTarget, CustomDomainCname: impl.cfg.CustomDomains.CnameTarget,
TokenSecret: impl.cfg.Auth.Cookie.Secret, TokenSecret: impl.cfg.Auth.Cookie.Secret,
Logger: l.Named("http.server"), Logger: l.Named("http.server"),
@@ -856,12 +909,22 @@ func (impl *Implm) Run(
wg.Go( wg.Go(
func() { func() {
if err := esignService.Run(esignServiceCtx, trustService.EmailPresenterConfigByOrganizationID); err != nil { if err := esignService.Run(esignServiceCtx, trustService.GetPortalEmailPresenterConfigByOrganizationID); err != nil {
cancel(fmt.Errorf("esign service crashed: %w", err)) cancel(fmt.Errorf("esign service crashed: %w", err))
} }
}, },
) )
certManagerServiceCtx, stopCertManagerService := context.WithCancel(context.Background())
wg.Go(
func() {
if err := certManagerService.Run(certManagerServiceCtx); err != nil {
cancel(fmt.Errorf("certificate manager service crashed: %w", err))
}
},
)
trackerPatternAnalysisWorker := cookiebanner.NewPatternAnalysisWorker(cookieBannerService, pgClient, l) trackerPatternAnalysisWorker := cookiebanner.NewPatternAnalysisWorker(cookieBannerService, pgClient, l)
trackerPatternAnalysisWorkerCtx, stopTrackerPatternAnalysisWorker := context.WithCancel(context.Background()) trackerPatternAnalysisWorkerCtx, stopTrackerPatternAnalysisWorker := context.WithCancel(context.Background())
@@ -1033,8 +1096,7 @@ func (impl *Implm) Run(
tp, tp,
pgClient, pgClient,
serverHandler.TrustCenterHandler(), serverHandler.TrustCenterHandler(),
acmeService, trustService,
proboService,
encryptionKey, encryptionKey,
); err != nil { ); err != nil {
cancel(fmt.Errorf("trust center server crashed: %w", err)) cancel(fmt.Errorf("trust center server crashed: %w", err))
@@ -1048,6 +1110,7 @@ func (impl *Implm) Run(
stopTrustCenterServer() stopTrustCenterServer()
stopWebhookWorker() stopWebhookWorker()
stopESignService() stopESignService()
stopCertManagerService()
stopTrackerPatternAnalysisWorker() stopTrackerPatternAnalysisWorker()
stopTrackerPolicyWorker() stopTrackerPolicyWorker()
stopTrackerMappingWorker() stopTrackerMappingWorker()
@@ -1186,7 +1249,7 @@ func (impl *Implm) runApiServer(
return ctx.Err() return ctx.Err()
} }
func newTrustCenterHTTPRedirectHandler(proboService *probo.Service, l *log.Logger) http.Handler { func newTrustCenterHTTPRedirectHandler(trustService *trust.Service, l *log.Logger) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
ctx := r.Context() ctx := r.Context()
@@ -1202,11 +1265,15 @@ func newTrustCenterHTTPRedirectHandler(proboService *probo.Service, l *log.Logge
return return
} }
// Check if this domain is a trust center domain // Check if this domain is a trust center custom domain
_, err := proboService.LoadOrganizationByDomain(ctx, domain) if _, err := trustService.GetPortalByDomainName(ctx, domain); err != nil {
if err != nil { if errors.Is(err, trust.ErrPageNotFound) || errors.Is(err, coredata.ErrResourceNotFound) {
// Not a trust center domain, return 404 // Not a trust center domain, return 404
httpserver.RenderError(w, http.StatusNotFound, errors.New("not found")) httpserver.RenderError(w, http.StatusNotFound, errors.New("not found"))
return
}
httpserver.RenderError(w, http.StatusInternalServerError, errors.New("internal server error"))
return return
} }
@@ -1236,8 +1303,7 @@ func (impl *Implm) runTrustCenterServer(
tp trace.TracerProvider, tp trace.TracerProvider,
pgClient *pg.Client, pgClient *pg.Client,
trustRouter http.Handler, trustRouter http.Handler,
acmeService *certmanager.ACMEService, trustService *trust.Service,
proboService *probo.Service,
encryptionKey cipher.EncryptionKey, encryptionKey cipher.EncryptionKey,
) error { ) error {
tracer := tp.Tracer("go.probo.inc/probo/pkg/probod") tracer := tp.Tracer("go.probo.inc/probo/pkg/probod")
@@ -1253,45 +1319,17 @@ func (impl *Implm) runTrustCenterServer(
l.ErrorCtx(ctx, "cannot warm certificate cache", log.Error(err)) l.ErrorCtx(ctx, "cannot warm certificate cache", log.Error(err))
} }
renewalInterval := time.Duration(impl.cfg.CustomDomains.RenewalInterval) * time.Second
if renewalInterval == 0 {
renewalInterval = time.Hour
}
renewer := certmanager.NewRenewer(pgClient, acmeService, encryptionKey, renewalInterval, l)
certProvisioningInterval := time.Duration(impl.cfg.CustomDomains.ProvisionInterval) * time.Second
if certProvisioningInterval == 0 {
certProvisioningInterval = 30 * time.Second
}
certProvisioner := certmanager.NewProvisioner(pgClient, acmeService, encryptionKey, impl.cfg.CustomDomains.CnameTarget, impl.cfg.CustomDomains.CAAIssuerDomain, certProvisioningInterval, impl.cfg.CustomDomains.ResolverAddr, l)
g, ctx := errgroup.WithContext(ctx) g, ctx := errgroup.WithContext(ctx)
l.Info("starting trust center services") l.Info("starting trust center services")
span.AddEvent("Trust center services starting") span.AddEvent("Trust center services starting")
g.Go(
func() error {
l.Info("starting certificate renewer")
return renewer.Run(ctx)
},
)
g.Go(
func() error {
l.Info("starting certificate provisioner")
return certProvisioner.Run(ctx)
},
)
httpACMEHandler := certmanager.NewACMEChallengeHandler( httpACMEHandler := certmanager.NewACMEChallengeHandler(
pgClient, pgClient,
l.Named("http_acme_handler"), l.Named("http_acme_handler"),
) )
httpRedirectHandler := newTrustCenterHTTPRedirectHandler(proboService, l.Named("http_redirect")) httpRedirectHandler := newTrustCenterHTTPRedirectHandler(trustService, l.Named("http_redirect"))
httpServer := httpserver.NewServer( httpServer := httpserver.NewServer(
impl.cfg.TrustCenter.HTTPAddr, impl.cfg.TrustCenter.HTTPAddr,

View File

@@ -34,6 +34,8 @@ import (
"go.probo.inc/probo/pkg/accessreview" "go.probo.inc/probo/pkg/accessreview"
"go.probo.inc/probo/pkg/agentrun" "go.probo.inc/probo/pkg/agentrun"
"go.probo.inc/probo/pkg/baseurl" "go.probo.inc/probo/pkg/baseurl"
"go.probo.inc/probo/pkg/complianceportal/management"
trust "go.probo.inc/probo/pkg/complianceportal/visitor"
"go.probo.inc/probo/pkg/connector" "go.probo.inc/probo/pkg/connector"
"go.probo.inc/probo/pkg/connector/provider" "go.probo.inc/probo/pkg/connector/provider"
"go.probo.inc/probo/pkg/cookiebanner" "go.probo.inc/probo/pkg/cookiebanner"
@@ -56,7 +58,6 @@ import (
"go.probo.inc/probo/pkg/server/gqlutils" "go.probo.inc/probo/pkg/server/gqlutils"
"go.probo.inc/probo/pkg/slack" "go.probo.inc/probo/pkg/slack"
"go.probo.inc/probo/pkg/thirdparty" "go.probo.inc/probo/pkg/thirdparty"
"go.probo.inc/probo/pkg/trust"
) )
type ( type (
@@ -69,6 +70,7 @@ type (
IAM *iam.Service IAM *iam.Service
Trust *trust.Service Trust *trust.Service
ESign *esign.Service ESign *esign.Service
CustomDomain *management.Service
AccessReview *accessreview.Service AccessReview *accessreview.Service
AgentRun *agentrun.Service AgentRun *agentrun.Service
Slack *slack.Service Slack *slack.Service
@@ -204,6 +206,7 @@ func NewServer(cfg Config) (*Server, error) {
cfg.ResourceAlias, cfg.ResourceAlias,
cfg.IAM, cfg.IAM,
cfg.ESign, cfg.ESign,
cfg.CustomDomain,
cfg.AccessReview, cfg.AccessReview,
cfg.AgentRun, cfg.AgentRun,
cfg.Mailman, cfg.Mailman,
@@ -236,6 +239,7 @@ func NewServer(cfg Config) (*Server, error) {
mcpHandler: mcp_v1.NewMux( mcpHandler: mcp_v1.NewMux(
cfg.Logger.Named("mcp.v1"), cfg.Logger.Named("mcp.v1"),
cfg.Probo, cfg.Probo,
cfg.CustomDomain,
cfg.ResourceAlias, cfg.ResourceAlias,
cfg.ThirdParty, cfg.ThirdParty,
cfg.IAM, cfg.IAM,
@@ -263,12 +267,12 @@ func NewServer(cfg Config) (*Server, error) {
return true return true
} }
_, err := cfg.Trust.GetByDomainName(ctx, host) _, err := cfg.Trust.GetPortalByDomainName(ctx, host)
return err == nil return err == nil
}, },
func(ctx context.Context, host string) bool { func(ctx context.Context, host string) bool {
_, err := cfg.Trust.GetByDomainName(ctx, host) _, err := cfg.Trust.GetPortalByDomainName(ctx, host)
return err == nil return err == nil
}, },
cfg.GraphQLLimits, cfg.GraphQLLimits,
@@ -314,7 +318,6 @@ func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) {
r.Mount("/console/v1", http.StripPrefix("/console/v1", s.consoleHandler)) r.Mount("/console/v1", http.StripPrefix("/console/v1", s.consoleHandler))
r.Mount("/connect/v1", http.StripPrefix("/connect/v1", s.connectHandler)) r.Mount("/connect/v1", http.StripPrefix("/connect/v1", s.connectHandler))
r.Mount("/files/v1", http.StripPrefix("/files/v1", s.filesHandler)) r.Mount("/files/v1", http.StripPrefix("/files/v1", s.filesHandler))
r.Mount("/trust/v1", http.StripPrefix("/trust/v1", s.compliancePageHandler))
r.Mount("/mcp/v1", http.StripPrefix("/mcp/v1", s.mcpHandler)) r.Mount("/mcp/v1", http.StripPrefix("/mcp/v1", s.mcpHandler))
r.Mount("/slack/v1", http.StripPrefix("/slack/v1", s.slackHandler)) r.Mount("/slack/v1", http.StripPrefix("/slack/v1", s.slackHandler))
}) })

View File

@@ -23,8 +23,6 @@ package server
import ( import (
"errors" "errors"
"net/http" "net/http"
"path"
"strings"
"github.com/go-chi/chi/v5" "github.com/go-chi/chi/v5"
"go.gearno.de/kit/httpserver" "go.gearno.de/kit/httpserver"
@@ -33,6 +31,8 @@ import (
"go.probo.inc/probo/pkg/accessreview" "go.probo.inc/probo/pkg/accessreview"
"go.probo.inc/probo/pkg/agentrun" "go.probo.inc/probo/pkg/agentrun"
"go.probo.inc/probo/pkg/baseurl" "go.probo.inc/probo/pkg/baseurl"
"go.probo.inc/probo/pkg/complianceportal/management"
trust "go.probo.inc/probo/pkg/complianceportal/visitor"
"go.probo.inc/probo/pkg/connector" "go.probo.inc/probo/pkg/connector"
"go.probo.inc/probo/pkg/connector/provider" "go.probo.inc/probo/pkg/connector/provider"
"go.probo.inc/probo/pkg/cookiebanner" "go.probo.inc/probo/pkg/cookiebanner"
@@ -47,14 +47,13 @@ import (
"go.probo.inc/probo/pkg/riskmanagement" "go.probo.inc/probo/pkg/riskmanagement"
"go.probo.inc/probo/pkg/securecookie" "go.probo.inc/probo/pkg/securecookie"
"go.probo.inc/probo/pkg/server/api" "go.probo.inc/probo/pkg/server/api"
"go.probo.inc/probo/pkg/server/api/compliancepage" "go.probo.inc/probo/pkg/server/api/complianceportal"
"go.probo.inc/probo/pkg/server/gqlutils" "go.probo.inc/probo/pkg/server/gqlutils"
"go.probo.inc/probo/pkg/server/mailactions" "go.probo.inc/probo/pkg/server/mailactions"
trust_web "go.probo.inc/probo/pkg/server/trust" trust_web "go.probo.inc/probo/pkg/server/trust"
console_web "go.probo.inc/probo/pkg/server/web" console_web "go.probo.inc/probo/pkg/server/web"
"go.probo.inc/probo/pkg/slack" "go.probo.inc/probo/pkg/slack"
"go.probo.inc/probo/pkg/thirdparty" "go.probo.inc/probo/pkg/thirdparty"
"go.probo.inc/probo/pkg/trust"
"go.probo.inc/probo/pkg/uri" "go.probo.inc/probo/pkg/uri"
) )
@@ -68,6 +67,7 @@ type Config struct {
IAM *iam.Service IAM *iam.Service
Trust *trust.Service Trust *trust.Service
ESign *esign.Service ESign *esign.Service
CustomDomain *management.Service
AccessReview *accessreview.Service AccessReview *accessreview.Service
AgentRun *agentrun.Service AgentRun *agentrun.Service
Slack *slack.Service Slack *slack.Service
@@ -109,6 +109,7 @@ func NewServer(cfg Config) (*Server, error) {
IAM: cfg.IAM, IAM: cfg.IAM,
Trust: cfg.Trust, Trust: cfg.Trust,
ESign: cfg.ESign, ESign: cfg.ESign,
CustomDomain: cfg.CustomDomain,
AccessReview: cfg.AccessReview, AccessReview: cfg.AccessReview,
AgentRun: cfg.AgentRun, AgentRun: cfg.AgentRun,
Slack: cfg.Slack, Slack: cfg.Slack,
@@ -157,12 +158,12 @@ func NewServer(cfg Config) (*Server, error) {
logger: cfg.Logger, logger: cfg.Logger,
} }
server.setupRoutes(cfg.BaseURL.String()) server.setupRoutes()
return server, nil return server, nil
} }
func (s *Server) setupRoutes(baseURL string) { func (s *Server) setupRoutes() {
// OIDC Discovery 1.0 §4 and RFC 8414 §3 both require the metadata // OIDC Discovery 1.0 §4 and RFC 8414 §3 both require the metadata
// document at the issuer root under well-known paths. // document at the issuer root under well-known paths.
s.router.Get("/.well-known/openid-configuration", s.oidcDiscoveryHandler) s.router.Get("/.well-known/openid-configuration", s.oidcDiscoveryHandler)
@@ -172,12 +173,6 @@ func (s *Server) setupRoutes(baseURL string) {
s.router.Mount("/api", http.StripPrefix("/api", s.apiServer)) s.router.Mount("/api", http.StripPrefix("/api", s.apiServer))
s.router.Mount("/mail-actions", http.StripPrefix("/mail-actions", s.mailActionsHandler)) s.router.Mount("/mail-actions", http.StripPrefix("/mail-actions", s.mailActionsHandler))
s.router.Route("/trust/{slugOrId}", func(r chi.Router) {
r.Use(compliancepage.NewIDMiddleware(s.trustService, baseURL))
r.Use(s.stripTrustPrefix)
r.Mount("/", s.trustCenterRouter())
})
s.router.Mount("/", s.consoleWebServer) s.router.Mount("/", s.consoleWebServer)
} }
@@ -224,31 +219,10 @@ func (s *Server) handleCustomDomain404(w http.ResponseWriter, r *http.Request) {
httpserver.RenderError(w, http.StatusNotFound, errors.New("not found")) httpserver.RenderError(w, http.StatusNotFound, errors.New("not found"))
} }
func (s *Server) stripTrustPrefix(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
slugOrId := chi.URLParam(r, "slugOrId")
prefix := "/trust/" + slugOrId
if r.URL.Path == prefix {
cleanPath := path.Clean(prefix) + "/"
http.Redirect(w, r, cleanPath, http.StatusMovedPermanently)
return
}
r.URL.Path = strings.TrimPrefix(r.URL.Path, prefix)
if r.URL.Path == "" {
r.URL.Path = "/"
}
next.ServeHTTP(w, r)
})
}
func (s *Server) trustCenterRouter() chi.Router { func (s *Server) trustCenterRouter() chi.Router {
r := chi.NewRouter() r := chi.NewRouter()
h := compliancepage.NewHandler(s.trustService) h := complianceportal.NewHandler(s.trustService)
r.Mount("/api/trust/v1", s.apiServer.CompliancePageHandler()) r.Mount("/api/trust/v1", s.apiServer.CompliancePageHandler())
r.Get("/llms.txt", h.HandleLLMsTxt) r.Get("/llms.txt", h.HandleLLMsTxt)
@@ -262,7 +236,7 @@ func (s *Server) trustCenterRouter() chi.Router {
func (s *Server) TrustCenterHandler() http.Handler { func (s *Server) TrustCenterHandler() http.Handler {
r := chi.NewRouter() r := chi.NewRouter()
r.Use(compliancepage.NewSNIMiddleware(s.trustService)) r.Use(complianceportal.NewSNIMiddleware(s.trustService))
r.Use(func(next http.Handler) http.Handler { r.Use(func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Strict-Transport-Security", "max-age=31536000; preload") w.Header().Set("Strict-Transport-Security", "max-age=31536000; preload")
@@ -280,21 +254,26 @@ func (s *Server) TrustCenterHandler() http.Handler {
func compliancePageHeadData(baseURL *baseurl.BaseURL, trustService *trust.Service) trust_web.HeadDataFunc { func compliancePageHeadData(baseURL *baseurl.BaseURL, trustService *trust.Service) trust_web.HeadDataFunc {
return func(r *http.Request) trust_web.HeadData { return func(r *http.Request) trust_web.HeadData {
tc := compliancepage.CompliancePageFromContext(r.Context()) tc := complianceportal.CompliancePageFromContext(r.Context())
if tc == nil { if tc == nil {
return trust_web.HeadData{Title: "Compliance Page"} return trust_web.HeadData{Title: "Compliance Page"}
} }
org, err := trustService.GetOrganizationByTrustCenterID(r.Context(), tc.ID) org, err := trustService.GetPortalOrganization(r.Context(), tc.ID)
if err != nil || org == nil { if err != nil || org == nil {
return trust_web.HeadData{Title: "Compliance Page"} return trust_web.HeadData{Title: "Compliance Page"}
} }
compliancePageBaseURL := compliancepage.CompliancePageBaseURLFromContext(r.Context()) compliancePageBaseURL := complianceportal.CompliancePageBaseURLFromContext(r.Context())
description := org.Name + " Compliance Page"
if tc.Description != nil && *tc.Description != "" {
description = *tc.Description
}
headData := trust_web.HeadData{ headData := trust_web.HeadData{
Title: org.Name + " — Compliance", Title: org.Name + " — Compliance",
Description: org.Name + " Compliance Page", Description: description,
OGURL: ref.UnrefOrZero(compliancePageBaseURL), OGURL: ref.UnrefOrZero(compliancePageBaseURL),
} }