@@ -37,14 +37,17 @@ func (impl *Implm) resolveAgentClient(
|
||||
r prometheus.Registerer,
|
||||
) (LLMAgentConfig, *llm.Client, error) {
|
||||
resolved := impl.cfg.Agents.ResolveAgent(agent)
|
||||
|
||||
providerCfg, ok := impl.cfg.Agents.Providers[resolved.Provider]
|
||||
if !ok {
|
||||
return LLMAgentConfig{}, nil, fmt.Errorf("unknown LLM provider %q for %s agent", resolved.Provider, name)
|
||||
}
|
||||
|
||||
client, err := buildLLMClient(providerCfg, l.Named("llm."+name), tp, r)
|
||||
if err != nil {
|
||||
return LLMAgentConfig{}, nil, fmt.Errorf("cannot create %s LLM client: %w", name, err)
|
||||
}
|
||||
|
||||
return resolved, client, nil
|
||||
}
|
||||
|
||||
@@ -66,6 +69,7 @@ func buildLLMClient(cfg LLMProviderConfig, l *log.Logger, tp trace.TracerProvide
|
||||
cfg.APIKey,
|
||||
llmopenai.WithHTTPClient(httpClient),
|
||||
)
|
||||
|
||||
return llm.NewClient(
|
||||
p,
|
||||
"openai",
|
||||
@@ -77,6 +81,7 @@ func buildLLMClient(cfg LLMProviderConfig, l *log.Logger, tp trace.TracerProvide
|
||||
cfg.APIKey,
|
||||
llmanthropic.WithHTTPClient(httpClient),
|
||||
)
|
||||
|
||||
return llm.NewClient(
|
||||
p,
|
||||
"anthropic",
|
||||
|
||||
@@ -191,6 +191,7 @@ func (impl *Implm) Run(
|
||||
tp trace.TracerProvider,
|
||||
) error {
|
||||
tracer := tp.Tracer("probod")
|
||||
|
||||
ctx, rootSpan := tracer.Start(parentCtx, "probod.Run")
|
||||
defer rootSpan.End()
|
||||
|
||||
@@ -208,6 +209,7 @@ func (impl *Implm) Run(
|
||||
}
|
||||
|
||||
wg := sync.WaitGroup{}
|
||||
|
||||
ctx, cancel := context.WithCancelCause(ctx)
|
||||
defer cancel(context.Canceled)
|
||||
|
||||
@@ -267,6 +269,7 @@ func (impl *Implm) Run(
|
||||
}
|
||||
|
||||
geolocService := geoloc.NewService(pgClient)
|
||||
|
||||
populated, err := geolocService.IsPopulated(ctx)
|
||||
if err != nil {
|
||||
l.ErrorCtx(ctx, "cannot check geoloc table", log.Error(err))
|
||||
@@ -281,10 +284,12 @@ func (impl *Implm) Run(
|
||||
|
||||
redirectURI := baseURL.WithPath(connector.CallbackPath).MustString()
|
||||
defaultConnectorRegistry := connector.NewConnectorRegistry()
|
||||
|
||||
for _, connectorCfg := range impl.cfg.Connectors {
|
||||
if oauth2c, ok := connectorCfg.Config.(*connector.OAuth2Connector); ok {
|
||||
connector.ApplyProviderDefaults(connectorCfg.Provider, redirectURI, oauth2c)
|
||||
}
|
||||
|
||||
if err := defaultConnectorRegistry.Register(connectorCfg.Provider, connectorCfg.Config); err != nil {
|
||||
return fmt.Errorf("cannot register connector: %w", err)
|
||||
}
|
||||
@@ -312,15 +317,20 @@ func (impl *Implm) Run(
|
||||
|
||||
fileManagerService := filemanager.NewService(s3Client)
|
||||
|
||||
var samlCert *x509.Certificate
|
||||
var samlKey *rsa.PrivateKey
|
||||
var (
|
||||
samlCert *x509.Certificate
|
||||
samlKey *rsa.PrivateKey
|
||||
)
|
||||
|
||||
if impl.cfg.Auth.SAML.Certificate != "" && impl.cfg.Auth.SAML.PrivateKey != "" {
|
||||
// Decode certificate
|
||||
certBlock, _ := pem.Decode([]byte(impl.cfg.Auth.SAML.Certificate))
|
||||
if certBlock == nil {
|
||||
return fmt.Errorf("cannot decode SAML certificate PEM block")
|
||||
}
|
||||
|
||||
var err error
|
||||
|
||||
samlCert, err = x509.ParseCertificate(certBlock.Bytes)
|
||||
if err != nil {
|
||||
return fmt.Errorf("cannot parse SAML certificate: %w", err)
|
||||
@@ -331,7 +341,9 @@ func (impl *Implm) Run(
|
||||
if err != nil {
|
||||
return fmt.Errorf("cannot decode SAML private key: %w", err)
|
||||
}
|
||||
|
||||
var ok bool
|
||||
|
||||
samlKey, ok = signer.(*rsa.PrivateKey)
|
||||
if !ok {
|
||||
return fmt.Errorf("SAML private key is not an RSA key")
|
||||
@@ -342,8 +354,11 @@ func (impl *Implm) Run(
|
||||
return fmt.Errorf("cannot configure OAuth2 server: at least one signing key is required")
|
||||
}
|
||||
|
||||
var oauth2SigningKeys oauth2server.SigningKeys
|
||||
var hasActive bool
|
||||
var (
|
||||
oauth2SigningKeys oauth2server.SigningKeys
|
||||
hasActive bool
|
||||
)
|
||||
|
||||
for _, keyCfg := range impl.cfg.Auth.OAuth2Server.SigningKeys {
|
||||
signer, err := pemutil.DecodePrivateKey([]byte(keyCfg.PrivateKey))
|
||||
if err != nil {
|
||||
@@ -432,6 +447,7 @@ func (impl *Implm) Run(
|
||||
if err != nil {
|
||||
return fmt.Errorf("cannot decode ACME account key: %w", err)
|
||||
}
|
||||
|
||||
l.Info("using configured ACME account key")
|
||||
}
|
||||
|
||||
@@ -569,6 +585,7 @@ func (impl *Implm) Run(
|
||||
|
||||
apiServerCtx, stopApiServer := context.WithCancel(context.Background())
|
||||
defer stopApiServer()
|
||||
|
||||
wg.Go(
|
||||
func() {
|
||||
if err := impl.runApiServer(apiServerCtx, l, r, tp, serverHandler); err != nil {
|
||||
@@ -596,6 +613,7 @@ func (impl *Implm) Run(
|
||||
worker.WithInterval(time.Duration(impl.cfg.Notifications.Mailer.MailerInterval)*time.Second),
|
||||
worker.WithMaxConcurrency(20),
|
||||
)
|
||||
|
||||
wg.Go(
|
||||
func() {
|
||||
if err := sendingWorker.Run(mailerCtx); err != nil {
|
||||
@@ -608,6 +626,7 @@ func (impl *Implm) Run(
|
||||
slackSender := slack.NewSender(pgClient, l.Named("slack-sender"), encryptionKey, slack.Config{
|
||||
Interval: time.Duration(impl.cfg.Notifications.Slack.SenderInterval) * time.Second,
|
||||
})
|
||||
|
||||
wg.Go(
|
||||
func() {
|
||||
if err := slackSender.Run(slackSenderCtx); err != nil {
|
||||
@@ -623,6 +642,7 @@ func (impl *Implm) Run(
|
||||
EncryptionKey: encryptionKey,
|
||||
Host: baseURL.String(),
|
||||
})
|
||||
|
||||
wg.Go(
|
||||
func() {
|
||||
if err := webhookSender.Run(webhookSenderCtx); err != nil {
|
||||
@@ -632,6 +652,7 @@ func (impl *Implm) Run(
|
||||
)
|
||||
|
||||
exportJobExporterCtx, stopExportJobExporter := context.WithCancel(context.Background())
|
||||
|
||||
wg.Go(
|
||||
func() {
|
||||
if err := impl.runExportJob(exportJobExporterCtx, proboService, l.Named("export-job-exporter")); err != nil {
|
||||
@@ -646,6 +667,7 @@ func (impl *Implm) Run(
|
||||
worker.WithInterval(30*time.Second),
|
||||
)
|
||||
documentPDFWorkerCtx, stopDocumentPDFWorker := context.WithCancel(context.Background())
|
||||
|
||||
wg.Go(
|
||||
func() {
|
||||
if err := documentPDFWorker.Run(documentPDFWorkerCtx); err != nil {
|
||||
@@ -655,6 +677,7 @@ func (impl *Implm) Run(
|
||||
)
|
||||
|
||||
accessReviewWorkerCtx, stopAccessReviewWorker := context.WithCancel(context.Background())
|
||||
|
||||
wg.Go(
|
||||
func() {
|
||||
if err := accessReviewService.Run(accessReviewWorkerCtx); err != nil {
|
||||
@@ -664,6 +687,7 @@ func (impl *Implm) Run(
|
||||
)
|
||||
|
||||
iamServiceCtx, stopIAMService := context.WithCancel(context.Background())
|
||||
|
||||
wg.Go(
|
||||
func() {
|
||||
if err := iamService.Run(iamServiceCtx); err != nil {
|
||||
@@ -673,6 +697,7 @@ func (impl *Implm) Run(
|
||||
)
|
||||
|
||||
esignServiceCtx, stopESignService := context.WithCancel(context.Background())
|
||||
|
||||
wg.Go(
|
||||
func() {
|
||||
if err := esignService.Run(esignServiceCtx, trustService.EmailPresenterConfigByOrganizationID); err != nil {
|
||||
@@ -683,6 +708,7 @@ func (impl *Implm) Run(
|
||||
|
||||
trackerPatternAnalysisWorker := cookiebanner.NewPatternAnalysisWorker(cookieBannerService, pgClient, l)
|
||||
trackerPatternAnalysisWorkerCtx, stopTrackerPatternAnalysisWorker := context.WithCancel(context.Background())
|
||||
|
||||
wg.Go(
|
||||
func() {
|
||||
if err := trackerPatternAnalysisWorker.Run(trackerPatternAnalysisWorkerCtx); err != nil {
|
||||
@@ -693,6 +719,7 @@ func (impl *Implm) Run(
|
||||
|
||||
trackerMappingWorker := cookiebanner.NewTrackerMappingWorker(pgClient, l, trackerMappingCfg)
|
||||
trackerMappingWorkerCtx, stopTrackerMappingWorker := context.WithCancel(context.Background())
|
||||
|
||||
wg.Go(
|
||||
func() {
|
||||
if err := trackerMappingWorker.Run(trackerMappingWorkerCtx); err != nil {
|
||||
@@ -703,6 +730,7 @@ func (impl *Implm) Run(
|
||||
|
||||
mailingListWorker := mailman.NewMailingListWorker(mailmanService, pgClient, l.Named("mailing-list-worker"))
|
||||
mailingListWorkerCtx, stopMailingListWorker := context.WithCancel(context.Background())
|
||||
|
||||
wg.Go(
|
||||
func() {
|
||||
if err := mailingListWorker.Run(mailingListWorkerCtx); err != nil {
|
||||
@@ -731,6 +759,7 @@ func (impl *Implm) Run(
|
||||
worker.WithMaxConcurrency(impl.cfg.EvidenceDescriber.MaxConcurrency),
|
||||
)
|
||||
evidenceDescriptionWorkerCtx, stopEvidenceDescriptionWorker := context.WithCancel(context.Background())
|
||||
|
||||
wg.Go(
|
||||
func() {
|
||||
if err := evidenceDescriptionWorker.Run(evidenceDescriptionWorkerCtx); err != nil {
|
||||
@@ -741,6 +770,7 @@ func (impl *Implm) Run(
|
||||
|
||||
trustCenterServerCtx, stopTrustCenterServer := context.WithCancel(context.Background())
|
||||
defer stopTrustCenterServer()
|
||||
|
||||
wg.Go(
|
||||
func() {
|
||||
if err := impl.runTrustCenterServer(
|
||||
@@ -811,6 +841,7 @@ func (impl *Implm) runApiServer(
|
||||
handler http.Handler,
|
||||
) error {
|
||||
tracer := tp.Tracer("go.probo.inc/probo/pkg/probod")
|
||||
|
||||
ctx, span := tracer.Start(ctx, "probod.runApiServer")
|
||||
defer span.End()
|
||||
|
||||
@@ -819,6 +850,7 @@ func (impl *Implm) runApiServer(
|
||||
span.RecordError(err)
|
||||
return fmt.Errorf("cannot build trusted proxy middleware: %w", err)
|
||||
}
|
||||
|
||||
handler = trustedProxyMiddleware(handler)
|
||||
|
||||
apiServer := httpserver.NewServer(
|
||||
@@ -853,14 +885,17 @@ func (impl *Implm) runApiServer(
|
||||
|
||||
l.Info("using proxy protocol", log.Any("trusted-proxies", impl.cfg.Api.ProxyProtocol.TrustedProxies))
|
||||
}
|
||||
|
||||
defer func() { _ = listener.Close() }()
|
||||
|
||||
serverErrCh := make(chan error, 1)
|
||||
|
||||
go func() {
|
||||
err := apiServer.Serve(listener)
|
||||
if err != nil && !errors.Is(err, http.ErrServerClosed) {
|
||||
serverErrCh <- fmt.Errorf("cannot server http request: %w", err)
|
||||
}
|
||||
|
||||
close(serverErrCh)
|
||||
}()
|
||||
|
||||
@@ -872,6 +907,7 @@ func (impl *Implm) runApiServer(
|
||||
if err != nil {
|
||||
span.RecordError(err)
|
||||
}
|
||||
|
||||
return err
|
||||
case <-ctx.Done():
|
||||
}
|
||||
@@ -888,6 +924,7 @@ func (impl *Implm) runApiServer(
|
||||
}
|
||||
|
||||
span.AddEvent("API server shutdown complete")
|
||||
|
||||
return ctx.Err()
|
||||
}
|
||||
|
||||
@@ -946,6 +983,7 @@ func (impl *Implm) runTrustCenterServer(
|
||||
encryptionKey cipher.EncryptionKey,
|
||||
) error {
|
||||
tracer := tp.Tracer("go.probo.inc/probo/pkg/probod")
|
||||
|
||||
ctx, span := tracer.Start(ctx, "probod.runTrustCenterServer")
|
||||
defer span.End()
|
||||
|
||||
@@ -968,6 +1006,7 @@ func (impl *Implm) runTrustCenterServer(
|
||||
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)
|
||||
@@ -1013,6 +1052,7 @@ func (impl *Implm) runTrustCenterServer(
|
||||
if err != nil {
|
||||
return fmt.Errorf("cannot listen on %q: %w", httpServer.Addr, err)
|
||||
}
|
||||
|
||||
defer func() { _ = listener.Close() }()
|
||||
|
||||
if len(impl.cfg.TrustCenter.ProxyProtocol.TrustedProxies) > 0 {
|
||||
@@ -1033,6 +1073,7 @@ func (impl *Implm) runTrustCenterServer(
|
||||
if err := httpServer.Serve(listener); err != nil && err != http.ErrServerClosed {
|
||||
return fmt.Errorf("cannot serve http requests: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
},
|
||||
)
|
||||
@@ -1075,10 +1116,12 @@ func (impl *Implm) runTrustCenterServer(
|
||||
if errors.As(err, &noSNIErr) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
if errors.Is(err, coredata.ErrResourceNotFound) {
|
||||
return nil, nil
|
||||
}
|
||||
}
|
||||
|
||||
return cert, err
|
||||
},
|
||||
MinVersion: tls.VersionTLS12,
|
||||
@@ -1101,6 +1144,7 @@ func (impl *Implm) runTrustCenterServer(
|
||||
if err != nil {
|
||||
return fmt.Errorf("cannot listen on %q: %w", httpsServer.Addr, err)
|
||||
}
|
||||
|
||||
defer func() { _ = listener.Close() }()
|
||||
|
||||
if len(impl.cfg.TrustCenter.ProxyProtocol.TrustedProxies) > 0 {
|
||||
|
||||
Reference in New Issue
Block a user