mirror of
https://github.com/netbirdio/netbird.git
synced 2026-07-30 17:56:11 -04:00
Read reverse-proxy service restrictions and target options in Postgres path
This commit is contained in:
@@ -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 {
|
||||
|
||||
@@ -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")
|
||||
})
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user