From 8d70cf1dc1c9a96ed01c3d3780773088cbf445f5 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Aur=C3=A9lien=20Sibiril?= <81782+aureliensibiril@users.noreply.github.com> Date: Fri, 10 Apr 2026 10:12:54 +0200 Subject: [PATCH] Stop leaking internal error in initiate handler MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The 500 response was wrapping the underlying error with %w, exposing internal details to the client. Log the full error, return a generic message. Signed-off-by: Aurélien Sibiril <81782+aureliensibiril@users.noreply.github.com> --- .../api/console/v1/connector_initiate.go | 52 +++++++++++++++++-- 1 file changed, 49 insertions(+), 3 deletions(-) diff --git a/pkg/server/api/console/v1/connector_initiate.go b/pkg/server/api/console/v1/connector_initiate.go index 7108b270c..1a5084458 100644 --- a/pkg/server/api/console/v1/connector_initiate.go +++ b/pkg/server/api/console/v1/connector_initiate.go @@ -29,6 +29,10 @@ import ( "go.probo.inc/probo/pkg/server/api/authn" ) +var ( + errInvalidReconnectConnector = errors.New("invalid reconnect connector") +) + func handleConnectorInitiate( logger *log.Logger, proboSvc *probo.Service, @@ -91,8 +95,12 @@ func handleConnectorInitiate( httpserver.RenderError(w, http.StatusBadRequest, fmt.Errorf("cannot reconnect: connector not found")) return } + if errors.Is(err, errInvalidReconnectConnector) { + httpserver.RenderError(w, http.StatusBadRequest, err) + return + } logger.ErrorCtx(r.Context(), "cannot look up existing connector", log.Error(err)) - httpserver.RenderError(w, http.StatusInternalServerError, fmt.Errorf("cannot look up existing connector: %w", err)) + httpserver.RenderError(w, http.StatusInternalServerError, fmt.Errorf("cannot look up existing connector")) return } @@ -131,9 +139,17 @@ func loadExistingConnector( if explicitID := r.URL.Query().Get("connector_id"); explicitID != "" { parsedID, err := gid.ParseGID(explicitID) if err != nil { - return nil, fmt.Errorf("cannot parse connector id: %w", err) + return nil, fmt.Errorf("%w: cannot parse connector id: %w", errInvalidReconnectConnector, err) } - return prb.Connectors.GetWithConnection(r.Context(), parsedID) + + found, err := prb.Connectors.GetWithConnection(r.Context(), parsedID) + if err != nil { + return nil, err + } + if err := validateReconnectConnector(found, organizationID, provider); err != nil { + return nil, err + } + return found, nil } found, err := prb.Connectors.GetByOrganizationIDAndProvider( @@ -146,3 +162,33 @@ func loadExistingConnector( } return found, err } + +func validateReconnectConnector( + c *coredata.Connector, + organizationID gid.GID, + provider string, +) error { + if c == nil { + return fmt.Errorf("%w: connector not found", errInvalidReconnectConnector) + } + + var connectorProvider coredata.ConnectorProvider + if err := connectorProvider.Scan(provider); err != nil { + return fmt.Errorf("%w: unsupported provider: %w", errInvalidReconnectConnector, err) + } + + if c.OrganizationID != organizationID { + return fmt.Errorf("%w: organization mismatch", errInvalidReconnectConnector) + } + if c.Provider != connectorProvider { + return fmt.Errorf("%w: provider mismatch", errInvalidReconnectConnector) + } + if c.Protocol != coredata.ConnectorProtocolOAuth2 { + return fmt.Errorf("%w: not an OAuth2 connector", errInvalidReconnectConnector) + } + if c.Connection == nil { + return fmt.Errorf("%w: connector has nil connection", errInvalidReconnectConnector) + } + + return nil +}