diff --git a/management/server/store/sql_store.go b/management/server/store/sql_store.go index 3ad870ad3..3350725f3 100644 --- a/management/server/store/sql_store.go +++ b/management/server/store/sql_store.go @@ -2259,15 +2259,18 @@ func (s *SqlStore) getPostureChecks(ctx context.Context, accountID string) ([]*p } func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpservice.Service, error) { - const serviceQuery = `SELECT id, account_id, name, domain, enabled, auth, - meta_created_at, meta_certificate_issued_at, meta_status, proxy_cluster, + const serviceQuery = `SELECT id, account_id, name, domain, enabled, auth, restrictions, + meta_created_at, meta_certificate_issued_at, meta_last_renewed_at, meta_status, proxy_cluster, pass_host_header, rewrite_redirects, session_private_key, session_public_key, mode, listen_port, port_auto_assigned, source, source_peer, terminated, 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 + 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) @@ -2278,8 +2281,9 @@ func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpserv 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 sql.NullTime + 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 @@ -2291,8 +2295,10 @@ func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpserv &s.Domain, &s.Enabled, &auth, + &restrictions, &createdAt, &certIssuedAt, + &lastRenewedAt, &status, &proxyCluster, &s.PassHostHeader, @@ -2318,6 +2324,12 @@ func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpserv } } + 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) @@ -2336,6 +2348,10 @@ func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpserv 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 } @@ -2392,6 +2408,10 @@ func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpserv 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, @@ -2403,6 +2423,19 @@ func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpserv &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 @@ -2410,6 +2443,33 @@ func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpserv 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 }) if err != nil { diff --git a/management/server/store/sql_store_service_test.go b/management/server/store/sql_store_service_test.go index 34999da4b..0e14fbdab 100644 --- a/management/server/store/sql_store_service_test.go +++ b/management/server/store/sql_store_service_test.go @@ -5,6 +5,7 @@ import ( "os" "runtime" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -44,3 +45,91 @@ func TestSqlStore_GetAccount_PrivateServiceRoundtrip(t *testing.T) { assert.Equal(t, []string{"grp-admins", "grp-ops"}, got.AccessGroups) }) } + +// TestSqlStore_GetAccount_ServiceTargetOptionsRoundtrip guards the Postgres pgx +// read path (getServices) against silently dropping columns present on the gorm +// model. Before the fix these fields loaded correctly on SQLite but came back +// zero-valued on Postgres because the hand-written SELECT and scan omitted them. +func TestSqlStore_GetAccount_ServiceTargetOptionsRoundtrip(t *testing.T) { + if os.Getenv("CI") == "true" && (runtime.GOOS == "darwin" || runtime.GOOS == "windows") { + t.Skip("skip CI tests on darwin and windows") + } + + runTestForAllEngines(t, "", func(t *testing.T, store Store) { + ctx := context.Background() + account := newAccountWithId(ctx, "account_svc_opts", "testuser", "") + require.NoError(t, store.SaveAccount(ctx, account)) + + renewedAt := time.Now().UTC().Truncate(time.Second) + targetPath := "/api" + svc := &rpservice.Service{ + ID: "svc-opts", + AccountID: account.Id, + Name: "opts-svc", + Domain: "opts.example", + Enabled: true, + Mode: rpservice.ModeHTTP, + Restrictions: rpservice.AccessRestrictions{ + AllowedCIDRs: []string{"10.0.0.0/8"}, + BlockedCountries: []string{"XX"}, + CrowdSecMode: "block", + }, + Meta: rpservice.Meta{ + LastRenewedAt: &renewedAt, + }, + Targets: []*rpservice.Target{ + { + AccountID: account.Id, + ServiceID: "svc-opts", + Path: &targetPath, + Host: "backend.internal", + Port: 8080, + Protocol: "http", + TargetId: "tgt-1", + Enabled: true, + ProxyProtocol: true, + Options: rpservice.TargetOptions{ + SkipTLSVerify: true, + RequestTimeout: 30 * time.Second, + SessionIdleTimeout: 5 * time.Minute, + PathRewrite: rpservice.PathRewritePreserve, + CustomHeaders: map[string]string{"X-Foo": "bar"}, + DirectUpstream: true, + CaptureMaxRequestBytes: 1024, + CaptureMaxResponseBytes: 2048, + CaptureContentTypes: []string{"application/json"}, + AgentNetwork: true, + DisableAccessLog: true, + }, + }, + }, + } + require.NoError(t, store.CreateService(ctx, svc)) + + loaded, err := store.GetAccount(ctx, account.Id) + require.NoError(t, err) + require.Len(t, loaded.Services, 1) + + got := loaded.Services[0] + assert.Equal(t, []string{"10.0.0.0/8"}, got.Restrictions.AllowedCIDRs, "restrictions allowed CIDRs") + assert.Equal(t, []string{"XX"}, got.Restrictions.BlockedCountries, "restrictions blocked countries") + assert.Equal(t, "block", got.Restrictions.CrowdSecMode, "restrictions crowdsec mode") + require.NotNil(t, got.Meta.LastRenewedAt, "meta last renewed at") + assert.WithinDuration(t, renewedAt, *got.Meta.LastRenewedAt, time.Second, "meta last renewed at") + + require.Len(t, got.Targets, 1) + tg := got.Targets[0] + assert.True(t, tg.ProxyProtocol, "target proxy protocol") + assert.True(t, tg.Options.SkipTLSVerify, "options skip TLS verify") + assert.Equal(t, 30*time.Second, tg.Options.RequestTimeout, "options request timeout") + assert.Equal(t, 5*time.Minute, tg.Options.SessionIdleTimeout, "options session idle timeout") + assert.Equal(t, rpservice.PathRewritePreserve, tg.Options.PathRewrite, "options path rewrite") + assert.Equal(t, map[string]string{"X-Foo": "bar"}, tg.Options.CustomHeaders, "options custom headers") + assert.True(t, tg.Options.DirectUpstream, "options direct upstream") + assert.Equal(t, int64(1024), tg.Options.CaptureMaxRequestBytes, "options capture max request bytes") + assert.Equal(t, int64(2048), tg.Options.CaptureMaxResponseBytes, "options capture max response bytes") + assert.Equal(t, []string{"application/json"}, tg.Options.CaptureContentTypes, "options capture content types") + assert.True(t, tg.Options.AgentNetwork, "options agent network") + assert.True(t, tg.Options.DisableAccessLog, "options disable access log") + }) +}