mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: Add high availability for multiple replicas (#4555)
* feat: HA tailnet coordinator * fixup! feat: HA tailnet coordinator * fixup! feat: HA tailnet coordinator * remove printlns * close all connections on coordinator * impelement high availability feature * fixup! impelement high availability feature * fixup! impelement high availability feature * fixup! impelement high availability feature * fixup! impelement high availability feature * Add replicas * Add DERP meshing to arbitrary addresses * Move packages to highavailability folder * Move coordinator to high availability package * Add flags for HA * Rename to replicasync * Denest packages for replicas * Add test for multiple replicas * Fix coordination test * Add HA to the helm chart * Rename function pointer * Add warnings for HA * Add the ability to block endpoints * Add flag to disable P2P connections * Wow, I made the tests pass * Add replicas endpoint * Ensure close kills replica * Update sql * Add database latency to high availability * Pipe TLS to DERP mesh * Fix DERP mesh with TLS * Add tests for TLS * Fix replica sync TLS * Fix RootCA for replica meshing * Remove ID from replicasync * Fix getting certificates for meshing * Remove excessive locking * Fix linting * Store mesh key in the database * Fix replica key for tests * Fix types gen * Fix unlocking unlocked * Fix race in tests * Update enterprise/derpmesh/derpmesh.go Co-authored-by: Colin Adler <colin1adler@gmail.com> * Rename to syncReplicas * Reuse http client * Delete old replicas on a CRON * Fix race condition in connection tests * Fix linting * Fix nil type * Move pubsub to in-memory for twenty test * Add comment for configuration tweaking * Fix leak with transport * Fix close leak in derpmesh * Fix race when creating server * Remove handler update * Skip test on Windows * Fix DERP mesh test * Wrap HTTP handler replacement in mutex * Fix error message for relay * Fix API handler for normal tests * Fix speedtest * Fix replica resend * Fix derpmesh send * Ping async * Increase wait time of template version jobd * Fix race when closing replica sync * Add name to client * Log the derpmap being used * Don't connect if DERP is empty * Improve agent coordinator logging * Fix lock in coordinator * Fix relay addr * Fix race when updating durations * Fix client publish race * Run pubsub loop in a queue * Store agent nodes in order * Fix coordinator locking * Check for closed pipe Co-authored-by: Colin Adler <colin1adler@gmail.com>
This commit is contained in:
co-authored by
Colin Adler
parent
dc3519e973
commit
2ba4a62a0d
@@ -57,7 +57,7 @@ func TestFeaturesList(t *testing.T) {
|
||||
var entitlements codersdk.Entitlements
|
||||
err := json.Unmarshal(buf.Bytes(), &entitlements)
|
||||
require.NoError(t, err, "unmarshal JSON output")
|
||||
assert.Len(t, entitlements.Features, 6)
|
||||
assert.Len(t, entitlements.Features, 7)
|
||||
assert.Empty(t, entitlements.Warnings)
|
||||
assert.Equal(t, codersdk.EntitlementNotEntitled,
|
||||
entitlements.Features[codersdk.FeatureUserLimit].Entitlement)
|
||||
@@ -71,6 +71,8 @@ func TestFeaturesList(t *testing.T) {
|
||||
entitlements.Features[codersdk.FeatureTemplateRBAC].Entitlement)
|
||||
assert.Equal(t, codersdk.EntitlementNotEntitled,
|
||||
entitlements.Features[codersdk.FeatureSCIM].Entitlement)
|
||||
assert.Equal(t, codersdk.EntitlementNotEntitled,
|
||||
entitlements.Features[codersdk.FeatureHighAvailability].Entitlement)
|
||||
assert.False(t, entitlements.HasLicense)
|
||||
assert.False(t, entitlements.Experimental)
|
||||
})
|
||||
|
||||
+45
-10
@@ -2,11 +2,20 @@ package cli
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"io"
|
||||
"net/url"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
"golang.org/x/xerrors"
|
||||
"tailscale.com/derp"
|
||||
"tailscale.com/types/key"
|
||||
|
||||
"github.com/coder/coder/cli/deployment"
|
||||
"github.com/coder/coder/cryptorand"
|
||||
"github.com/coder/coder/enterprise/coderd"
|
||||
"github.com/coder/coder/tailnet"
|
||||
|
||||
agpl "github.com/coder/coder/cli"
|
||||
agplcoderd "github.com/coder/coder/coderd"
|
||||
@@ -14,23 +23,49 @@ import (
|
||||
|
||||
func server() *cobra.Command {
|
||||
dflags := deployment.Flags()
|
||||
cmd := agpl.Server(dflags, func(ctx context.Context, options *agplcoderd.Options) (*agplcoderd.API, error) {
|
||||
cmd := agpl.Server(dflags, func(ctx context.Context, options *agplcoderd.Options) (*agplcoderd.API, io.Closer, error) {
|
||||
if dflags.DerpServerRelayAddress.Value != "" {
|
||||
_, err := url.Parse(dflags.DerpServerRelayAddress.Value)
|
||||
if err != nil {
|
||||
return nil, nil, xerrors.Errorf("derp-server-relay-address must be a valid HTTP URL: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
options.DERPServer = derp.NewServer(key.NewNode(), tailnet.Logger(options.Logger.Named("derp")))
|
||||
meshKey, err := options.Database.GetDERPMeshKey(ctx)
|
||||
if err != nil {
|
||||
if !errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, nil, xerrors.Errorf("get mesh key: %w", err)
|
||||
}
|
||||
meshKey, err = cryptorand.String(32)
|
||||
if err != nil {
|
||||
return nil, nil, xerrors.Errorf("generate mesh key: %w", err)
|
||||
}
|
||||
err = options.Database.InsertDERPMeshKey(ctx, meshKey)
|
||||
if err != nil {
|
||||
return nil, nil, xerrors.Errorf("insert mesh key: %w", err)
|
||||
}
|
||||
}
|
||||
options.DERPServer.SetMeshKey(meshKey)
|
||||
|
||||
o := &coderd.Options{
|
||||
AuditLogging: dflags.AuditLogging.Value,
|
||||
BrowserOnly: dflags.BrowserOnly.Value,
|
||||
SCIMAPIKey: []byte(dflags.SCIMAuthHeader.Value),
|
||||
UserWorkspaceQuota: dflags.UserWorkspaceQuota.Value,
|
||||
RBACEnabled: true,
|
||||
Options: options,
|
||||
AuditLogging: dflags.AuditLogging.Value,
|
||||
BrowserOnly: dflags.BrowserOnly.Value,
|
||||
SCIMAPIKey: []byte(dflags.SCIMAuthHeader.Value),
|
||||
UserWorkspaceQuota: dflags.UserWorkspaceQuota.Value,
|
||||
RBAC: true,
|
||||
DERPServerRelayAddress: dflags.DerpServerRelayAddress.Value,
|
||||
DERPServerRegionID: dflags.DerpServerRegionID.Value,
|
||||
|
||||
Options: options,
|
||||
}
|
||||
api, err := coderd.New(ctx, o)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, nil, err
|
||||
}
|
||||
return api.AGPL, nil
|
||||
return api.AGPL, api, nil
|
||||
})
|
||||
|
||||
deployment.AttachFlags(cmd.Flags(), dflags, true)
|
||||
|
||||
return cmd
|
||||
}
|
||||
|
||||
@@ -28,7 +28,7 @@ func TestCheckACLPermissions(t *testing.T) {
|
||||
// Create adminClient, member, and org adminClient
|
||||
adminUser := coderdtest.CreateFirstUser(t, adminClient)
|
||||
_ = coderdenttest.AddLicense(t, adminClient, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
memberClient := coderdtest.CreateAnotherUser(t, adminClient, adminUser.OrganizationID)
|
||||
|
||||
+104
-8
@@ -3,6 +3,8 @@ package coderd
|
||||
import (
|
||||
"context"
|
||||
"crypto/ed25519"
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -23,6 +25,10 @@ import (
|
||||
"github.com/coder/coder/enterprise/audit"
|
||||
"github.com/coder/coder/enterprise/audit/backends"
|
||||
"github.com/coder/coder/enterprise/coderd/license"
|
||||
"github.com/coder/coder/enterprise/derpmesh"
|
||||
"github.com/coder/coder/enterprise/replicasync"
|
||||
"github.com/coder/coder/enterprise/tailnet"
|
||||
agpltailnet "github.com/coder/coder/tailnet"
|
||||
)
|
||||
|
||||
// New constructs an Enterprise coderd API instance.
|
||||
@@ -47,6 +53,7 @@ func New(ctx context.Context, options *Options) (*API, error) {
|
||||
Options: options,
|
||||
cancelEntitlementsLoop: cancelFunc,
|
||||
}
|
||||
|
||||
oauthConfigs := &httpmw.OAuth2Configs{
|
||||
Github: options.GithubOAuth2Config,
|
||||
OIDC: options.OIDCConfig,
|
||||
@@ -59,6 +66,10 @@ func New(ctx context.Context, options *Options) (*API, error) {
|
||||
|
||||
api.AGPL.APIHandler.Group(func(r chi.Router) {
|
||||
r.Get("/entitlements", api.serveEntitlements)
|
||||
r.Route("/replicas", func(r chi.Router) {
|
||||
r.Use(apiKeyMiddleware)
|
||||
r.Get("/", api.replicas)
|
||||
})
|
||||
r.Route("/licenses", func(r chi.Router) {
|
||||
r.Use(apiKeyMiddleware)
|
||||
r.Post("/", api.postLicense)
|
||||
@@ -117,7 +128,40 @@ func New(ctx context.Context, options *Options) (*API, error) {
|
||||
})
|
||||
}
|
||||
|
||||
err := api.updateEntitlements(ctx)
|
||||
meshRootCA := x509.NewCertPool()
|
||||
for _, certificate := range options.TLSCertificates {
|
||||
for _, certificatePart := range certificate.Certificate {
|
||||
certificate, err := x509.ParseCertificate(certificatePart)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("parse certificate %s: %w", certificate.Subject.CommonName, err)
|
||||
}
|
||||
meshRootCA.AddCert(certificate)
|
||||
}
|
||||
}
|
||||
// This TLS configuration spoofs access from the access URL hostname
|
||||
// assuming that the certificates provided will cover that hostname.
|
||||
//
|
||||
// Replica sync and DERP meshing require accessing replicas via their
|
||||
// internal IP addresses, and if TLS is configured we use the same
|
||||
// certificates.
|
||||
meshTLSConfig := &tls.Config{
|
||||
MinVersion: tls.VersionTLS12,
|
||||
Certificates: options.TLSCertificates,
|
||||
RootCAs: meshRootCA,
|
||||
ServerName: options.AccessURL.Hostname(),
|
||||
}
|
||||
var err error
|
||||
api.replicaManager, err = replicasync.New(ctx, options.Logger, options.Database, options.Pubsub, &replicasync.Options{
|
||||
RelayAddress: options.DERPServerRelayAddress,
|
||||
RegionID: int32(options.DERPServerRegionID),
|
||||
TLSConfig: meshTLSConfig,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("initialize replica: %w", err)
|
||||
}
|
||||
api.derpMesh = derpmesh.New(options.Logger.Named("derpmesh"), api.DERPServer, meshTLSConfig)
|
||||
|
||||
err = api.updateEntitlements(ctx)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("update entitlements: %w", err)
|
||||
}
|
||||
@@ -129,13 +173,17 @@ func New(ctx context.Context, options *Options) (*API, error) {
|
||||
type Options struct {
|
||||
*coderd.Options
|
||||
|
||||
RBACEnabled bool
|
||||
RBAC bool
|
||||
AuditLogging bool
|
||||
// Whether to block non-browser connections.
|
||||
BrowserOnly bool
|
||||
SCIMAPIKey []byte
|
||||
UserWorkspaceQuota int
|
||||
|
||||
// Used for high availability.
|
||||
DERPServerRelayAddress string
|
||||
DERPServerRegionID int
|
||||
|
||||
EntitlementsUpdateInterval time.Duration
|
||||
Keys map[string]ed25519.PublicKey
|
||||
}
|
||||
@@ -144,6 +192,11 @@ type API struct {
|
||||
AGPL *coderd.API
|
||||
*Options
|
||||
|
||||
// Detects multiple Coder replicas running at the same time.
|
||||
replicaManager *replicasync.Manager
|
||||
// Meshes DERP connections from multiple replicas.
|
||||
derpMesh *derpmesh.Mesh
|
||||
|
||||
cancelEntitlementsLoop func()
|
||||
entitlementsMu sync.RWMutex
|
||||
entitlements codersdk.Entitlements
|
||||
@@ -151,6 +204,8 @@ type API struct {
|
||||
|
||||
func (api *API) Close() error {
|
||||
api.cancelEntitlementsLoop()
|
||||
_ = api.replicaManager.Close()
|
||||
_ = api.derpMesh.Close()
|
||||
return api.AGPL.Close()
|
||||
}
|
||||
|
||||
@@ -158,12 +213,13 @@ func (api *API) updateEntitlements(ctx context.Context) error {
|
||||
api.entitlementsMu.Lock()
|
||||
defer api.entitlementsMu.Unlock()
|
||||
|
||||
entitlements, err := license.Entitlements(ctx, api.Database, api.Logger, api.Keys, map[string]bool{
|
||||
codersdk.FeatureAuditLog: api.AuditLogging,
|
||||
codersdk.FeatureBrowserOnly: api.BrowserOnly,
|
||||
codersdk.FeatureSCIM: len(api.SCIMAPIKey) != 0,
|
||||
codersdk.FeatureWorkspaceQuota: api.UserWorkspaceQuota != 0,
|
||||
codersdk.FeatureTemplateRBAC: api.RBACEnabled,
|
||||
entitlements, err := license.Entitlements(ctx, api.Database, api.Logger, len(api.replicaManager.All()), api.Keys, map[string]bool{
|
||||
codersdk.FeatureAuditLog: api.AuditLogging,
|
||||
codersdk.FeatureBrowserOnly: api.BrowserOnly,
|
||||
codersdk.FeatureSCIM: len(api.SCIMAPIKey) != 0,
|
||||
codersdk.FeatureWorkspaceQuota: api.UserWorkspaceQuota != 0,
|
||||
codersdk.FeatureHighAvailability: api.DERPServerRelayAddress != "",
|
||||
codersdk.FeatureTemplateRBAC: api.RBAC,
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -209,6 +265,46 @@ func (api *API) updateEntitlements(ctx context.Context) error {
|
||||
api.AGPL.WorkspaceQuotaEnforcer.Store(&enforcer)
|
||||
}
|
||||
|
||||
if changed, enabled := featureChanged(codersdk.FeatureHighAvailability); changed {
|
||||
coordinator := agpltailnet.NewCoordinator()
|
||||
if enabled {
|
||||
haCoordinator, err := tailnet.NewCoordinator(api.Logger, api.Pubsub)
|
||||
if err != nil {
|
||||
api.Logger.Error(ctx, "unable to set up high availability coordinator", slog.Error(err))
|
||||
// If we try to setup the HA coordinator and it fails, nothing
|
||||
// is actually changing.
|
||||
changed = false
|
||||
} else {
|
||||
coordinator = haCoordinator
|
||||
}
|
||||
|
||||
api.replicaManager.SetCallback(func() {
|
||||
addresses := make([]string, 0)
|
||||
for _, replica := range api.replicaManager.Regional() {
|
||||
addresses = append(addresses, replica.RelayAddress)
|
||||
}
|
||||
api.derpMesh.SetAddresses(addresses, false)
|
||||
_ = api.updateEntitlements(ctx)
|
||||
})
|
||||
} else {
|
||||
api.derpMesh.SetAddresses([]string{}, false)
|
||||
api.replicaManager.SetCallback(func() {
|
||||
// If the amount of replicas change, so should our entitlements.
|
||||
// This is to display a warning in the UI if the user is unlicensed.
|
||||
_ = api.updateEntitlements(ctx)
|
||||
})
|
||||
}
|
||||
|
||||
// Recheck changed in case the HA coordinator failed to set up.
|
||||
if changed {
|
||||
oldCoordinator := *api.AGPL.TailnetCoordinator.Swap(&coordinator)
|
||||
err := oldCoordinator.Close()
|
||||
if err != nil {
|
||||
api.Logger.Error(ctx, "close old tailnet coordinator", slog.Error(err))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
api.entitlements = entitlements
|
||||
|
||||
return nil
|
||||
|
||||
@@ -41,9 +41,9 @@ func TestEntitlements(t *testing.T) {
|
||||
})
|
||||
_ = coderdtest.CreateFirstUser(t, client)
|
||||
coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
UserLimit: 100,
|
||||
AuditLog: true,
|
||||
TemplateRBACEnabled: true,
|
||||
UserLimit: 100,
|
||||
AuditLog: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
res, err := client.Entitlements(context.Background())
|
||||
require.NoError(t, err)
|
||||
@@ -85,7 +85,7 @@ func TestEntitlements(t *testing.T) {
|
||||
assert.False(t, res.HasLicense)
|
||||
al = res.Features[codersdk.FeatureAuditLog]
|
||||
assert.Equal(t, codersdk.EntitlementNotEntitled, al.Entitlement)
|
||||
assert.True(t, al.Enabled)
|
||||
assert.False(t, al.Enabled)
|
||||
})
|
||||
t.Run("Pubsub", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -4,7 +4,9 @@ import (
|
||||
"context"
|
||||
"crypto/ed25519"
|
||||
"crypto/rand"
|
||||
"crypto/tls"
|
||||
"io"
|
||||
"net/http"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -60,19 +62,21 @@ func NewWithAPI(t *testing.T, options *Options) (*codersdk.Client, io.Closer, *c
|
||||
if options.Options == nil {
|
||||
options.Options = &coderdtest.Options{}
|
||||
}
|
||||
srv, cancelFunc, oop := coderdtest.NewOptions(t, options.Options)
|
||||
setHandler, cancelFunc, oop := coderdtest.NewOptions(t, options.Options)
|
||||
coderAPI, err := coderd.New(context.Background(), &coderd.Options{
|
||||
RBACEnabled: true,
|
||||
RBAC: true,
|
||||
AuditLogging: options.AuditLogging,
|
||||
BrowserOnly: options.BrowserOnly,
|
||||
SCIMAPIKey: options.SCIMAPIKey,
|
||||
DERPServerRelayAddress: oop.AccessURL.String(),
|
||||
DERPServerRegionID: oop.DERPMap.RegionIDs()[0],
|
||||
UserWorkspaceQuota: options.UserWorkspaceQuota,
|
||||
Options: oop,
|
||||
EntitlementsUpdateInterval: options.EntitlementsUpdateInterval,
|
||||
Keys: Keys,
|
||||
})
|
||||
assert.NoError(t, err)
|
||||
srv.Config.Handler = coderAPI.AGPL.RootHandler
|
||||
setHandler(coderAPI.AGPL.RootHandler)
|
||||
var provisionerCloser io.Closer = nopcloser{}
|
||||
if options.IncludeProvisionerDaemon {
|
||||
provisionerCloser = coderdtest.NewProvisionerDaemon(t, coderAPI.AGPL)
|
||||
@@ -83,22 +87,32 @@ func NewWithAPI(t *testing.T, options *Options) (*codersdk.Client, io.Closer, *c
|
||||
_ = provisionerCloser.Close()
|
||||
_ = coderAPI.Close()
|
||||
})
|
||||
return codersdk.New(coderAPI.AccessURL), provisionerCloser, coderAPI
|
||||
client := codersdk.New(coderAPI.AccessURL)
|
||||
client.HTTPClient = &http.Client{
|
||||
Transport: &http.Transport{
|
||||
TLSClientConfig: &tls.Config{
|
||||
//nolint:gosec
|
||||
InsecureSkipVerify: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
return client, provisionerCloser, coderAPI
|
||||
}
|
||||
|
||||
type LicenseOptions struct {
|
||||
AccountType string
|
||||
AccountID string
|
||||
Trial bool
|
||||
AllFeatures bool
|
||||
GraceAt time.Time
|
||||
ExpiresAt time.Time
|
||||
UserLimit int64
|
||||
AuditLog bool
|
||||
BrowserOnly bool
|
||||
SCIM bool
|
||||
WorkspaceQuota bool
|
||||
TemplateRBACEnabled bool
|
||||
AccountType string
|
||||
AccountID string
|
||||
Trial bool
|
||||
AllFeatures bool
|
||||
GraceAt time.Time
|
||||
ExpiresAt time.Time
|
||||
UserLimit int64
|
||||
AuditLog bool
|
||||
BrowserOnly bool
|
||||
SCIM bool
|
||||
WorkspaceQuota bool
|
||||
TemplateRBAC bool
|
||||
HighAvailability bool
|
||||
}
|
||||
|
||||
// AddLicense generates a new license with the options provided and inserts it.
|
||||
@@ -134,9 +148,13 @@ func GenerateLicense(t *testing.T, options LicenseOptions) string {
|
||||
if options.WorkspaceQuota {
|
||||
workspaceQuota = 1
|
||||
}
|
||||
highAvailability := int64(0)
|
||||
if options.HighAvailability {
|
||||
highAvailability = 1
|
||||
}
|
||||
|
||||
rbacEnabled := int64(0)
|
||||
if options.TemplateRBACEnabled {
|
||||
if options.TemplateRBAC {
|
||||
rbacEnabled = 1
|
||||
}
|
||||
|
||||
@@ -154,12 +172,13 @@ func GenerateLicense(t *testing.T, options LicenseOptions) string {
|
||||
Version: license.CurrentVersion,
|
||||
AllFeatures: options.AllFeatures,
|
||||
Features: license.Features{
|
||||
UserLimit: options.UserLimit,
|
||||
AuditLog: auditLog,
|
||||
BrowserOnly: browserOnly,
|
||||
SCIM: scim,
|
||||
WorkspaceQuota: workspaceQuota,
|
||||
TemplateRBAC: rbacEnabled,
|
||||
UserLimit: options.UserLimit,
|
||||
AuditLog: auditLog,
|
||||
BrowserOnly: browserOnly,
|
||||
SCIM: scim,
|
||||
WorkspaceQuota: workspaceQuota,
|
||||
HighAvailability: highAvailability,
|
||||
TemplateRBAC: rbacEnabled,
|
||||
},
|
||||
}
|
||||
tok := jwt.NewWithClaims(jwt.SigningMethodEdDSA, c)
|
||||
|
||||
@@ -33,7 +33,7 @@ func TestAuthorizeAllEndpoints(t *testing.T) {
|
||||
ctx, _ := testutil.Context(t)
|
||||
admin := coderdtest.CreateFirstUser(t, client)
|
||||
license := coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
group, err := client.CreateGroup(ctx, admin.OrganizationID, codersdk.CreateGroupRequest{
|
||||
Name: "testgroup",
|
||||
@@ -58,6 +58,10 @@ func TestAuthorizeAllEndpoints(t *testing.T) {
|
||||
AssertAction: rbac.ActionRead,
|
||||
AssertObject: rbac.ResourceLicense,
|
||||
}
|
||||
assertRoute["GET:/api/v2/replicas"] = coderdtest.RouteCheck{
|
||||
AssertAction: rbac.ActionRead,
|
||||
AssertObject: rbac.ResourceReplicas,
|
||||
}
|
||||
assertRoute["DELETE:/api/v2/licenses/{id}"] = coderdtest.RouteCheck{
|
||||
AssertAction: rbac.ActionDelete,
|
||||
AssertObject: rbac.ResourceLicense,
|
||||
|
||||
@@ -24,7 +24,7 @@ func TestCreateGroup(t *testing.T) {
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
ctx, _ := testutil.Context(t)
|
||||
group, err := client.CreateGroup(ctx, user.OrganizationID, codersdk.CreateGroupRequest{
|
||||
@@ -43,7 +43,7 @@ func TestCreateGroup(t *testing.T) {
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
ctx, _ := testutil.Context(t)
|
||||
_, err := client.CreateGroup(ctx, user.OrganizationID, codersdk.CreateGroupRequest{
|
||||
@@ -67,7 +67,7 @@ func TestCreateGroup(t *testing.T) {
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
ctx, _ := testutil.Context(t)
|
||||
_, err := client.CreateGroup(ctx, user.OrganizationID, codersdk.CreateGroupRequest{
|
||||
@@ -90,7 +90,7 @@ func TestPatchGroup(t *testing.T) {
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
ctx, _ := testutil.Context(t)
|
||||
group, err := client.CreateGroup(ctx, user.OrganizationID, codersdk.CreateGroupRequest{
|
||||
@@ -112,7 +112,7 @@ func TestPatchGroup(t *testing.T) {
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
_, user2 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
_, user3 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
@@ -138,7 +138,7 @@ func TestPatchGroup(t *testing.T) {
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
_, user2 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
_, user3 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
@@ -173,7 +173,7 @@ func TestPatchGroup(t *testing.T) {
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
ctx, _ := testutil.Context(t)
|
||||
group, err := client.CreateGroup(ctx, user.OrganizationID, codersdk.CreateGroupRequest{
|
||||
@@ -197,7 +197,7 @@ func TestPatchGroup(t *testing.T) {
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
ctx, _ := testutil.Context(t)
|
||||
group, err := client.CreateGroup(ctx, user.OrganizationID, codersdk.CreateGroupRequest{
|
||||
@@ -221,7 +221,7 @@ func TestPatchGroup(t *testing.T) {
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
_, user2 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
ctx, _ := testutil.Context(t)
|
||||
@@ -247,7 +247,7 @@ func TestPatchGroup(t *testing.T) {
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
ctx, _ := testutil.Context(t)
|
||||
group, err := client.CreateGroup(ctx, user.OrganizationID, codersdk.CreateGroupRequest{
|
||||
@@ -276,7 +276,7 @@ func TestGroup(t *testing.T) {
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
ctx, _ := testutil.Context(t)
|
||||
group, err := client.CreateGroup(ctx, user.OrganizationID, codersdk.CreateGroupRequest{
|
||||
@@ -296,7 +296,7 @@ func TestGroup(t *testing.T) {
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
_, user2 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
_, user3 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
@@ -326,7 +326,7 @@ func TestGroup(t *testing.T) {
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
client1, _ := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
|
||||
@@ -347,7 +347,7 @@ func TestGroup(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
_, user1 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
@@ -380,7 +380,7 @@ func TestGroup(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
_, user1 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
@@ -421,7 +421,7 @@ func TestGroups(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
_, user2 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
_, user3 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
@@ -467,7 +467,7 @@ func TestDeleteGroup(t *testing.T) {
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
ctx, _ := testutil.Context(t)
|
||||
group1, err := client.CreateGroup(ctx, user.OrganizationID, codersdk.CreateGroupRequest{
|
||||
@@ -492,7 +492,7 @@ func TestDeleteGroup(t *testing.T) {
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
ctx, _ := testutil.Context(t)
|
||||
err := client.DeleteGroup(ctx, user.OrganizationID)
|
||||
|
||||
@@ -17,12 +17,20 @@ import (
|
||||
)
|
||||
|
||||
// Entitlements processes licenses to return whether features are enabled or not.
|
||||
func Entitlements(ctx context.Context, db database.Store, logger slog.Logger, keys map[string]ed25519.PublicKey, enablements map[string]bool) (codersdk.Entitlements, error) {
|
||||
func Entitlements(
|
||||
ctx context.Context,
|
||||
db database.Store,
|
||||
logger slog.Logger,
|
||||
replicaCount int,
|
||||
keys map[string]ed25519.PublicKey,
|
||||
enablements map[string]bool,
|
||||
) (codersdk.Entitlements, error) {
|
||||
now := time.Now()
|
||||
// Default all entitlements to be disabled.
|
||||
entitlements := codersdk.Entitlements{
|
||||
Features: map[string]codersdk.Feature{},
|
||||
Warnings: []string{},
|
||||
Errors: []string{},
|
||||
}
|
||||
for _, featureName := range codersdk.FeatureNames {
|
||||
entitlements.Features[featureName] = codersdk.Feature{
|
||||
@@ -96,6 +104,12 @@ func Entitlements(ctx context.Context, db database.Store, logger slog.Logger, ke
|
||||
Enabled: enablements[codersdk.FeatureWorkspaceQuota],
|
||||
}
|
||||
}
|
||||
if claims.Features.HighAvailability > 0 {
|
||||
entitlements.Features[codersdk.FeatureHighAvailability] = codersdk.Feature{
|
||||
Entitlement: entitlement,
|
||||
Enabled: enablements[codersdk.FeatureHighAvailability],
|
||||
}
|
||||
}
|
||||
if claims.Features.TemplateRBAC > 0 {
|
||||
entitlements.Features[codersdk.FeatureTemplateRBAC] = codersdk.Feature{
|
||||
Entitlement: entitlement,
|
||||
@@ -132,6 +146,10 @@ func Entitlements(ctx context.Context, db database.Store, logger slog.Logger, ke
|
||||
if featureName == codersdk.FeatureUserLimit {
|
||||
continue
|
||||
}
|
||||
// High availability has it's own warnings based on replica count!
|
||||
if featureName == codersdk.FeatureHighAvailability {
|
||||
continue
|
||||
}
|
||||
feature := entitlements.Features[featureName]
|
||||
if !feature.Enabled {
|
||||
continue
|
||||
@@ -141,9 +159,6 @@ func Entitlements(ctx context.Context, db database.Store, logger slog.Logger, ke
|
||||
case codersdk.EntitlementNotEntitled:
|
||||
entitlements.Warnings = append(entitlements.Warnings,
|
||||
fmt.Sprintf("%s is enabled but your license is not entitled to this feature.", niceName))
|
||||
// Disable the feature and add a warning...
|
||||
feature.Enabled = false
|
||||
entitlements.Features[featureName] = feature
|
||||
case codersdk.EntitlementGracePeriod:
|
||||
entitlements.Warnings = append(entitlements.Warnings,
|
||||
fmt.Sprintf("%s is enabled but your license for this feature is expired.", niceName))
|
||||
@@ -152,6 +167,32 @@ func Entitlements(ctx context.Context, db database.Store, logger slog.Logger, ke
|
||||
}
|
||||
}
|
||||
|
||||
if replicaCount > 1 {
|
||||
feature := entitlements.Features[codersdk.FeatureHighAvailability]
|
||||
|
||||
switch feature.Entitlement {
|
||||
case codersdk.EntitlementNotEntitled:
|
||||
if entitlements.HasLicense {
|
||||
entitlements.Errors = append(entitlements.Warnings,
|
||||
"You have multiple replicas but your license is not entitled to high availability. You will be unable to connect to workspaces.")
|
||||
} else {
|
||||
entitlements.Errors = append(entitlements.Warnings,
|
||||
"You have multiple replicas but high availability is an Enterprise feature. You will be unable to connect to workspaces.")
|
||||
}
|
||||
case codersdk.EntitlementGracePeriod:
|
||||
entitlements.Warnings = append(entitlements.Warnings,
|
||||
"You have multiple replicas but your license for high availability is expired. Reduce to one replica or workspace connections will stop working.")
|
||||
}
|
||||
}
|
||||
|
||||
for _, featureName := range codersdk.FeatureNames {
|
||||
feature := entitlements.Features[featureName]
|
||||
if feature.Entitlement == codersdk.EntitlementNotEntitled {
|
||||
feature.Enabled = false
|
||||
entitlements.Features[featureName] = feature
|
||||
}
|
||||
}
|
||||
|
||||
return entitlements, nil
|
||||
}
|
||||
|
||||
@@ -171,12 +212,13 @@ var (
|
||||
)
|
||||
|
||||
type Features struct {
|
||||
UserLimit int64 `json:"user_limit"`
|
||||
AuditLog int64 `json:"audit_log"`
|
||||
BrowserOnly int64 `json:"browser_only"`
|
||||
SCIM int64 `json:"scim"`
|
||||
WorkspaceQuota int64 `json:"workspace_quota"`
|
||||
TemplateRBAC int64 `json:"template_rbac"`
|
||||
UserLimit int64 `json:"user_limit"`
|
||||
AuditLog int64 `json:"audit_log"`
|
||||
BrowserOnly int64 `json:"browser_only"`
|
||||
SCIM int64 `json:"scim"`
|
||||
WorkspaceQuota int64 `json:"workspace_quota"`
|
||||
TemplateRBAC int64 `json:"template_rbac"`
|
||||
HighAvailability int64 `json:"high_availability"`
|
||||
}
|
||||
|
||||
type Claims struct {
|
||||
|
||||
@@ -20,17 +20,18 @@ import (
|
||||
func TestEntitlements(t *testing.T) {
|
||||
t.Parallel()
|
||||
all := map[string]bool{
|
||||
codersdk.FeatureAuditLog: true,
|
||||
codersdk.FeatureBrowserOnly: true,
|
||||
codersdk.FeatureSCIM: true,
|
||||
codersdk.FeatureWorkspaceQuota: true,
|
||||
codersdk.FeatureTemplateRBAC: true,
|
||||
codersdk.FeatureAuditLog: true,
|
||||
codersdk.FeatureBrowserOnly: true,
|
||||
codersdk.FeatureSCIM: true,
|
||||
codersdk.FeatureWorkspaceQuota: true,
|
||||
codersdk.FeatureHighAvailability: true,
|
||||
codersdk.FeatureTemplateRBAC: true,
|
||||
}
|
||||
|
||||
t.Run("Defaults", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db := databasefake.New()
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, coderdenttest.Keys, map[string]bool{})
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, 1, coderdenttest.Keys, all)
|
||||
require.NoError(t, err)
|
||||
require.False(t, entitlements.HasLicense)
|
||||
require.False(t, entitlements.Trial)
|
||||
@@ -46,7 +47,7 @@ func TestEntitlements(t *testing.T) {
|
||||
JWT: coderdenttest.GenerateLicense(t, coderdenttest.LicenseOptions{}),
|
||||
Exp: time.Now().Add(time.Hour),
|
||||
})
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, coderdenttest.Keys, map[string]bool{})
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, 1, coderdenttest.Keys, map[string]bool{})
|
||||
require.NoError(t, err)
|
||||
require.True(t, entitlements.HasLicense)
|
||||
require.False(t, entitlements.Trial)
|
||||
@@ -60,16 +61,17 @@ func TestEntitlements(t *testing.T) {
|
||||
db := databasefake.New()
|
||||
db.InsertLicense(context.Background(), database.InsertLicenseParams{
|
||||
JWT: coderdenttest.GenerateLicense(t, coderdenttest.LicenseOptions{
|
||||
UserLimit: 100,
|
||||
AuditLog: true,
|
||||
BrowserOnly: true,
|
||||
SCIM: true,
|
||||
WorkspaceQuota: true,
|
||||
TemplateRBACEnabled: true,
|
||||
UserLimit: 100,
|
||||
AuditLog: true,
|
||||
BrowserOnly: true,
|
||||
SCIM: true,
|
||||
WorkspaceQuota: true,
|
||||
HighAvailability: true,
|
||||
TemplateRBAC: true,
|
||||
}),
|
||||
Exp: time.Now().Add(time.Hour),
|
||||
})
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, coderdenttest.Keys, map[string]bool{})
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, 1, coderdenttest.Keys, map[string]bool{})
|
||||
require.NoError(t, err)
|
||||
require.True(t, entitlements.HasLicense)
|
||||
require.False(t, entitlements.Trial)
|
||||
@@ -82,18 +84,19 @@ func TestEntitlements(t *testing.T) {
|
||||
db := databasefake.New()
|
||||
db.InsertLicense(context.Background(), database.InsertLicenseParams{
|
||||
JWT: coderdenttest.GenerateLicense(t, coderdenttest.LicenseOptions{
|
||||
UserLimit: 100,
|
||||
AuditLog: true,
|
||||
BrowserOnly: true,
|
||||
SCIM: true,
|
||||
WorkspaceQuota: true,
|
||||
TemplateRBACEnabled: true,
|
||||
GraceAt: time.Now().Add(-time.Hour),
|
||||
ExpiresAt: time.Now().Add(time.Hour),
|
||||
UserLimit: 100,
|
||||
AuditLog: true,
|
||||
BrowserOnly: true,
|
||||
SCIM: true,
|
||||
WorkspaceQuota: true,
|
||||
HighAvailability: true,
|
||||
TemplateRBAC: true,
|
||||
GraceAt: time.Now().Add(-time.Hour),
|
||||
ExpiresAt: time.Now().Add(time.Hour),
|
||||
}),
|
||||
Exp: time.Now().Add(time.Hour),
|
||||
})
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, coderdenttest.Keys, all)
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, 1, coderdenttest.Keys, all)
|
||||
require.NoError(t, err)
|
||||
require.True(t, entitlements.HasLicense)
|
||||
require.False(t, entitlements.Trial)
|
||||
@@ -101,6 +104,9 @@ func TestEntitlements(t *testing.T) {
|
||||
if featureName == codersdk.FeatureUserLimit {
|
||||
continue
|
||||
}
|
||||
if featureName == codersdk.FeatureHighAvailability {
|
||||
continue
|
||||
}
|
||||
niceName := strings.Title(strings.ReplaceAll(featureName, "_", " "))
|
||||
require.Equal(t, codersdk.EntitlementGracePeriod, entitlements.Features[featureName].Entitlement)
|
||||
require.Contains(t, entitlements.Warnings, fmt.Sprintf("%s is enabled but your license for this feature is expired.", niceName))
|
||||
@@ -113,7 +119,7 @@ func TestEntitlements(t *testing.T) {
|
||||
JWT: coderdenttest.GenerateLicense(t, coderdenttest.LicenseOptions{}),
|
||||
Exp: time.Now().Add(time.Hour),
|
||||
})
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, coderdenttest.Keys, all)
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, 1, coderdenttest.Keys, all)
|
||||
require.NoError(t, err)
|
||||
require.True(t, entitlements.HasLicense)
|
||||
require.False(t, entitlements.Trial)
|
||||
@@ -121,6 +127,9 @@ func TestEntitlements(t *testing.T) {
|
||||
if featureName == codersdk.FeatureUserLimit {
|
||||
continue
|
||||
}
|
||||
if featureName == codersdk.FeatureHighAvailability {
|
||||
continue
|
||||
}
|
||||
niceName := strings.Title(strings.ReplaceAll(featureName, "_", " "))
|
||||
// Ensures features that are not entitled are properly disabled.
|
||||
require.False(t, entitlements.Features[featureName].Enabled)
|
||||
@@ -139,7 +148,7 @@ func TestEntitlements(t *testing.T) {
|
||||
}),
|
||||
Exp: time.Now().Add(time.Hour),
|
||||
})
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, coderdenttest.Keys, map[string]bool{})
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, 1, coderdenttest.Keys, map[string]bool{})
|
||||
require.NoError(t, err)
|
||||
require.True(t, entitlements.HasLicense)
|
||||
require.Contains(t, entitlements.Warnings, "Your deployment has 2 active users but is only licensed for 1.")
|
||||
@@ -161,7 +170,7 @@ func TestEntitlements(t *testing.T) {
|
||||
}),
|
||||
Exp: time.Now().Add(time.Hour),
|
||||
})
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, coderdenttest.Keys, map[string]bool{})
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, 1, coderdenttest.Keys, map[string]bool{})
|
||||
require.NoError(t, err)
|
||||
require.True(t, entitlements.HasLicense)
|
||||
require.Empty(t, entitlements.Warnings)
|
||||
@@ -184,7 +193,7 @@ func TestEntitlements(t *testing.T) {
|
||||
}),
|
||||
})
|
||||
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, coderdenttest.Keys, map[string]bool{})
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, 1, coderdenttest.Keys, map[string]bool{})
|
||||
require.NoError(t, err)
|
||||
require.True(t, entitlements.HasLicense)
|
||||
require.False(t, entitlements.Trial)
|
||||
@@ -199,7 +208,7 @@ func TestEntitlements(t *testing.T) {
|
||||
AllFeatures: true,
|
||||
}),
|
||||
})
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, coderdenttest.Keys, all)
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, 1, coderdenttest.Keys, all)
|
||||
require.NoError(t, err)
|
||||
require.True(t, entitlements.HasLicense)
|
||||
require.False(t, entitlements.Trial)
|
||||
@@ -211,4 +220,52 @@ func TestEntitlements(t *testing.T) {
|
||||
require.Equal(t, codersdk.EntitlementEntitled, entitlements.Features[featureName].Entitlement)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("MultipleReplicasNoLicense", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db := databasefake.New()
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, 2, coderdenttest.Keys, all)
|
||||
require.NoError(t, err)
|
||||
require.False(t, entitlements.HasLicense)
|
||||
require.Len(t, entitlements.Errors, 1)
|
||||
require.Equal(t, "You have multiple replicas but high availability is an Enterprise feature. You will be unable to connect to workspaces.", entitlements.Errors[0])
|
||||
})
|
||||
|
||||
t.Run("MultipleReplicasNotEntitled", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db := databasefake.New()
|
||||
db.InsertLicense(context.Background(), database.InsertLicenseParams{
|
||||
Exp: time.Now().Add(time.Hour),
|
||||
JWT: coderdenttest.GenerateLicense(t, coderdenttest.LicenseOptions{
|
||||
AuditLog: true,
|
||||
}),
|
||||
})
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, 2, coderdenttest.Keys, map[string]bool{
|
||||
codersdk.FeatureHighAvailability: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.True(t, entitlements.HasLicense)
|
||||
require.Len(t, entitlements.Errors, 1)
|
||||
require.Equal(t, "You have multiple replicas but your license is not entitled to high availability. You will be unable to connect to workspaces.", entitlements.Errors[0])
|
||||
})
|
||||
|
||||
t.Run("MultipleReplicasGrace", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db := databasefake.New()
|
||||
db.InsertLicense(context.Background(), database.InsertLicenseParams{
|
||||
JWT: coderdenttest.GenerateLicense(t, coderdenttest.LicenseOptions{
|
||||
HighAvailability: true,
|
||||
GraceAt: time.Now().Add(-time.Hour),
|
||||
ExpiresAt: time.Now().Add(time.Hour),
|
||||
}),
|
||||
Exp: time.Now().Add(time.Hour),
|
||||
})
|
||||
entitlements, err := license.Entitlements(context.Background(), db, slog.Logger{}, 2, coderdenttest.Keys, map[string]bool{
|
||||
codersdk.FeatureHighAvailability: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.True(t, entitlements.HasLicense)
|
||||
require.Len(t, entitlements.Warnings, 1)
|
||||
require.Equal(t, "You have multiple replicas but your license for high availability is expired. Reduce to one replica or workspace connections will stop working.", entitlements.Warnings[0])
|
||||
})
|
||||
}
|
||||
|
||||
@@ -78,21 +78,21 @@ func TestGetLicense(t *testing.T) {
|
||||
defer cancel()
|
||||
|
||||
coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
AccountID: "testing",
|
||||
AuditLog: true,
|
||||
SCIM: true,
|
||||
BrowserOnly: true,
|
||||
TemplateRBACEnabled: true,
|
||||
AccountID: "testing",
|
||||
AuditLog: true,
|
||||
SCIM: true,
|
||||
BrowserOnly: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
AccountID: "testing2",
|
||||
AuditLog: true,
|
||||
SCIM: true,
|
||||
BrowserOnly: true,
|
||||
Trial: true,
|
||||
UserLimit: 200,
|
||||
TemplateRBACEnabled: false,
|
||||
AccountID: "testing2",
|
||||
AuditLog: true,
|
||||
SCIM: true,
|
||||
BrowserOnly: true,
|
||||
Trial: true,
|
||||
UserLimit: 200,
|
||||
TemplateRBAC: false,
|
||||
})
|
||||
|
||||
licenses, err := client.Licenses(ctx)
|
||||
@@ -101,23 +101,25 @@ func TestGetLicense(t *testing.T) {
|
||||
assert.Equal(t, int32(1), licenses[0].ID)
|
||||
assert.Equal(t, "testing", licenses[0].Claims["account_id"])
|
||||
assert.Equal(t, map[string]interface{}{
|
||||
codersdk.FeatureUserLimit: json.Number("0"),
|
||||
codersdk.FeatureAuditLog: json.Number("1"),
|
||||
codersdk.FeatureSCIM: json.Number("1"),
|
||||
codersdk.FeatureBrowserOnly: json.Number("1"),
|
||||
codersdk.FeatureWorkspaceQuota: json.Number("0"),
|
||||
codersdk.FeatureTemplateRBAC: json.Number("1"),
|
||||
codersdk.FeatureUserLimit: json.Number("0"),
|
||||
codersdk.FeatureAuditLog: json.Number("1"),
|
||||
codersdk.FeatureSCIM: json.Number("1"),
|
||||
codersdk.FeatureBrowserOnly: json.Number("1"),
|
||||
codersdk.FeatureWorkspaceQuota: json.Number("0"),
|
||||
codersdk.FeatureHighAvailability: json.Number("0"),
|
||||
codersdk.FeatureTemplateRBAC: json.Number("1"),
|
||||
}, licenses[0].Claims["features"])
|
||||
assert.Equal(t, int32(2), licenses[1].ID)
|
||||
assert.Equal(t, "testing2", licenses[1].Claims["account_id"])
|
||||
assert.Equal(t, true, licenses[1].Claims["trial"])
|
||||
assert.Equal(t, map[string]interface{}{
|
||||
codersdk.FeatureUserLimit: json.Number("200"),
|
||||
codersdk.FeatureAuditLog: json.Number("1"),
|
||||
codersdk.FeatureSCIM: json.Number("1"),
|
||||
codersdk.FeatureBrowserOnly: json.Number("1"),
|
||||
codersdk.FeatureWorkspaceQuota: json.Number("0"),
|
||||
codersdk.FeatureTemplateRBAC: json.Number("0"),
|
||||
codersdk.FeatureUserLimit: json.Number("200"),
|
||||
codersdk.FeatureAuditLog: json.Number("1"),
|
||||
codersdk.FeatureSCIM: json.Number("1"),
|
||||
codersdk.FeatureBrowserOnly: json.Number("1"),
|
||||
codersdk.FeatureWorkspaceQuota: json.Number("0"),
|
||||
codersdk.FeatureHighAvailability: json.Number("0"),
|
||||
codersdk.FeatureTemplateRBAC: json.Number("0"),
|
||||
}, licenses[1].Claims["features"])
|
||||
})
|
||||
}
|
||||
|
||||
@@ -0,0 +1,37 @@
|
||||
package coderd
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/coder/coder/coderd/database"
|
||||
"github.com/coder/coder/coderd/httpapi"
|
||||
"github.com/coder/coder/coderd/rbac"
|
||||
"github.com/coder/coder/codersdk"
|
||||
)
|
||||
|
||||
// replicas returns the number of replicas that are active in Coder.
|
||||
func (api *API) replicas(rw http.ResponseWriter, r *http.Request) {
|
||||
if !api.AGPL.Authorize(r, rbac.ActionRead, rbac.ResourceReplicas) {
|
||||
httpapi.ResourceNotFound(rw)
|
||||
return
|
||||
}
|
||||
|
||||
replicas := api.replicaManager.All()
|
||||
res := make([]codersdk.Replica, 0, len(replicas))
|
||||
for _, replica := range replicas {
|
||||
res = append(res, convertReplica(replica))
|
||||
}
|
||||
httpapi.Write(r.Context(), rw, http.StatusOK, res)
|
||||
}
|
||||
|
||||
func convertReplica(replica database.Replica) codersdk.Replica {
|
||||
return codersdk.Replica{
|
||||
ID: replica.ID,
|
||||
Hostname: replica.Hostname,
|
||||
CreatedAt: replica.CreatedAt,
|
||||
RelayAddress: replica.RelayAddress,
|
||||
RegionID: replica.RegionID,
|
||||
Error: replica.Error,
|
||||
DatabaseLatency: replica.DatabaseLatency,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,138 @@
|
||||
package coderd_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"cdr.dev/slog"
|
||||
"cdr.dev/slog/sloggers/slogtest"
|
||||
|
||||
"github.com/coder/coder/coderd/coderdtest"
|
||||
"github.com/coder/coder/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/codersdk"
|
||||
"github.com/coder/coder/enterprise/coderd/coderdenttest"
|
||||
"github.com/coder/coder/testutil"
|
||||
)
|
||||
|
||||
func TestReplicas(t *testing.T) {
|
||||
t.Parallel()
|
||||
t.Run("ErrorWithoutLicense", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, pubsub := dbtestutil.NewDB(t)
|
||||
firstClient := coderdenttest.New(t, &coderdenttest.Options{
|
||||
Options: &coderdtest.Options{
|
||||
IncludeProvisionerDaemon: true,
|
||||
Database: db,
|
||||
Pubsub: pubsub,
|
||||
},
|
||||
})
|
||||
_ = coderdtest.CreateFirstUser(t, firstClient)
|
||||
secondClient, _, secondAPI := coderdenttest.NewWithAPI(t, &coderdenttest.Options{
|
||||
Options: &coderdtest.Options{
|
||||
Database: db,
|
||||
Pubsub: pubsub,
|
||||
},
|
||||
})
|
||||
secondClient.SessionToken = firstClient.SessionToken
|
||||
ents, err := secondClient.Entitlements(context.Background())
|
||||
require.NoError(t, err)
|
||||
require.Len(t, ents.Errors, 1)
|
||||
_ = secondAPI.Close()
|
||||
|
||||
ents, err = firstClient.Entitlements(context.Background())
|
||||
require.NoError(t, err)
|
||||
require.Len(t, ents.Warnings, 0)
|
||||
})
|
||||
t.Run("ConnectAcrossMultiple", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, pubsub := dbtestutil.NewDB(t)
|
||||
firstClient := coderdenttest.New(t, &coderdenttest.Options{
|
||||
Options: &coderdtest.Options{
|
||||
IncludeProvisionerDaemon: true,
|
||||
Database: db,
|
||||
Pubsub: pubsub,
|
||||
},
|
||||
})
|
||||
firstUser := coderdtest.CreateFirstUser(t, firstClient)
|
||||
coderdenttest.AddLicense(t, firstClient, coderdenttest.LicenseOptions{
|
||||
HighAvailability: true,
|
||||
})
|
||||
|
||||
secondClient := coderdenttest.New(t, &coderdenttest.Options{
|
||||
Options: &coderdtest.Options{
|
||||
Database: db,
|
||||
Pubsub: pubsub,
|
||||
},
|
||||
})
|
||||
secondClient.SessionToken = firstClient.SessionToken
|
||||
replicas, err := secondClient.Replicas(context.Background())
|
||||
require.NoError(t, err)
|
||||
require.Len(t, replicas, 2)
|
||||
|
||||
_, agent := setupWorkspaceAgent(t, firstClient, firstUser, 0)
|
||||
conn, err := secondClient.DialWorkspaceAgent(context.Background(), agent.ID, &codersdk.DialWorkspaceAgentOptions{
|
||||
BlockEndpoints: true,
|
||||
Logger: slogtest.Make(t, nil).Leveled(slog.LevelDebug),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Eventually(t, func() bool {
|
||||
ctx, cancelFunc := context.WithTimeout(context.Background(), testutil.WaitShort)
|
||||
defer cancelFunc()
|
||||
_, err = conn.Ping(ctx)
|
||||
return err == nil
|
||||
}, testutil.WaitLong, testutil.IntervalFast)
|
||||
_ = conn.Close()
|
||||
})
|
||||
t.Run("ConnectAcrossMultipleTLS", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, pubsub := dbtestutil.NewDB(t)
|
||||
certificates := []tls.Certificate{testutil.GenerateTLSCertificate(t, "localhost")}
|
||||
firstClient := coderdenttest.New(t, &coderdenttest.Options{
|
||||
Options: &coderdtest.Options{
|
||||
IncludeProvisionerDaemon: true,
|
||||
Database: db,
|
||||
Pubsub: pubsub,
|
||||
TLSCertificates: certificates,
|
||||
},
|
||||
})
|
||||
firstUser := coderdtest.CreateFirstUser(t, firstClient)
|
||||
coderdenttest.AddLicense(t, firstClient, coderdenttest.LicenseOptions{
|
||||
HighAvailability: true,
|
||||
})
|
||||
|
||||
secondClient := coderdenttest.New(t, &coderdenttest.Options{
|
||||
Options: &coderdtest.Options{
|
||||
Database: db,
|
||||
Pubsub: pubsub,
|
||||
TLSCertificates: certificates,
|
||||
},
|
||||
})
|
||||
secondClient.SessionToken = firstClient.SessionToken
|
||||
replicas, err := secondClient.Replicas(context.Background())
|
||||
require.NoError(t, err)
|
||||
require.Len(t, replicas, 2)
|
||||
|
||||
_, agent := setupWorkspaceAgent(t, firstClient, firstUser, 0)
|
||||
conn, err := secondClient.DialWorkspaceAgent(context.Background(), agent.ID, &codersdk.DialWorkspaceAgentOptions{
|
||||
BlockEndpoints: true,
|
||||
Logger: slogtest.Make(t, nil).Named("client").Leveled(slog.LevelDebug),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Eventually(t, func() bool {
|
||||
ctx, cancelFunc := context.WithTimeout(context.Background(), testutil.IntervalSlow)
|
||||
defer cancelFunc()
|
||||
_, err = conn.Ping(ctx)
|
||||
return err == nil
|
||||
}, testutil.WaitLong, testutil.IntervalFast)
|
||||
_ = conn.Close()
|
||||
replicas, err = secondClient.Replicas(context.Background())
|
||||
require.NoError(t, err)
|
||||
require.Len(t, replicas, 2)
|
||||
for _, replica := range replicas {
|
||||
require.Empty(t, replica.Error)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -23,7 +23,7 @@ func TestTemplateACL(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
_, user2 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
@@ -64,7 +64,7 @@ func TestTemplateACL(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
_, user1 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
@@ -88,7 +88,7 @@ func TestTemplateACL(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
client1, _ := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
@@ -138,7 +138,7 @@ func TestTemplateACL(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
_, user1 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
@@ -176,7 +176,7 @@ func TestTemplateACL(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
_, user1 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
@@ -214,7 +214,7 @@ func TestTemplateACL(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
version := coderdtest.CreateTemplateVersion(t, client, user.OrganizationID, nil)
|
||||
@@ -262,7 +262,7 @@ func TestTemplateACL(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
client1, user1 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
@@ -318,7 +318,7 @@ func TestUpdateTemplateACL(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
_, user2 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
@@ -361,7 +361,7 @@ func TestUpdateTemplateACL(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
_, user2 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
@@ -422,7 +422,7 @@ func TestUpdateTemplateACL(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
version := coderdtest.CreateTemplateVersion(t, client, user.OrganizationID, nil)
|
||||
@@ -447,7 +447,7 @@ func TestUpdateTemplateACL(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
version := coderdtest.CreateTemplateVersion(t, client, user.OrganizationID, nil)
|
||||
@@ -472,7 +472,7 @@ func TestUpdateTemplateACL(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
_, user2 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
@@ -498,7 +498,7 @@ func TestUpdateTemplateACL(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
client2, user2 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
@@ -533,7 +533,7 @@ func TestUpdateTemplateACL(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
client2, user2 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
@@ -575,7 +575,7 @@ func TestUpdateTemplateACL(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
version := coderdtest.CreateTemplateVersion(t, client, user.OrganizationID, nil)
|
||||
@@ -597,7 +597,7 @@ func TestUpdateTemplateACL(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
client1, user1 := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
@@ -662,7 +662,7 @@ func TestUpdateTemplateACL(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
client1, _ := coderdtest.CreateAnotherUserWithUser(t, client, user.OrganizationID)
|
||||
|
||||
@@ -2,6 +2,7 @@ package coderd_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"testing"
|
||||
@@ -9,7 +10,6 @@ import (
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"cdr.dev/slog"
|
||||
"cdr.dev/slog/sloggers/slogtest"
|
||||
"github.com/coder/coder/agent"
|
||||
"github.com/coder/coder/coderd/coderdtest"
|
||||
@@ -42,7 +42,7 @@ func TestBlockNonBrowser(t *testing.T) {
|
||||
BrowserOnly: true,
|
||||
})
|
||||
_, agent := setupWorkspaceAgent(t, client, user, 0)
|
||||
_, err := client.DialWorkspaceAgentTailnet(context.Background(), slog.Logger{}, agent.ID)
|
||||
_, err := client.DialWorkspaceAgent(context.Background(), agent.ID, nil)
|
||||
var apiErr *codersdk.Error
|
||||
require.ErrorAs(t, err, &apiErr)
|
||||
require.Equal(t, http.StatusConflict, apiErr.StatusCode())
|
||||
@@ -59,7 +59,7 @@ func TestBlockNonBrowser(t *testing.T) {
|
||||
BrowserOnly: false,
|
||||
})
|
||||
_, agent := setupWorkspaceAgent(t, client, user, 0)
|
||||
conn, err := client.DialWorkspaceAgentTailnet(context.Background(), slog.Logger{}, agent.ID)
|
||||
conn, err := client.DialWorkspaceAgent(context.Background(), agent.ID, nil)
|
||||
require.NoError(t, err)
|
||||
_ = conn.Close()
|
||||
})
|
||||
@@ -109,6 +109,14 @@ func setupWorkspaceAgent(t *testing.T, client *codersdk.Client, user codersdk.Cr
|
||||
workspace := coderdtest.CreateWorkspace(t, client, user.OrganizationID, template.ID)
|
||||
coderdtest.AwaitWorkspaceBuildJob(t, client, workspace.LatestBuild.ID)
|
||||
agentClient := codersdk.New(client.URL)
|
||||
agentClient.HTTPClient = &http.Client{
|
||||
Transport: &http.Transport{
|
||||
TLSClientConfig: &tls.Config{
|
||||
//nolint:gosec
|
||||
InsecureSkipVerify: true,
|
||||
},
|
||||
},
|
||||
}
|
||||
agentClient.SessionToken = authToken
|
||||
agentCloser := agent.New(agent.Options{
|
||||
FetchMetadata: agentClient.WorkspaceAgentMetadata,
|
||||
|
||||
@@ -26,7 +26,7 @@ func TestCreateWorkspace(t *testing.T) {
|
||||
client := coderdenttest.New(t, nil)
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
_ = coderdenttest.AddLicense(t, client, coderdenttest.LicenseOptions{
|
||||
TemplateRBACEnabled: true,
|
||||
TemplateRBAC: true,
|
||||
})
|
||||
|
||||
version := coderdtest.CreateTemplateVersion(t, client, user.OrganizationID, nil)
|
||||
|
||||
@@ -0,0 +1,165 @@
|
||||
package derpmesh
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"net"
|
||||
"net/url"
|
||||
"sync"
|
||||
|
||||
"golang.org/x/xerrors"
|
||||
"tailscale.com/derp"
|
||||
"tailscale.com/derp/derphttp"
|
||||
"tailscale.com/types/key"
|
||||
|
||||
"github.com/coder/coder/tailnet"
|
||||
|
||||
"cdr.dev/slog"
|
||||
)
|
||||
|
||||
// New constructs a new mesh for DERP servers.
|
||||
func New(logger slog.Logger, server *derp.Server, tlsConfig *tls.Config) *Mesh {
|
||||
return &Mesh{
|
||||
logger: logger,
|
||||
server: server,
|
||||
tlsConfig: tlsConfig,
|
||||
ctx: context.Background(),
|
||||
closed: make(chan struct{}),
|
||||
active: make(map[string]context.CancelFunc),
|
||||
}
|
||||
}
|
||||
|
||||
type Mesh struct {
|
||||
logger slog.Logger
|
||||
server *derp.Server
|
||||
ctx context.Context
|
||||
tlsConfig *tls.Config
|
||||
|
||||
mutex sync.Mutex
|
||||
closed chan struct{}
|
||||
active map[string]context.CancelFunc
|
||||
}
|
||||
|
||||
// SetAddresses performs a diff of the incoming addresses and adds
|
||||
// or removes DERP clients from the mesh.
|
||||
//
|
||||
// Connect is only used for testing to ensure DERPs are meshed before
|
||||
// exchanging messages.
|
||||
// nolint:revive
|
||||
func (m *Mesh) SetAddresses(addresses []string, connect bool) {
|
||||
total := make(map[string]struct{}, 0)
|
||||
for _, address := range addresses {
|
||||
addressURL, err := url.Parse(address)
|
||||
if err != nil {
|
||||
m.logger.Error(m.ctx, "invalid address", slog.F("address", err), slog.Error(err))
|
||||
continue
|
||||
}
|
||||
derpURL, err := addressURL.Parse("/derp")
|
||||
if err != nil {
|
||||
m.logger.Error(m.ctx, "parse derp", slog.F("address", err), slog.Error(err))
|
||||
continue
|
||||
}
|
||||
address = derpURL.String()
|
||||
|
||||
total[address] = struct{}{}
|
||||
added, err := m.addAddress(address, connect)
|
||||
if err != nil {
|
||||
m.logger.Error(m.ctx, "failed to add address", slog.F("address", address), slog.Error(err))
|
||||
continue
|
||||
}
|
||||
if added {
|
||||
m.logger.Debug(m.ctx, "added mesh address", slog.F("address", address))
|
||||
}
|
||||
}
|
||||
|
||||
m.mutex.Lock()
|
||||
for address := range m.active {
|
||||
_, found := total[address]
|
||||
if found {
|
||||
continue
|
||||
}
|
||||
removed := m.removeAddress(address)
|
||||
if removed {
|
||||
m.logger.Debug(m.ctx, "removed mesh address", slog.F("address", address))
|
||||
}
|
||||
}
|
||||
m.mutex.Unlock()
|
||||
}
|
||||
|
||||
// addAddress begins meshing with a new address. It returns false if the address is already being meshed with.
|
||||
// It's expected that this is a full HTTP address with a path.
|
||||
// e.g. http://127.0.0.1:8080/derp
|
||||
// nolint:revive
|
||||
func (m *Mesh) addAddress(address string, connect bool) (bool, error) {
|
||||
m.mutex.Lock()
|
||||
defer m.mutex.Unlock()
|
||||
if m.isClosed() {
|
||||
return false, nil
|
||||
}
|
||||
_, isActive := m.active[address]
|
||||
if isActive {
|
||||
return false, nil
|
||||
}
|
||||
client, err := derphttp.NewClient(m.server.PrivateKey(), address, tailnet.Logger(m.logger.Named("client")))
|
||||
if err != nil {
|
||||
return false, xerrors.Errorf("create derp client: %w", err)
|
||||
}
|
||||
client.TLSConfig = m.tlsConfig
|
||||
client.MeshKey = m.server.MeshKey()
|
||||
client.SetURLDialer(func(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||
var dialer net.Dialer
|
||||
return dialer.DialContext(ctx, network, addr)
|
||||
})
|
||||
if connect {
|
||||
_ = client.Connect(m.ctx)
|
||||
}
|
||||
ctx, cancelFunc := context.WithCancel(m.ctx)
|
||||
closed := make(chan struct{})
|
||||
closeFunc := func() {
|
||||
cancelFunc()
|
||||
_ = client.Close()
|
||||
<-closed
|
||||
}
|
||||
m.active[address] = closeFunc
|
||||
go func() {
|
||||
defer close(closed)
|
||||
client.RunWatchConnectionLoop(ctx, m.server.PublicKey(), tailnet.Logger(m.logger.Named("loop")), func(np key.NodePublic) {
|
||||
m.server.AddPacketForwarder(np, client)
|
||||
}, func(np key.NodePublic) {
|
||||
m.server.RemovePacketForwarder(np, client)
|
||||
})
|
||||
}()
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// removeAddress stops meshing with a given address.
|
||||
func (m *Mesh) removeAddress(address string) bool {
|
||||
cancelFunc, isActive := m.active[address]
|
||||
if isActive {
|
||||
cancelFunc()
|
||||
}
|
||||
return isActive
|
||||
}
|
||||
|
||||
// Close ends all active meshes with the DERP server.
|
||||
func (m *Mesh) Close() error {
|
||||
m.mutex.Lock()
|
||||
defer m.mutex.Unlock()
|
||||
if m.isClosed() {
|
||||
return nil
|
||||
}
|
||||
close(m.closed)
|
||||
for _, cancelFunc := range m.active {
|
||||
cancelFunc()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Mesh) isClosed() bool {
|
||||
select {
|
||||
case <-m.closed:
|
||||
return true
|
||||
default:
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,219 @@
|
||||
package derpmesh_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/goleak"
|
||||
"tailscale.com/derp"
|
||||
"tailscale.com/derp/derphttp"
|
||||
"tailscale.com/types/key"
|
||||
|
||||
"cdr.dev/slog"
|
||||
"cdr.dev/slog/sloggers/slogtest"
|
||||
"github.com/coder/coder/enterprise/derpmesh"
|
||||
"github.com/coder/coder/tailnet"
|
||||
"github.com/coder/coder/testutil"
|
||||
)
|
||||
|
||||
func TestMain(m *testing.M) {
|
||||
goleak.VerifyTestMain(m)
|
||||
}
|
||||
|
||||
func TestDERPMesh(t *testing.T) {
|
||||
t.Parallel()
|
||||
commonName := "something.org"
|
||||
rawCert := testutil.GenerateTLSCertificate(t, commonName)
|
||||
certificate, err := x509.ParseCertificate(rawCert.Certificate[0])
|
||||
require.NoError(t, err)
|
||||
pool := x509.NewCertPool()
|
||||
pool.AddCert(certificate)
|
||||
tlsConfig := &tls.Config{
|
||||
MinVersion: tls.VersionTLS12,
|
||||
ServerName: commonName,
|
||||
RootCAs: pool,
|
||||
Certificates: []tls.Certificate{rawCert},
|
||||
}
|
||||
|
||||
t.Run("ExchangeMessages", func(t *testing.T) {
|
||||
// This tests messages passing through multiple DERP servers.
|
||||
t.Parallel()
|
||||
firstServer, firstServerURL := startDERP(t, tlsConfig)
|
||||
defer firstServer.Close()
|
||||
secondServer, secondServerURL := startDERP(t, tlsConfig)
|
||||
firstMesh := derpmesh.New(slogtest.Make(t, nil).Named("first").Leveled(slog.LevelDebug), firstServer, tlsConfig)
|
||||
firstMesh.SetAddresses([]string{secondServerURL}, true)
|
||||
secondMesh := derpmesh.New(slogtest.Make(t, nil).Named("second").Leveled(slog.LevelDebug), secondServer, tlsConfig)
|
||||
secondMesh.SetAddresses([]string{firstServerURL}, true)
|
||||
defer firstMesh.Close()
|
||||
defer secondMesh.Close()
|
||||
|
||||
first := key.NewNode()
|
||||
second := key.NewNode()
|
||||
firstClient, err := derphttp.NewClient(first, secondServerURL, tailnet.Logger(slogtest.Make(t, nil)))
|
||||
require.NoError(t, err)
|
||||
firstClient.TLSConfig = tlsConfig
|
||||
secondClient, err := derphttp.NewClient(second, firstServerURL, tailnet.Logger(slogtest.Make(t, nil)))
|
||||
require.NoError(t, err)
|
||||
secondClient.TLSConfig = tlsConfig
|
||||
err = secondClient.Connect(context.Background())
|
||||
require.NoError(t, err)
|
||||
|
||||
closed := make(chan struct{})
|
||||
ctx, cancelFunc := context.WithCancel(context.Background())
|
||||
defer cancelFunc()
|
||||
sent := []byte("hello world")
|
||||
go func() {
|
||||
defer close(closed)
|
||||
ticker := time.NewTicker(50 * time.Millisecond)
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
}
|
||||
err = firstClient.Send(second.Public(), sent)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
}()
|
||||
|
||||
got := recvData(t, secondClient)
|
||||
require.Equal(t, sent, got)
|
||||
cancelFunc()
|
||||
<-closed
|
||||
})
|
||||
t.Run("RemoveAddress", func(t *testing.T) {
|
||||
// This tests messages passing through multiple DERP servers.
|
||||
t.Parallel()
|
||||
server, serverURL := startDERP(t, tlsConfig)
|
||||
mesh := derpmesh.New(slogtest.Make(t, nil).Named("first").Leveled(slog.LevelDebug), server, tlsConfig)
|
||||
mesh.SetAddresses([]string{"http://fake.com"}, false)
|
||||
// This should trigger a removal...
|
||||
mesh.SetAddresses([]string{}, false)
|
||||
defer mesh.Close()
|
||||
|
||||
first := key.NewNode()
|
||||
second := key.NewNode()
|
||||
firstClient, err := derphttp.NewClient(first, serverURL, tailnet.Logger(slogtest.Make(t, nil)))
|
||||
require.NoError(t, err)
|
||||
firstClient.TLSConfig = tlsConfig
|
||||
secondClient, err := derphttp.NewClient(second, serverURL, tailnet.Logger(slogtest.Make(t, nil)))
|
||||
require.NoError(t, err)
|
||||
secondClient.TLSConfig = tlsConfig
|
||||
err = secondClient.Connect(context.Background())
|
||||
require.NoError(t, err)
|
||||
|
||||
closed := make(chan struct{})
|
||||
ctx, cancelFunc := context.WithCancel(context.Background())
|
||||
defer cancelFunc()
|
||||
sent := []byte("hello world")
|
||||
go func() {
|
||||
defer close(closed)
|
||||
ticker := time.NewTicker(50 * time.Millisecond)
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
}
|
||||
err = firstClient.Send(second.Public(), sent)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
}()
|
||||
got := recvData(t, secondClient)
|
||||
require.Equal(t, sent, got)
|
||||
cancelFunc()
|
||||
<-closed
|
||||
})
|
||||
t.Run("TwentyMeshes", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
meshes := make([]*derpmesh.Mesh, 0, 20)
|
||||
serverURLs := make([]string, 0, 20)
|
||||
for i := 0; i < 20; i++ {
|
||||
server, url := startDERP(t, tlsConfig)
|
||||
mesh := derpmesh.New(slogtest.Make(t, nil).Named("mesh").Leveled(slog.LevelDebug), server, tlsConfig)
|
||||
t.Cleanup(func() {
|
||||
_ = server.Close()
|
||||
_ = mesh.Close()
|
||||
})
|
||||
serverURLs = append(serverURLs, url)
|
||||
meshes = append(meshes, mesh)
|
||||
}
|
||||
for _, mesh := range meshes {
|
||||
mesh.SetAddresses(serverURLs, true)
|
||||
}
|
||||
|
||||
first := key.NewNode()
|
||||
second := key.NewNode()
|
||||
firstClient, err := derphttp.NewClient(first, serverURLs[9], tailnet.Logger(slogtest.Make(t, nil)))
|
||||
require.NoError(t, err)
|
||||
firstClient.TLSConfig = tlsConfig
|
||||
secondClient, err := derphttp.NewClient(second, serverURLs[16], tailnet.Logger(slogtest.Make(t, nil)))
|
||||
require.NoError(t, err)
|
||||
secondClient.TLSConfig = tlsConfig
|
||||
err = secondClient.Connect(context.Background())
|
||||
require.NoError(t, err)
|
||||
|
||||
closed := make(chan struct{})
|
||||
ctx, cancelFunc := context.WithCancel(context.Background())
|
||||
defer cancelFunc()
|
||||
sent := []byte("hello world")
|
||||
go func() {
|
||||
defer close(closed)
|
||||
ticker := time.NewTicker(50 * time.Millisecond)
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
}
|
||||
err = firstClient.Send(second.Public(), sent)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
}()
|
||||
|
||||
got := recvData(t, secondClient)
|
||||
require.Equal(t, sent, got)
|
||||
cancelFunc()
|
||||
<-closed
|
||||
})
|
||||
}
|
||||
|
||||
func recvData(t *testing.T, client *derphttp.Client) []byte {
|
||||
for {
|
||||
msg, err := client.Recv()
|
||||
if errors.Is(err, io.EOF) {
|
||||
return nil
|
||||
}
|
||||
assert.NoError(t, err)
|
||||
t.Logf("derp: %T", msg)
|
||||
switch msg := msg.(type) {
|
||||
case derp.ReceivedPacket:
|
||||
return msg.Data
|
||||
default:
|
||||
// Drop all others!
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func startDERP(t *testing.T, tlsConfig *tls.Config) (*derp.Server, string) {
|
||||
logf := tailnet.Logger(slogtest.Make(t, nil))
|
||||
d := derp.NewServer(key.NewNode(), logf)
|
||||
d.SetMeshKey("some-key")
|
||||
server := httptest.NewUnstartedServer(derphttp.Handler(d))
|
||||
server.TLS = tlsConfig
|
||||
server.StartTLS()
|
||||
t.Cleanup(func() {
|
||||
_ = d.Close()
|
||||
})
|
||||
t.Cleanup(server.Close)
|
||||
return d, server.URL
|
||||
}
|
||||
@@ -0,0 +1,391 @@
|
||||
package replicasync
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"cdr.dev/slog"
|
||||
|
||||
"github.com/coder/coder/buildinfo"
|
||||
"github.com/coder/coder/coderd/database"
|
||||
)
|
||||
|
||||
var (
|
||||
PubsubEvent = "replica"
|
||||
)
|
||||
|
||||
type Options struct {
|
||||
CleanupInterval time.Duration
|
||||
UpdateInterval time.Duration
|
||||
PeerTimeout time.Duration
|
||||
RelayAddress string
|
||||
RegionID int32
|
||||
TLSConfig *tls.Config
|
||||
}
|
||||
|
||||
// New registers the replica with the database and periodically updates to ensure
|
||||
// it's healthy. It contacts all other alive replicas to ensure they are reachable.
|
||||
func New(ctx context.Context, logger slog.Logger, db database.Store, pubsub database.Pubsub, options *Options) (*Manager, error) {
|
||||
if options == nil {
|
||||
options = &Options{}
|
||||
}
|
||||
if options.PeerTimeout == 0 {
|
||||
options.PeerTimeout = 3 * time.Second
|
||||
}
|
||||
if options.UpdateInterval == 0 {
|
||||
options.UpdateInterval = 5 * time.Second
|
||||
}
|
||||
if options.CleanupInterval == 0 {
|
||||
// The cleanup interval can be quite long, because it's
|
||||
// primary purpose is to clean up dead replicas.
|
||||
options.CleanupInterval = 30 * time.Minute
|
||||
}
|
||||
hostname, err := os.Hostname()
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("get hostname: %w", err)
|
||||
}
|
||||
databaseLatency, err := db.Ping(ctx)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("ping database: %w", err)
|
||||
}
|
||||
id := uuid.New()
|
||||
replica, err := db.InsertReplica(ctx, database.InsertReplicaParams{
|
||||
ID: id,
|
||||
CreatedAt: database.Now(),
|
||||
StartedAt: database.Now(),
|
||||
UpdatedAt: database.Now(),
|
||||
Hostname: hostname,
|
||||
RegionID: options.RegionID,
|
||||
RelayAddress: options.RelayAddress,
|
||||
Version: buildinfo.Version(),
|
||||
DatabaseLatency: int32(databaseLatency.Microseconds()),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("insert replica: %w", err)
|
||||
}
|
||||
err = pubsub.Publish(PubsubEvent, []byte(id.String()))
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("publish new replica: %w", err)
|
||||
}
|
||||
ctx, cancelFunc := context.WithCancel(ctx)
|
||||
manager := &Manager{
|
||||
id: id,
|
||||
options: options,
|
||||
db: db,
|
||||
pubsub: pubsub,
|
||||
self: replica,
|
||||
logger: logger,
|
||||
closed: make(chan struct{}),
|
||||
closeCancel: cancelFunc,
|
||||
}
|
||||
err = manager.syncReplicas(ctx)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("run replica: %w", err)
|
||||
}
|
||||
peers := manager.Regional()
|
||||
if len(peers) > 0 {
|
||||
self := manager.Self()
|
||||
if self.RelayAddress == "" {
|
||||
return nil, xerrors.Errorf("a relay address must be specified when running multiple replicas in the same region")
|
||||
}
|
||||
}
|
||||
|
||||
err = manager.subscribe(ctx)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("subscribe: %w", err)
|
||||
}
|
||||
manager.closeWait.Add(1)
|
||||
go manager.loop(ctx)
|
||||
return manager, nil
|
||||
}
|
||||
|
||||
// Manager keeps the replica up to date and in sync with other replicas.
|
||||
type Manager struct {
|
||||
id uuid.UUID
|
||||
options *Options
|
||||
db database.Store
|
||||
pubsub database.Pubsub
|
||||
logger slog.Logger
|
||||
|
||||
closeWait sync.WaitGroup
|
||||
closeMutex sync.Mutex
|
||||
closed chan (struct{})
|
||||
closeCancel context.CancelFunc
|
||||
|
||||
self database.Replica
|
||||
mutex sync.Mutex
|
||||
peers []database.Replica
|
||||
callback func()
|
||||
}
|
||||
|
||||
// updateInterval is used to determine a replicas state.
|
||||
// If the replica was updated > the time, it's considered healthy.
|
||||
// If the replica was updated < the time, it's considered stale.
|
||||
func (m *Manager) updateInterval() time.Time {
|
||||
return database.Now().Add(-3 * m.options.UpdateInterval)
|
||||
}
|
||||
|
||||
// loop runs the replica update sequence on an update interval.
|
||||
func (m *Manager) loop(ctx context.Context) {
|
||||
defer m.closeWait.Done()
|
||||
updateTicker := time.NewTicker(m.options.UpdateInterval)
|
||||
defer updateTicker.Stop()
|
||||
deleteTicker := time.NewTicker(m.options.CleanupInterval)
|
||||
defer deleteTicker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-deleteTicker.C:
|
||||
err := m.db.DeleteReplicasUpdatedBefore(ctx, m.updateInterval())
|
||||
if err != nil {
|
||||
m.logger.Warn(ctx, "delete old replicas", slog.Error(err))
|
||||
}
|
||||
continue
|
||||
case <-updateTicker.C:
|
||||
}
|
||||
err := m.syncReplicas(ctx)
|
||||
if err != nil && !errors.Is(err, context.Canceled) {
|
||||
m.logger.Warn(ctx, "run replica update loop", slog.Error(err))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// subscribe listens for new replica information!
|
||||
func (m *Manager) subscribe(ctx context.Context) error {
|
||||
var (
|
||||
needsUpdate = false
|
||||
updating = false
|
||||
updateMutex = sync.Mutex{}
|
||||
)
|
||||
|
||||
// This loop will continually update nodes as updates are processed.
|
||||
// The intent is to always be up to date without spamming the run
|
||||
// function, so if a new update comes in while one is being processed,
|
||||
// it will reprocess afterwards.
|
||||
var update func()
|
||||
update = func() {
|
||||
err := m.syncReplicas(ctx)
|
||||
if err != nil && !errors.Is(err, context.Canceled) {
|
||||
m.logger.Warn(ctx, "run replica from subscribe", slog.Error(err))
|
||||
}
|
||||
updateMutex.Lock()
|
||||
if needsUpdate {
|
||||
needsUpdate = false
|
||||
updateMutex.Unlock()
|
||||
update()
|
||||
return
|
||||
}
|
||||
updating = false
|
||||
updateMutex.Unlock()
|
||||
}
|
||||
cancelFunc, err := m.pubsub.Subscribe(PubsubEvent, func(ctx context.Context, message []byte) {
|
||||
updateMutex.Lock()
|
||||
defer updateMutex.Unlock()
|
||||
id, err := uuid.Parse(string(message))
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
// Don't process updates for ourself!
|
||||
if id == m.id {
|
||||
return
|
||||
}
|
||||
if updating {
|
||||
needsUpdate = true
|
||||
return
|
||||
}
|
||||
updating = true
|
||||
go update()
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
go func() {
|
||||
<-ctx.Done()
|
||||
cancelFunc()
|
||||
}()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Manager) syncReplicas(ctx context.Context) error {
|
||||
m.closeMutex.Lock()
|
||||
m.closeWait.Add(1)
|
||||
m.closeMutex.Unlock()
|
||||
defer m.closeWait.Done()
|
||||
// Expect replicas to update once every three times the interval...
|
||||
// If they don't, assume death!
|
||||
replicas, err := m.db.GetReplicasUpdatedAfter(ctx, m.updateInterval())
|
||||
if err != nil {
|
||||
return xerrors.Errorf("get replicas: %w", err)
|
||||
}
|
||||
|
||||
m.mutex.Lock()
|
||||
m.peers = make([]database.Replica, 0, len(replicas))
|
||||
for _, replica := range replicas {
|
||||
if replica.ID == m.id {
|
||||
continue
|
||||
}
|
||||
m.peers = append(m.peers, replica)
|
||||
}
|
||||
m.mutex.Unlock()
|
||||
|
||||
client := http.Client{
|
||||
Timeout: m.options.PeerTimeout,
|
||||
Transport: &http.Transport{
|
||||
TLSClientConfig: m.options.TLSConfig,
|
||||
},
|
||||
}
|
||||
defer client.CloseIdleConnections()
|
||||
var wg sync.WaitGroup
|
||||
var mu sync.Mutex
|
||||
failed := make([]string, 0)
|
||||
for _, peer := range m.Regional() {
|
||||
wg.Add(1)
|
||||
go func(peer database.Replica) {
|
||||
defer wg.Done()
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, peer.RelayAddress, nil)
|
||||
if err != nil {
|
||||
m.logger.Warn(ctx, "create http request for relay probe",
|
||||
slog.F("relay_address", peer.RelayAddress), slog.Error(err))
|
||||
return
|
||||
}
|
||||
res, err := client.Do(req)
|
||||
if err != nil {
|
||||
mu.Lock()
|
||||
failed = append(failed, fmt.Sprintf("relay %s (%s): %s", peer.Hostname, peer.RelayAddress, err))
|
||||
mu.Unlock()
|
||||
return
|
||||
}
|
||||
_ = res.Body.Close()
|
||||
}(peer)
|
||||
}
|
||||
wg.Wait()
|
||||
replicaError := ""
|
||||
if len(failed) > 0 {
|
||||
replicaError = fmt.Sprintf("Failed to dial peers: %s", strings.Join(failed, ", "))
|
||||
}
|
||||
|
||||
databaseLatency, err := m.db.Ping(ctx)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("ping database: %w", err)
|
||||
}
|
||||
|
||||
replica, err := m.db.UpdateReplica(ctx, database.UpdateReplicaParams{
|
||||
ID: m.self.ID,
|
||||
UpdatedAt: database.Now(),
|
||||
StartedAt: m.self.StartedAt,
|
||||
StoppedAt: m.self.StoppedAt,
|
||||
RelayAddress: m.self.RelayAddress,
|
||||
RegionID: m.self.RegionID,
|
||||
Hostname: m.self.Hostname,
|
||||
Version: m.self.Version,
|
||||
Error: replicaError,
|
||||
DatabaseLatency: int32(databaseLatency.Microseconds()),
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("update replica: %w", err)
|
||||
}
|
||||
m.mutex.Lock()
|
||||
defer m.mutex.Unlock()
|
||||
if m.self.Error != replica.Error {
|
||||
// Publish an update occurred!
|
||||
err = m.pubsub.Publish(PubsubEvent, []byte(m.self.ID.String()))
|
||||
if err != nil {
|
||||
return xerrors.Errorf("publish replica update: %w", err)
|
||||
}
|
||||
}
|
||||
m.self = replica
|
||||
if m.callback != nil {
|
||||
go m.callback()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Self represents the current replica.
|
||||
func (m *Manager) Self() database.Replica {
|
||||
m.mutex.Lock()
|
||||
defer m.mutex.Unlock()
|
||||
return m.self
|
||||
}
|
||||
|
||||
// All returns every replica, including itself.
|
||||
func (m *Manager) All() []database.Replica {
|
||||
m.mutex.Lock()
|
||||
defer m.mutex.Unlock()
|
||||
return append(m.peers[:], m.self)
|
||||
}
|
||||
|
||||
// Regional returns all replicas in the same region excluding itself.
|
||||
func (m *Manager) Regional() []database.Replica {
|
||||
m.mutex.Lock()
|
||||
defer m.mutex.Unlock()
|
||||
replicas := make([]database.Replica, 0)
|
||||
for _, replica := range m.peers {
|
||||
if replica.RegionID != m.self.RegionID {
|
||||
continue
|
||||
}
|
||||
replicas = append(replicas, replica)
|
||||
}
|
||||
return replicas
|
||||
}
|
||||
|
||||
// SetCallback sets a function to execute whenever new peers
|
||||
// are refreshed or updated.
|
||||
func (m *Manager) SetCallback(callback func()) {
|
||||
m.mutex.Lock()
|
||||
defer m.mutex.Unlock()
|
||||
m.callback = callback
|
||||
// Instantly call the callback to inform replicas!
|
||||
go callback()
|
||||
}
|
||||
|
||||
func (m *Manager) Close() error {
|
||||
m.closeMutex.Lock()
|
||||
select {
|
||||
case <-m.closed:
|
||||
m.closeMutex.Unlock()
|
||||
return nil
|
||||
default:
|
||||
}
|
||||
close(m.closed)
|
||||
m.closeCancel()
|
||||
m.closeWait.Wait()
|
||||
m.closeMutex.Unlock()
|
||||
m.mutex.Lock()
|
||||
defer m.mutex.Unlock()
|
||||
ctx, cancelFunc := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancelFunc()
|
||||
_, err := m.db.UpdateReplica(ctx, database.UpdateReplicaParams{
|
||||
ID: m.self.ID,
|
||||
UpdatedAt: database.Now(),
|
||||
StartedAt: m.self.StartedAt,
|
||||
StoppedAt: sql.NullTime{
|
||||
Time: database.Now(),
|
||||
Valid: true,
|
||||
},
|
||||
RelayAddress: m.self.RelayAddress,
|
||||
RegionID: m.self.RegionID,
|
||||
Hostname: m.self.Hostname,
|
||||
Version: m.self.Version,
|
||||
Error: m.self.Error,
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("update replica: %w", err)
|
||||
}
|
||||
err = m.pubsub.Publish(PubsubEvent, []byte(m.self.ID.String()))
|
||||
if err != nil {
|
||||
return xerrors.Errorf("publish replica update: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,239 @@
|
||||
package replicasync_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/goleak"
|
||||
|
||||
"cdr.dev/slog/sloggers/slogtest"
|
||||
"github.com/coder/coder/coderd/database"
|
||||
"github.com/coder/coder/coderd/database/databasefake"
|
||||
"github.com/coder/coder/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/enterprise/replicasync"
|
||||
"github.com/coder/coder/testutil"
|
||||
)
|
||||
|
||||
func TestMain(m *testing.M) {
|
||||
goleak.VerifyTestMain(m)
|
||||
}
|
||||
|
||||
func TestReplica(t *testing.T) {
|
||||
t.Parallel()
|
||||
t.Run("CreateOnNew", func(t *testing.T) {
|
||||
// This ensures that a new replica is created on New.
|
||||
t.Parallel()
|
||||
db, pubsub := dbtestutil.NewDB(t)
|
||||
closeChan := make(chan struct{}, 1)
|
||||
cancel, err := pubsub.Subscribe(replicasync.PubsubEvent, func(ctx context.Context, message []byte) {
|
||||
closeChan <- struct{}{}
|
||||
})
|
||||
require.NoError(t, err)
|
||||
defer cancel()
|
||||
server, err := replicasync.New(context.Background(), slogtest.Make(t, nil), db, pubsub, nil)
|
||||
require.NoError(t, err)
|
||||
<-closeChan
|
||||
_ = server.Close()
|
||||
require.NoError(t, err)
|
||||
})
|
||||
t.Run("ErrorsWithoutRelayAddress", func(t *testing.T) {
|
||||
// Ensures that the replica reports a successful status for
|
||||
// accessing all of its peers.
|
||||
t.Parallel()
|
||||
db, pubsub := dbtestutil.NewDB(t)
|
||||
_, err := db.InsertReplica(context.Background(), database.InsertReplicaParams{
|
||||
ID: uuid.New(),
|
||||
CreatedAt: database.Now(),
|
||||
StartedAt: database.Now(),
|
||||
UpdatedAt: database.Now(),
|
||||
Hostname: "something",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
_, err = replicasync.New(context.Background(), slogtest.Make(t, nil), db, pubsub, nil)
|
||||
require.Error(t, err)
|
||||
require.Equal(t, "a relay address must be specified when running multiple replicas in the same region", err.Error())
|
||||
})
|
||||
t.Run("ConnectsToPeerReplica", func(t *testing.T) {
|
||||
// Ensures that the replica reports a successful status for
|
||||
// accessing all of its peers.
|
||||
t.Parallel()
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer srv.Close()
|
||||
db, pubsub := dbtestutil.NewDB(t)
|
||||
peer, err := db.InsertReplica(context.Background(), database.InsertReplicaParams{
|
||||
ID: uuid.New(),
|
||||
CreatedAt: database.Now(),
|
||||
StartedAt: database.Now(),
|
||||
UpdatedAt: database.Now(),
|
||||
Hostname: "something",
|
||||
RelayAddress: srv.URL,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
server, err := replicasync.New(context.Background(), slogtest.Make(t, nil), db, pubsub, &replicasync.Options{
|
||||
RelayAddress: "http://169.254.169.254",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, server.Regional(), 1)
|
||||
require.Equal(t, peer.ID, server.Regional()[0].ID)
|
||||
require.Empty(t, server.Self().Error)
|
||||
_ = server.Close()
|
||||
})
|
||||
t.Run("ConnectsToPeerReplicaTLS", func(t *testing.T) {
|
||||
// Ensures that the replica reports a successful status for
|
||||
// accessing all of its peers.
|
||||
t.Parallel()
|
||||
rawCert := testutil.GenerateTLSCertificate(t, "hello.org")
|
||||
certificate, err := x509.ParseCertificate(rawCert.Certificate[0])
|
||||
require.NoError(t, err)
|
||||
pool := x509.NewCertPool()
|
||||
pool.AddCert(certificate)
|
||||
// nolint:gosec
|
||||
tlsConfig := &tls.Config{
|
||||
Certificates: []tls.Certificate{rawCert},
|
||||
ServerName: "hello.org",
|
||||
RootCAs: pool,
|
||||
}
|
||||
srv := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
srv.TLS = tlsConfig
|
||||
srv.StartTLS()
|
||||
defer srv.Close()
|
||||
db, pubsub := dbtestutil.NewDB(t)
|
||||
peer, err := db.InsertReplica(context.Background(), database.InsertReplicaParams{
|
||||
ID: uuid.New(),
|
||||
CreatedAt: database.Now(),
|
||||
StartedAt: database.Now(),
|
||||
UpdatedAt: database.Now(),
|
||||
Hostname: "something",
|
||||
RelayAddress: srv.URL,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
server, err := replicasync.New(context.Background(), slogtest.Make(t, nil), db, pubsub, &replicasync.Options{
|
||||
RelayAddress: "http://169.254.169.254",
|
||||
TLSConfig: tlsConfig,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, server.Regional(), 1)
|
||||
require.Equal(t, peer.ID, server.Regional()[0].ID)
|
||||
require.Empty(t, server.Self().Error)
|
||||
_ = server.Close()
|
||||
})
|
||||
t.Run("ConnectsToFakePeerWithError", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, pubsub := dbtestutil.NewDB(t)
|
||||
peer, err := db.InsertReplica(context.Background(), database.InsertReplicaParams{
|
||||
ID: uuid.New(),
|
||||
CreatedAt: database.Now().Add(time.Minute),
|
||||
StartedAt: database.Now().Add(time.Minute),
|
||||
UpdatedAt: database.Now().Add(time.Minute),
|
||||
Hostname: "something",
|
||||
// Fake address to dial!
|
||||
RelayAddress: "http://127.0.0.1:1",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
server, err := replicasync.New(context.Background(), slogtest.Make(t, nil), db, pubsub, &replicasync.Options{
|
||||
PeerTimeout: 1 * time.Millisecond,
|
||||
RelayAddress: "http://127.0.0.1:1",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, server.Regional(), 1)
|
||||
require.Equal(t, peer.ID, server.Regional()[0].ID)
|
||||
require.NotEmpty(t, server.Self().Error)
|
||||
require.Contains(t, server.Self().Error, "Failed to dial peers")
|
||||
_ = server.Close()
|
||||
})
|
||||
t.Run("RefreshOnPublish", func(t *testing.T) {
|
||||
// Refresh when a new replica appears!
|
||||
t.Parallel()
|
||||
db, pubsub := dbtestutil.NewDB(t)
|
||||
server, err := replicasync.New(context.Background(), slogtest.Make(t, nil), db, pubsub, nil)
|
||||
require.NoError(t, err)
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer srv.Close()
|
||||
peer, err := db.InsertReplica(context.Background(), database.InsertReplicaParams{
|
||||
ID: uuid.New(),
|
||||
RelayAddress: srv.URL,
|
||||
UpdatedAt: database.Now(),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
// Publish multiple times to ensure it can handle that case.
|
||||
err = pubsub.Publish(replicasync.PubsubEvent, []byte(peer.ID.String()))
|
||||
require.NoError(t, err)
|
||||
err = pubsub.Publish(replicasync.PubsubEvent, []byte(peer.ID.String()))
|
||||
require.NoError(t, err)
|
||||
require.Eventually(t, func() bool {
|
||||
return len(server.Regional()) == 1
|
||||
}, testutil.WaitShort, testutil.IntervalFast)
|
||||
_ = server.Close()
|
||||
})
|
||||
t.Run("DeletesOld", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
db, pubsub := dbtestutil.NewDB(t)
|
||||
_, err := db.InsertReplica(context.Background(), database.InsertReplicaParams{
|
||||
ID: uuid.New(),
|
||||
UpdatedAt: database.Now().Add(-time.Hour),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
server, err := replicasync.New(context.Background(), slogtest.Make(t, nil), db, pubsub, &replicasync.Options{
|
||||
RelayAddress: "google.com",
|
||||
CleanupInterval: time.Millisecond,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
defer server.Close()
|
||||
require.Eventually(t, func() bool {
|
||||
return len(server.Regional()) == 0
|
||||
}, testutil.WaitShort, testutil.IntervalFast)
|
||||
})
|
||||
t.Run("TwentyConcurrent", func(t *testing.T) {
|
||||
// Ensures that twenty concurrent replicas can spawn and all
|
||||
// discover each other in parallel!
|
||||
t.Parallel()
|
||||
// This doesn't use the database fake because creating
|
||||
// this many PostgreSQL connections takes some
|
||||
// configuration tweaking.
|
||||
db := databasefake.New()
|
||||
pubsub := database.NewPubsubInMemory()
|
||||
logger := slogtest.Make(t, nil)
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer srv.Close()
|
||||
var wg sync.WaitGroup
|
||||
count := 20
|
||||
wg.Add(count)
|
||||
for i := 0; i < count; i++ {
|
||||
server, err := replicasync.New(context.Background(), logger, db, pubsub, &replicasync.Options{
|
||||
RelayAddress: srv.URL,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() {
|
||||
_ = server.Close()
|
||||
})
|
||||
done := false
|
||||
server.SetCallback(func() {
|
||||
if len(server.All()) != count {
|
||||
return
|
||||
}
|
||||
if done {
|
||||
return
|
||||
}
|
||||
done = true
|
||||
wg.Done()
|
||||
})
|
||||
}
|
||||
wg.Wait()
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,575 @@
|
||||
package tailnet
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"cdr.dev/slog"
|
||||
"github.com/coder/coder/coderd/database"
|
||||
agpl "github.com/coder/coder/tailnet"
|
||||
)
|
||||
|
||||
// NewCoordinator creates a new high availability coordinator
|
||||
// that uses PostgreSQL pubsub to exchange handshakes.
|
||||
func NewCoordinator(logger slog.Logger, pubsub database.Pubsub) (agpl.Coordinator, error) {
|
||||
ctx, cancelFunc := context.WithCancel(context.Background())
|
||||
coord := &haCoordinator{
|
||||
id: uuid.New(),
|
||||
log: logger,
|
||||
pubsub: pubsub,
|
||||
closeFunc: cancelFunc,
|
||||
close: make(chan struct{}),
|
||||
nodes: map[uuid.UUID]*agpl.Node{},
|
||||
agentSockets: map[uuid.UUID]net.Conn{},
|
||||
agentToConnectionSockets: map[uuid.UUID]map[uuid.UUID]net.Conn{},
|
||||
}
|
||||
|
||||
if err := coord.runPubsub(ctx); err != nil {
|
||||
return nil, xerrors.Errorf("run coordinator pubsub: %w", err)
|
||||
}
|
||||
|
||||
return coord, nil
|
||||
}
|
||||
|
||||
type haCoordinator struct {
|
||||
id uuid.UUID
|
||||
log slog.Logger
|
||||
mutex sync.RWMutex
|
||||
pubsub database.Pubsub
|
||||
close chan struct{}
|
||||
closeFunc context.CancelFunc
|
||||
|
||||
// nodes maps agent and connection IDs their respective node.
|
||||
nodes map[uuid.UUID]*agpl.Node
|
||||
// agentSockets maps agent IDs to their open websocket.
|
||||
agentSockets map[uuid.UUID]net.Conn
|
||||
// agentToConnectionSockets maps agent IDs to connection IDs of conns that
|
||||
// are subscribed to updates for that agent.
|
||||
agentToConnectionSockets map[uuid.UUID]map[uuid.UUID]net.Conn
|
||||
}
|
||||
|
||||
// Node returns an in-memory node by ID.
|
||||
func (c *haCoordinator) Node(id uuid.UUID) *agpl.Node {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
node := c.nodes[id]
|
||||
return node
|
||||
}
|
||||
|
||||
// ServeClient accepts a WebSocket connection that wants to connect to an agent
|
||||
// with the specified ID.
|
||||
func (c *haCoordinator) ServeClient(conn net.Conn, id uuid.UUID, agent uuid.UUID) error {
|
||||
c.mutex.Lock()
|
||||
// When a new connection is requested, we update it with the latest
|
||||
// node of the agent. This allows the connection to establish.
|
||||
node, ok := c.nodes[agent]
|
||||
c.mutex.Unlock()
|
||||
if ok {
|
||||
data, err := json.Marshal([]*agpl.Node{node})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("marshal node: %w", err)
|
||||
}
|
||||
_, err = conn.Write(data)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("write nodes: %w", err)
|
||||
}
|
||||
} else {
|
||||
err := c.publishClientHello(agent)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("publish client hello: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
c.mutex.Lock()
|
||||
connectionSockets, ok := c.agentToConnectionSockets[agent]
|
||||
if !ok {
|
||||
connectionSockets = map[uuid.UUID]net.Conn{}
|
||||
c.agentToConnectionSockets[agent] = connectionSockets
|
||||
}
|
||||
|
||||
// Insert this connection into a map so the agent can publish node updates.
|
||||
connectionSockets[id] = conn
|
||||
c.mutex.Unlock()
|
||||
|
||||
defer func() {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
// Clean all traces of this connection from the map.
|
||||
delete(c.nodes, id)
|
||||
connectionSockets, ok := c.agentToConnectionSockets[agent]
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
delete(connectionSockets, id)
|
||||
if len(connectionSockets) != 0 {
|
||||
return
|
||||
}
|
||||
delete(c.agentToConnectionSockets, agent)
|
||||
}()
|
||||
|
||||
decoder := json.NewDecoder(conn)
|
||||
// Indefinitely handle messages from the client websocket.
|
||||
for {
|
||||
err := c.handleNextClientMessage(id, agent, decoder)
|
||||
if err != nil {
|
||||
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrClosedPipe) {
|
||||
return nil
|
||||
}
|
||||
return xerrors.Errorf("handle next client message: %w", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *haCoordinator) handleNextClientMessage(id, agent uuid.UUID, decoder *json.Decoder) error {
|
||||
var node agpl.Node
|
||||
err := decoder.Decode(&node)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("read json: %w", err)
|
||||
}
|
||||
|
||||
c.mutex.Lock()
|
||||
// Update the node of this client in our in-memory map. If an agent entirely
|
||||
// shuts down and reconnects, it needs to be aware of all clients attempting
|
||||
// to establish connections.
|
||||
c.nodes[id] = &node
|
||||
// Write the new node from this client to the actively connected agent.
|
||||
agentSocket, ok := c.agentSockets[agent]
|
||||
c.mutex.Unlock()
|
||||
if !ok {
|
||||
// If we don't own the agent locally, send it over pubsub to a node that
|
||||
// owns the agent.
|
||||
err := c.publishNodesToAgent(agent, []*agpl.Node{&node})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("publish node to agent")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Write the new node from this client to the actively
|
||||
// connected agent.
|
||||
data, err := json.Marshal([]*agpl.Node{&node})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("marshal nodes: %w", err)
|
||||
}
|
||||
|
||||
_, err = agentSocket.Write(data)
|
||||
if err != nil {
|
||||
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrClosedPipe) {
|
||||
return nil
|
||||
}
|
||||
return xerrors.Errorf("write json: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ServeAgent accepts a WebSocket connection to an agent that listens to
|
||||
// incoming connections and publishes node updates.
|
||||
func (c *haCoordinator) ServeAgent(conn net.Conn, id uuid.UUID) error {
|
||||
// Tell clients on other instances to send a callmemaybe to us.
|
||||
err := c.publishAgentHello(id)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("publish agent hello: %w", err)
|
||||
}
|
||||
|
||||
// Publish all nodes on this instance that want to connect to this agent.
|
||||
nodes := c.nodesSubscribedToAgent(id)
|
||||
if len(nodes) > 0 {
|
||||
data, err := json.Marshal(nodes)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("marshal json: %w", err)
|
||||
}
|
||||
_, err = conn.Write(data)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("write nodes: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// If an old agent socket is connected, we close it
|
||||
// to avoid any leaks. This shouldn't ever occur because
|
||||
// we expect one agent to be running.
|
||||
c.mutex.Lock()
|
||||
oldAgentSocket, ok := c.agentSockets[id]
|
||||
if ok {
|
||||
_ = oldAgentSocket.Close()
|
||||
}
|
||||
c.agentSockets[id] = conn
|
||||
c.mutex.Unlock()
|
||||
defer func() {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
delete(c.agentSockets, id)
|
||||
delete(c.nodes, id)
|
||||
}()
|
||||
|
||||
decoder := json.NewDecoder(conn)
|
||||
for {
|
||||
node, err := c.handleAgentUpdate(id, decoder)
|
||||
if err != nil {
|
||||
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrClosedPipe) {
|
||||
return nil
|
||||
}
|
||||
return xerrors.Errorf("handle next agent message: %w", err)
|
||||
}
|
||||
|
||||
err = c.publishAgentToNodes(id, node)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("publish agent to nodes: %w", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *haCoordinator) nodesSubscribedToAgent(agentID uuid.UUID) []*agpl.Node {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
sockets, ok := c.agentToConnectionSockets[agentID]
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
nodes := make([]*agpl.Node, 0, len(sockets))
|
||||
for targetID := range sockets {
|
||||
node, ok := c.nodes[targetID]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
nodes = append(nodes, node)
|
||||
}
|
||||
|
||||
return nodes
|
||||
}
|
||||
|
||||
func (c *haCoordinator) handleClientHello(id uuid.UUID) error {
|
||||
c.mutex.Lock()
|
||||
node, ok := c.nodes[id]
|
||||
c.mutex.Unlock()
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
return c.publishAgentToNodes(id, node)
|
||||
}
|
||||
|
||||
func (c *haCoordinator) handleAgentUpdate(id uuid.UUID, decoder *json.Decoder) (*agpl.Node, error) {
|
||||
var node agpl.Node
|
||||
err := decoder.Decode(&node)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("read json: %w", err)
|
||||
}
|
||||
|
||||
c.mutex.Lock()
|
||||
oldNode := c.nodes[id]
|
||||
if oldNode != nil {
|
||||
if oldNode.AsOf.After(node.AsOf) {
|
||||
c.mutex.Unlock()
|
||||
return oldNode, nil
|
||||
}
|
||||
}
|
||||
c.nodes[id] = &node
|
||||
connectionSockets, ok := c.agentToConnectionSockets[id]
|
||||
if !ok {
|
||||
c.mutex.Unlock()
|
||||
return &node, nil
|
||||
}
|
||||
|
||||
data, err := json.Marshal([]*agpl.Node{&node})
|
||||
if err != nil {
|
||||
c.mutex.Unlock()
|
||||
return nil, xerrors.Errorf("marshal nodes: %w", err)
|
||||
}
|
||||
|
||||
// Publish the new node to every listening socket.
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(len(connectionSockets))
|
||||
for _, connectionSocket := range connectionSockets {
|
||||
connectionSocket := connectionSocket
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
_ = connectionSocket.SetWriteDeadline(time.Now().Add(5 * time.Second))
|
||||
_, _ = connectionSocket.Write(data)
|
||||
}()
|
||||
}
|
||||
c.mutex.Unlock()
|
||||
wg.Wait()
|
||||
return &node, nil
|
||||
}
|
||||
|
||||
// Close closes all of the open connections in the coordinator and stops the
|
||||
// coordinator from accepting new connections.
|
||||
func (c *haCoordinator) Close() error {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
select {
|
||||
case <-c.close:
|
||||
return nil
|
||||
default:
|
||||
}
|
||||
close(c.close)
|
||||
c.closeFunc()
|
||||
|
||||
wg := sync.WaitGroup{}
|
||||
|
||||
wg.Add(len(c.agentSockets))
|
||||
for _, socket := range c.agentSockets {
|
||||
socket := socket
|
||||
go func() {
|
||||
_ = socket.Close()
|
||||
wg.Done()
|
||||
}()
|
||||
}
|
||||
|
||||
for _, connMap := range c.agentToConnectionSockets {
|
||||
wg.Add(len(connMap))
|
||||
for _, socket := range connMap {
|
||||
socket := socket
|
||||
go func() {
|
||||
_ = socket.Close()
|
||||
wg.Done()
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
wg.Wait()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *haCoordinator) publishNodesToAgent(recipient uuid.UUID, nodes []*agpl.Node) error {
|
||||
msg, err := c.formatCallMeMaybe(recipient, nodes)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("format publish message: %w", err)
|
||||
}
|
||||
|
||||
err = c.pubsub.Publish("wireguard_peers", msg)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("publish message: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *haCoordinator) publishAgentHello(id uuid.UUID) error {
|
||||
msg, err := c.formatAgentHello(id)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("format publish message: %w", err)
|
||||
}
|
||||
|
||||
err = c.pubsub.Publish("wireguard_peers", msg)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("publish message: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *haCoordinator) publishClientHello(id uuid.UUID) error {
|
||||
msg, err := c.formatClientHello(id)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("format client hello: %w", err)
|
||||
}
|
||||
err = c.pubsub.Publish("wireguard_peers", msg)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("publish client hello: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *haCoordinator) publishAgentToNodes(id uuid.UUID, node *agpl.Node) error {
|
||||
msg, err := c.formatAgentUpdate(id, node)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("format publish message: %w", err)
|
||||
}
|
||||
|
||||
err = c.pubsub.Publish("wireguard_peers", msg)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("publish message: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *haCoordinator) runPubsub(ctx context.Context) error {
|
||||
messageQueue := make(chan []byte, 64)
|
||||
cancelSub, err := c.pubsub.Subscribe("wireguard_peers", func(ctx context.Context, message []byte) {
|
||||
select {
|
||||
case messageQueue <- message:
|
||||
case <-ctx.Done():
|
||||
return
|
||||
}
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("subscribe wireguard peers")
|
||||
}
|
||||
go func() {
|
||||
for {
|
||||
var message []byte
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case message = <-messageQueue:
|
||||
}
|
||||
c.handlePubsubMessage(ctx, message)
|
||||
}
|
||||
}()
|
||||
|
||||
go func() {
|
||||
defer cancelSub()
|
||||
<-c.close
|
||||
}()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *haCoordinator) handlePubsubMessage(ctx context.Context, message []byte) {
|
||||
sp := bytes.Split(message, []byte("|"))
|
||||
if len(sp) != 4 {
|
||||
c.log.Error(ctx, "invalid wireguard peer message", slog.F("msg", string(message)))
|
||||
return
|
||||
}
|
||||
|
||||
var (
|
||||
coordinatorID = sp[0]
|
||||
eventType = sp[1]
|
||||
agentID = sp[2]
|
||||
nodeJSON = sp[3]
|
||||
)
|
||||
|
||||
sender, err := uuid.ParseBytes(coordinatorID)
|
||||
if err != nil {
|
||||
c.log.Error(ctx, "invalid sender id", slog.F("id", string(coordinatorID)), slog.F("msg", string(message)))
|
||||
return
|
||||
}
|
||||
|
||||
// We sent this message!
|
||||
if sender == c.id {
|
||||
return
|
||||
}
|
||||
|
||||
switch string(eventType) {
|
||||
case "callmemaybe":
|
||||
agentUUID, err := uuid.ParseBytes(agentID)
|
||||
if err != nil {
|
||||
c.log.Error(ctx, "invalid agent id", slog.F("id", string(agentID)))
|
||||
return
|
||||
}
|
||||
|
||||
c.mutex.Lock()
|
||||
agentSocket, ok := c.agentSockets[agentUUID]
|
||||
if !ok {
|
||||
c.mutex.Unlock()
|
||||
return
|
||||
}
|
||||
c.mutex.Unlock()
|
||||
|
||||
// We get a single node over pubsub, so turn into an array.
|
||||
_, err = agentSocket.Write(nodeJSON)
|
||||
if err != nil {
|
||||
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrClosedPipe) {
|
||||
return
|
||||
}
|
||||
c.log.Error(ctx, "send callmemaybe to agent", slog.Error(err))
|
||||
return
|
||||
}
|
||||
case "clienthello":
|
||||
agentUUID, err := uuid.ParseBytes(agentID)
|
||||
if err != nil {
|
||||
c.log.Error(ctx, "invalid agent id", slog.F("id", string(agentID)))
|
||||
return
|
||||
}
|
||||
|
||||
err = c.handleClientHello(agentUUID)
|
||||
if err != nil {
|
||||
c.log.Error(ctx, "handle agent request node", slog.Error(err))
|
||||
return
|
||||
}
|
||||
case "agenthello":
|
||||
agentUUID, err := uuid.ParseBytes(agentID)
|
||||
if err != nil {
|
||||
c.log.Error(ctx, "invalid agent id", slog.F("id", string(agentID)))
|
||||
return
|
||||
}
|
||||
|
||||
nodes := c.nodesSubscribedToAgent(agentUUID)
|
||||
if len(nodes) > 0 {
|
||||
err := c.publishNodesToAgent(agentUUID, nodes)
|
||||
if err != nil {
|
||||
c.log.Error(ctx, "publish nodes to agent", slog.Error(err))
|
||||
return
|
||||
}
|
||||
}
|
||||
case "agentupdate":
|
||||
agentUUID, err := uuid.ParseBytes(agentID)
|
||||
if err != nil {
|
||||
c.log.Error(ctx, "invalid agent id", slog.F("id", string(agentID)))
|
||||
return
|
||||
}
|
||||
|
||||
decoder := json.NewDecoder(bytes.NewReader(nodeJSON))
|
||||
_, err = c.handleAgentUpdate(agentUUID, decoder)
|
||||
if err != nil {
|
||||
c.log.Error(ctx, "handle agent update", slog.Error(err))
|
||||
return
|
||||
}
|
||||
default:
|
||||
c.log.Error(ctx, "unknown peer event", slog.F("name", string(eventType)))
|
||||
}
|
||||
}
|
||||
|
||||
// format: <coordinator id>|callmemaybe|<recipient id>|<node json>
|
||||
func (c *haCoordinator) formatCallMeMaybe(recipient uuid.UUID, nodes []*agpl.Node) ([]byte, error) {
|
||||
buf := bytes.Buffer{}
|
||||
|
||||
buf.WriteString(c.id.String() + "|")
|
||||
buf.WriteString("callmemaybe|")
|
||||
buf.WriteString(recipient.String() + "|")
|
||||
err := json.NewEncoder(&buf).Encode(nodes)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("encode node: %w", err)
|
||||
}
|
||||
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
|
||||
// format: <coordinator id>|agenthello|<node id>|
|
||||
func (c *haCoordinator) formatAgentHello(id uuid.UUID) ([]byte, error) {
|
||||
buf := bytes.Buffer{}
|
||||
|
||||
buf.WriteString(c.id.String() + "|")
|
||||
buf.WriteString("agenthello|")
|
||||
buf.WriteString(id.String() + "|")
|
||||
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
|
||||
// format: <coordinator id>|clienthello|<agent id>|
|
||||
func (c *haCoordinator) formatClientHello(id uuid.UUID) ([]byte, error) {
|
||||
buf := bytes.Buffer{}
|
||||
|
||||
buf.WriteString(c.id.String() + "|")
|
||||
buf.WriteString("clienthello|")
|
||||
buf.WriteString(id.String() + "|")
|
||||
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
|
||||
// format: <coordinator id>|agentupdate|<node id>|<node json>
|
||||
func (c *haCoordinator) formatAgentUpdate(id uuid.UUID, node *agpl.Node) ([]byte, error) {
|
||||
buf := bytes.Buffer{}
|
||||
|
||||
buf.WriteString(c.id.String() + "|")
|
||||
buf.WriteString("agentupdate|")
|
||||
buf.WriteString(id.String() + "|")
|
||||
err := json.NewEncoder(&buf).Encode(node)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("encode node: %w", err)
|
||||
}
|
||||
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
@@ -0,0 +1,261 @@
|
||||
package tailnet_test
|
||||
|
||||
import (
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"cdr.dev/slog/sloggers/slogtest"
|
||||
|
||||
"github.com/coder/coder/coderd/database"
|
||||
"github.com/coder/coder/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/enterprise/tailnet"
|
||||
agpl "github.com/coder/coder/tailnet"
|
||||
"github.com/coder/coder/testutil"
|
||||
)
|
||||
|
||||
func TestCoordinatorSingle(t *testing.T) {
|
||||
t.Parallel()
|
||||
t.Run("ClientWithoutAgent", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
coordinator, err := tailnet.NewCoordinator(slogtest.Make(t, nil), database.NewPubsubInMemory())
|
||||
require.NoError(t, err)
|
||||
defer coordinator.Close()
|
||||
|
||||
client, server := net.Pipe()
|
||||
sendNode, errChan := agpl.ServeCoordinator(client, func(node []*agpl.Node) error {
|
||||
return nil
|
||||
})
|
||||
id := uuid.New()
|
||||
closeChan := make(chan struct{})
|
||||
go func() {
|
||||
err := coordinator.ServeClient(server, id, uuid.New())
|
||||
assert.NoError(t, err)
|
||||
close(closeChan)
|
||||
}()
|
||||
sendNode(&agpl.Node{})
|
||||
require.Eventually(t, func() bool {
|
||||
return coordinator.Node(id) != nil
|
||||
}, testutil.WaitShort, testutil.IntervalFast)
|
||||
|
||||
err = client.Close()
|
||||
require.NoError(t, err)
|
||||
<-errChan
|
||||
<-closeChan
|
||||
})
|
||||
|
||||
t.Run("AgentWithoutClients", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
coordinator, err := tailnet.NewCoordinator(slogtest.Make(t, nil), database.NewPubsubInMemory())
|
||||
require.NoError(t, err)
|
||||
defer coordinator.Close()
|
||||
|
||||
client, server := net.Pipe()
|
||||
sendNode, errChan := agpl.ServeCoordinator(client, func(node []*agpl.Node) error {
|
||||
return nil
|
||||
})
|
||||
id := uuid.New()
|
||||
closeChan := make(chan struct{})
|
||||
go func() {
|
||||
err := coordinator.ServeAgent(server, id)
|
||||
assert.NoError(t, err)
|
||||
close(closeChan)
|
||||
}()
|
||||
sendNode(&agpl.Node{})
|
||||
require.Eventually(t, func() bool {
|
||||
return coordinator.Node(id) != nil
|
||||
}, testutil.WaitShort, testutil.IntervalFast)
|
||||
err = client.Close()
|
||||
require.NoError(t, err)
|
||||
<-errChan
|
||||
<-closeChan
|
||||
})
|
||||
|
||||
t.Run("AgentWithClient", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
coordinator, err := tailnet.NewCoordinator(slogtest.Make(t, nil), database.NewPubsubInMemory())
|
||||
require.NoError(t, err)
|
||||
defer coordinator.Close()
|
||||
|
||||
agentWS, agentServerWS := net.Pipe()
|
||||
defer agentWS.Close()
|
||||
agentNodeChan := make(chan []*agpl.Node)
|
||||
sendAgentNode, agentErrChan := agpl.ServeCoordinator(agentWS, func(nodes []*agpl.Node) error {
|
||||
agentNodeChan <- nodes
|
||||
return nil
|
||||
})
|
||||
agentID := uuid.New()
|
||||
closeAgentChan := make(chan struct{})
|
||||
go func() {
|
||||
err := coordinator.ServeAgent(agentServerWS, agentID)
|
||||
assert.NoError(t, err)
|
||||
close(closeAgentChan)
|
||||
}()
|
||||
sendAgentNode(&agpl.Node{})
|
||||
require.Eventually(t, func() bool {
|
||||
return coordinator.Node(agentID) != nil
|
||||
}, testutil.WaitShort, testutil.IntervalFast)
|
||||
|
||||
clientWS, clientServerWS := net.Pipe()
|
||||
defer clientWS.Close()
|
||||
defer clientServerWS.Close()
|
||||
clientNodeChan := make(chan []*agpl.Node)
|
||||
sendClientNode, clientErrChan := agpl.ServeCoordinator(clientWS, func(nodes []*agpl.Node) error {
|
||||
clientNodeChan <- nodes
|
||||
return nil
|
||||
})
|
||||
clientID := uuid.New()
|
||||
closeClientChan := make(chan struct{})
|
||||
go func() {
|
||||
err := coordinator.ServeClient(clientServerWS, clientID, agentID)
|
||||
assert.NoError(t, err)
|
||||
close(closeClientChan)
|
||||
}()
|
||||
agentNodes := <-clientNodeChan
|
||||
require.Len(t, agentNodes, 1)
|
||||
sendClientNode(&agpl.Node{})
|
||||
clientNodes := <-agentNodeChan
|
||||
require.Len(t, clientNodes, 1)
|
||||
|
||||
// Ensure an update to the agent node reaches the client!
|
||||
sendAgentNode(&agpl.Node{})
|
||||
agentNodes = <-clientNodeChan
|
||||
require.Len(t, agentNodes, 1)
|
||||
|
||||
// Close the agent WebSocket so a new one can connect.
|
||||
err = agentWS.Close()
|
||||
require.NoError(t, err)
|
||||
<-agentErrChan
|
||||
<-closeAgentChan
|
||||
|
||||
// Create a new agent connection. This is to simulate a reconnect!
|
||||
agentWS, agentServerWS = net.Pipe()
|
||||
defer agentWS.Close()
|
||||
agentNodeChan = make(chan []*agpl.Node)
|
||||
_, agentErrChan = agpl.ServeCoordinator(agentWS, func(nodes []*agpl.Node) error {
|
||||
agentNodeChan <- nodes
|
||||
return nil
|
||||
})
|
||||
closeAgentChan = make(chan struct{})
|
||||
go func() {
|
||||
err := coordinator.ServeAgent(agentServerWS, agentID)
|
||||
assert.NoError(t, err)
|
||||
close(closeAgentChan)
|
||||
}()
|
||||
// Ensure the existing listening client sends it's node immediately!
|
||||
clientNodes = <-agentNodeChan
|
||||
require.Len(t, clientNodes, 1)
|
||||
|
||||
err = agentWS.Close()
|
||||
require.NoError(t, err)
|
||||
<-agentErrChan
|
||||
<-closeAgentChan
|
||||
|
||||
err = clientWS.Close()
|
||||
require.NoError(t, err)
|
||||
<-clientErrChan
|
||||
<-closeClientChan
|
||||
})
|
||||
}
|
||||
|
||||
func TestCoordinatorHA(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("AgentWithClient", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
_, pubsub := dbtestutil.NewDB(t)
|
||||
|
||||
coordinator1, err := tailnet.NewCoordinator(slogtest.Make(t, nil), pubsub)
|
||||
require.NoError(t, err)
|
||||
defer coordinator1.Close()
|
||||
|
||||
agentWS, agentServerWS := net.Pipe()
|
||||
defer agentWS.Close()
|
||||
agentNodeChan := make(chan []*agpl.Node)
|
||||
sendAgentNode, agentErrChan := agpl.ServeCoordinator(agentWS, func(nodes []*agpl.Node) error {
|
||||
agentNodeChan <- nodes
|
||||
return nil
|
||||
})
|
||||
agentID := uuid.New()
|
||||
closeAgentChan := make(chan struct{})
|
||||
go func() {
|
||||
err := coordinator1.ServeAgent(agentServerWS, agentID)
|
||||
assert.NoError(t, err)
|
||||
close(closeAgentChan)
|
||||
}()
|
||||
sendAgentNode(&agpl.Node{})
|
||||
require.Eventually(t, func() bool {
|
||||
return coordinator1.Node(agentID) != nil
|
||||
}, testutil.WaitShort, testutil.IntervalFast)
|
||||
|
||||
coordinator2, err := tailnet.NewCoordinator(slogtest.Make(t, nil), pubsub)
|
||||
require.NoError(t, err)
|
||||
defer coordinator2.Close()
|
||||
|
||||
clientWS, clientServerWS := net.Pipe()
|
||||
defer clientWS.Close()
|
||||
defer clientServerWS.Close()
|
||||
clientNodeChan := make(chan []*agpl.Node)
|
||||
sendClientNode, clientErrChan := agpl.ServeCoordinator(clientWS, func(nodes []*agpl.Node) error {
|
||||
clientNodeChan <- nodes
|
||||
return nil
|
||||
})
|
||||
clientID := uuid.New()
|
||||
closeClientChan := make(chan struct{})
|
||||
go func() {
|
||||
err := coordinator2.ServeClient(clientServerWS, clientID, agentID)
|
||||
assert.NoError(t, err)
|
||||
close(closeClientChan)
|
||||
}()
|
||||
agentNodes := <-clientNodeChan
|
||||
require.Len(t, agentNodes, 1)
|
||||
sendClientNode(&agpl.Node{})
|
||||
_ = sendClientNode
|
||||
clientNodes := <-agentNodeChan
|
||||
require.Len(t, clientNodes, 1)
|
||||
|
||||
// Ensure an update to the agent node reaches the client!
|
||||
sendAgentNode(&agpl.Node{})
|
||||
agentNodes = <-clientNodeChan
|
||||
require.Len(t, agentNodes, 1)
|
||||
|
||||
// Close the agent WebSocket so a new one can connect.
|
||||
require.NoError(t, agentWS.Close())
|
||||
require.NoError(t, agentServerWS.Close())
|
||||
<-agentErrChan
|
||||
<-closeAgentChan
|
||||
|
||||
// Create a new agent connection. This is to simulate a reconnect!
|
||||
agentWS, agentServerWS = net.Pipe()
|
||||
defer agentWS.Close()
|
||||
agentNodeChan = make(chan []*agpl.Node)
|
||||
_, agentErrChan = agpl.ServeCoordinator(agentWS, func(nodes []*agpl.Node) error {
|
||||
agentNodeChan <- nodes
|
||||
return nil
|
||||
})
|
||||
closeAgentChan = make(chan struct{})
|
||||
go func() {
|
||||
err := coordinator1.ServeAgent(agentServerWS, agentID)
|
||||
assert.NoError(t, err)
|
||||
close(closeAgentChan)
|
||||
}()
|
||||
// Ensure the existing listening client sends it's node immediately!
|
||||
clientNodes = <-agentNodeChan
|
||||
require.Len(t, clientNodes, 1)
|
||||
|
||||
err = agentWS.Close()
|
||||
require.NoError(t, err)
|
||||
<-agentErrChan
|
||||
<-closeAgentChan
|
||||
|
||||
err = clientWS.Close()
|
||||
require.NoError(t, err)
|
||||
<-clientErrChan
|
||||
<-closeClientChan
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user