From 982794d46e9c09e99b12715b447971d979f7f274 Mon Sep 17 00:00:00 2001 From: Viktor Liu Date: Fri, 24 Jul 2026 15:53:29 +0200 Subject: [PATCH] Split getServices into scanService and getServiceTargets helpers --- management/server/store/sql_store.go | 386 ++++++++++++++------------- 1 file changed, 197 insertions(+), 189 deletions(-) diff --git a/management/server/store/sql_store.go b/management/server/store/sql_store.go index 3350725f3..c2125e75b 100644 --- a/management/server/store/sql_store.go +++ b/management/server/store/sql_store.go @@ -2266,125 +2266,12 @@ func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpserv private, access_groups FROM services WHERE account_id = $1` - const targetsQuery = `SELECT id, account_id, service_id, path, host, port, protocol, - target_id, target_type, enabled, proxy_protocol, - skip_tls_verify, request_timeout, session_idle_timeout, path_rewrite, custom_headers, - direct_upstream, middlewares, capture_max_request_bytes, capture_max_response_bytes, - capture_content_types, agent_network, disable_access_log - FROM targets WHERE service_id = ANY($1)` - serviceRows, err := s.pool.Query(ctx, serviceQuery, accountID) if err != nil { return nil, err } - services, err := pgx.CollectRows(serviceRows, func(row pgx.CollectableRow) (*rpservice.Service, error) { - var s rpservice.Service - var auth []byte - var restrictions []byte - var accessGroups []byte - var createdAt, certIssuedAt, lastRenewedAt sql.NullTime - var status, proxyCluster, sessionPrivateKey, sessionPublicKey sql.NullString - var mode, source, sourcePeer sql.NullString - var terminated, portAutoAssigned, private sql.NullBool - var listenPort sql.NullInt64 - err := row.Scan( - &s.ID, - &s.AccountID, - &s.Name, - &s.Domain, - &s.Enabled, - &auth, - &restrictions, - &createdAt, - &certIssuedAt, - &lastRenewedAt, - &status, - &proxyCluster, - &s.PassHostHeader, - &s.RewriteRedirects, - &sessionPrivateKey, - &sessionPublicKey, - &mode, - &listenPort, - &portAutoAssigned, - &source, - &sourcePeer, - &terminated, - &private, - &accessGroups, - ) - if err != nil { - return nil, err - } - - if auth != nil { - if err := json.Unmarshal(auth, &s.Auth); err != nil { - return nil, err - } - } - - if len(restrictions) > 0 { - if err := json.Unmarshal(restrictions, &s.Restrictions); err != nil { - return nil, fmt.Errorf("unmarshal restrictions: %w", err) - } - } - - if len(accessGroups) > 0 { - if err := json.Unmarshal(accessGroups, &s.AccessGroups); err != nil { - return nil, fmt.Errorf("unmarshal access_groups: %w", err) - } - } - - if private.Valid { - s.Private = private.Bool - } - - s.Meta = rpservice.Meta{} - if createdAt.Valid { - s.Meta.CreatedAt = createdAt.Time - } - if certIssuedAt.Valid { - t := certIssuedAt.Time - s.Meta.CertificateIssuedAt = &t - } - if lastRenewedAt.Valid { - t := lastRenewedAt.Time - s.Meta.LastRenewedAt = &t - } - if status.Valid { - s.Meta.Status = status.String - } - if proxyCluster.Valid { - s.ProxyCluster = proxyCluster.String - } - if sessionPrivateKey.Valid { - s.SessionPrivateKey = sessionPrivateKey.String - } - if sessionPublicKey.Valid { - s.SessionPublicKey = sessionPublicKey.String - } - if mode.Valid { - s.Mode = mode.String - } - if source.Valid { - s.Source = source.String - } - if sourcePeer.Valid { - s.SourcePeer = sourcePeer.String - } - if terminated.Valid { - s.Terminated = terminated.Bool - } - if portAutoAssigned.Valid { - s.PortAutoAssigned = portAutoAssigned.Bool - } - if listenPort.Valid { - s.ListenPort = uint16(listenPort.Int64) - } - s.Targets = []*rpservice.Target{} - return &s, nil - }) + services, err := pgx.CollectRows(serviceRows, scanService) if err != nil { return nil, err } @@ -2395,83 +2282,12 @@ func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpserv serviceIDs := make([]string, len(services)) serviceMap := make(map[string]*rpservice.Service) - for i, s := range services { - serviceIDs[i] = s.ID - serviceMap[s.ID] = s + for i, svc := range services { + serviceIDs[i] = svc.ID + serviceMap[svc.ID] = svc } - targetRows, err := s.pool.Query(ctx, targetsQuery, serviceIDs) - if err != nil { - return nil, err - } - - targets, err := pgx.CollectRows(targetRows, func(row pgx.CollectableRow) (*rpservice.Target, error) { - var t rpservice.Target - var path sql.NullString - var pathRewrite sql.NullString - var proxyProtocol, skipTLSVerify, directUpstream, agentNetwork, disableAccessLog sql.NullBool - var requestTimeout, sessionIdleTimeout, captureMaxRequestBytes, captureMaxResponseBytes sql.NullInt64 - var customHeaders, middlewares, captureContentTypes []byte - err := row.Scan( - &t.ID, - &t.AccountID, - &t.ServiceID, - &path, - &t.Host, - &t.Port, - &t.Protocol, - &t.TargetId, - &t.TargetType, - &t.Enabled, - &proxyProtocol, - &skipTLSVerify, - &requestTimeout, - &sessionIdleTimeout, - &pathRewrite, - &customHeaders, - &directUpstream, - &middlewares, - &captureMaxRequestBytes, - &captureMaxResponseBytes, - &captureContentTypes, - &agentNetwork, - &disableAccessLog, - ) - if err != nil { - return nil, err - } - if path.Valid { - t.Path = &path.String - } - - t.ProxyProtocol = proxyProtocol.Bool - t.Options.SkipTLSVerify = skipTLSVerify.Bool - t.Options.RequestTimeout = time.Duration(requestTimeout.Int64) - t.Options.SessionIdleTimeout = time.Duration(sessionIdleTimeout.Int64) - t.Options.PathRewrite = rpservice.PathRewriteMode(pathRewrite.String) - t.Options.DirectUpstream = directUpstream.Bool - t.Options.CaptureMaxRequestBytes = captureMaxRequestBytes.Int64 - t.Options.CaptureMaxResponseBytes = captureMaxResponseBytes.Int64 - t.Options.AgentNetwork = agentNetwork.Bool - t.Options.DisableAccessLog = disableAccessLog.Bool - - if len(customHeaders) > 0 { - if err := json.Unmarshal(customHeaders, &t.Options.CustomHeaders); err != nil { - return nil, fmt.Errorf("unmarshal custom_headers: %w", err) - } - } - if len(middlewares) > 0 { - if err := json.Unmarshal(middlewares, &t.Options.Middlewares); err != nil { - return nil, fmt.Errorf("unmarshal middlewares: %w", err) - } - } - if len(captureContentTypes) > 0 { - if err := json.Unmarshal(captureContentTypes, &t.Options.CaptureContentTypes); err != nil { - return nil, fmt.Errorf("unmarshal capture_content_types: %w", err) - } - } - return &t, nil - }) + targets, err := s.getServiceTargets(ctx, serviceIDs) if err != nil { return nil, err } @@ -2485,6 +2301,198 @@ func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpserv return services, nil } +func scanService(row pgx.CollectableRow) (*rpservice.Service, error) { + var s rpservice.Service + var auth []byte + var restrictions []byte + var accessGroups []byte + var createdAt, certIssuedAt, lastRenewedAt sql.NullTime + var status, proxyCluster, sessionPrivateKey, sessionPublicKey sql.NullString + var mode, source, sourcePeer sql.NullString + var terminated, portAutoAssigned, private sql.NullBool + var listenPort sql.NullInt64 + err := row.Scan( + &s.ID, + &s.AccountID, + &s.Name, + &s.Domain, + &s.Enabled, + &auth, + &restrictions, + &createdAt, + &certIssuedAt, + &lastRenewedAt, + &status, + &proxyCluster, + &s.PassHostHeader, + &s.RewriteRedirects, + &sessionPrivateKey, + &sessionPublicKey, + &mode, + &listenPort, + &portAutoAssigned, + &source, + &sourcePeer, + &terminated, + &private, + &accessGroups, + ) + if err != nil { + return nil, err + } + + if auth != nil { + if err := json.Unmarshal(auth, &s.Auth); err != nil { + return nil, err + } + } + + if len(restrictions) > 0 { + if err := json.Unmarshal(restrictions, &s.Restrictions); err != nil { + return nil, fmt.Errorf("unmarshal restrictions: %w", err) + } + } + + if len(accessGroups) > 0 { + if err := json.Unmarshal(accessGroups, &s.AccessGroups); err != nil { + return nil, fmt.Errorf("unmarshal access_groups: %w", err) + } + } + + if private.Valid { + s.Private = private.Bool + } + + s.Meta = rpservice.Meta{} + if createdAt.Valid { + s.Meta.CreatedAt = createdAt.Time + } + if certIssuedAt.Valid { + t := certIssuedAt.Time + s.Meta.CertificateIssuedAt = &t + } + if lastRenewedAt.Valid { + t := lastRenewedAt.Time + s.Meta.LastRenewedAt = &t + } + if status.Valid { + s.Meta.Status = status.String + } + if proxyCluster.Valid { + s.ProxyCluster = proxyCluster.String + } + if sessionPrivateKey.Valid { + s.SessionPrivateKey = sessionPrivateKey.String + } + if sessionPublicKey.Valid { + s.SessionPublicKey = sessionPublicKey.String + } + if mode.Valid { + s.Mode = mode.String + } + if source.Valid { + s.Source = source.String + } + if sourcePeer.Valid { + s.SourcePeer = sourcePeer.String + } + if terminated.Valid { + s.Terminated = terminated.Bool + } + if portAutoAssigned.Valid { + s.PortAutoAssigned = portAutoAssigned.Bool + } + if listenPort.Valid { + s.ListenPort = uint16(listenPort.Int64) + } + s.Targets = []*rpservice.Target{} + return &s, nil +} + +func (s *SqlStore) getServiceTargets(ctx context.Context, serviceIDs []string) ([]*rpservice.Target, error) { + const targetsQuery = `SELECT id, account_id, service_id, path, host, port, protocol, + target_id, target_type, enabled, proxy_protocol, + skip_tls_verify, request_timeout, session_idle_timeout, path_rewrite, custom_headers, + direct_upstream, middlewares, capture_max_request_bytes, capture_max_response_bytes, + capture_content_types, agent_network, disable_access_log + FROM targets WHERE service_id = ANY($1)` + + rows, err := s.pool.Query(ctx, targetsQuery, serviceIDs) + if err != nil { + return nil, err + } + + return pgx.CollectRows(rows, scanTarget) +} + +func scanTarget(row pgx.CollectableRow) (*rpservice.Target, error) { + var t rpservice.Target + var path sql.NullString + var pathRewrite sql.NullString + var proxyProtocol, skipTLSVerify, directUpstream, agentNetwork, disableAccessLog sql.NullBool + var requestTimeout, sessionIdleTimeout, captureMaxRequestBytes, captureMaxResponseBytes sql.NullInt64 + var customHeaders, middlewares, captureContentTypes []byte + err := row.Scan( + &t.ID, + &t.AccountID, + &t.ServiceID, + &path, + &t.Host, + &t.Port, + &t.Protocol, + &t.TargetId, + &t.TargetType, + &t.Enabled, + &proxyProtocol, + &skipTLSVerify, + &requestTimeout, + &sessionIdleTimeout, + &pathRewrite, + &customHeaders, + &directUpstream, + &middlewares, + &captureMaxRequestBytes, + &captureMaxResponseBytes, + &captureContentTypes, + &agentNetwork, + &disableAccessLog, + ) + if err != nil { + return nil, err + } + if path.Valid { + t.Path = &path.String + } + + t.ProxyProtocol = proxyProtocol.Bool + t.Options.SkipTLSVerify = skipTLSVerify.Bool + t.Options.RequestTimeout = time.Duration(requestTimeout.Int64) + t.Options.SessionIdleTimeout = time.Duration(sessionIdleTimeout.Int64) + t.Options.PathRewrite = rpservice.PathRewriteMode(pathRewrite.String) + t.Options.DirectUpstream = directUpstream.Bool + t.Options.CaptureMaxRequestBytes = captureMaxRequestBytes.Int64 + t.Options.CaptureMaxResponseBytes = captureMaxResponseBytes.Int64 + t.Options.AgentNetwork = agentNetwork.Bool + t.Options.DisableAccessLog = disableAccessLog.Bool + + if len(customHeaders) > 0 { + if err := json.Unmarshal(customHeaders, &t.Options.CustomHeaders); err != nil { + return nil, fmt.Errorf("unmarshal custom_headers: %w", err) + } + } + if len(middlewares) > 0 { + if err := json.Unmarshal(middlewares, &t.Options.Middlewares); err != nil { + return nil, fmt.Errorf("unmarshal middlewares: %w", err) + } + } + if len(captureContentTypes) > 0 { + if err := json.Unmarshal(captureContentTypes, &t.Options.CaptureContentTypes); err != nil { + return nil, fmt.Errorf("unmarshal capture_content_types: %w", err) + } + } + return &t, nil +} + func (s *SqlStore) getNetworks(ctx context.Context, accountID string) ([]*networkTypes.Network, error) { const query = `SELECT id, account_id, public_id, name, description FROM networks WHERE account_id = $1` rows, err := s.pool.Query(ctx, query, accountID)