Files
Paul Lorenz 1b2ea0c224 Merge pull request #3836 from openziti/oidc-auth-test
OIDC authentication fablab test
2026-08-31 22:39:24 -04:00

1365 lines
41 KiB
Go

/*
Copyright NetFoundry Inc.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
https://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
package tests
import (
"context"
"crypto"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/sha1"
cryptoTls "crypto/tls"
"crypto/x509"
"encoding/json"
"encoding/pem"
"fmt"
"io"
"net"
"net/http"
"net/http/cookiejar"
"net/url"
"os"
"runtime"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/Jeffail/gabs"
"github.com/go-openapi/strfmt"
"github.com/google/uuid"
"github.com/gorilla/websocket"
"github.com/michaelquigley/pfxlog"
"github.com/openziti/channel/v5"
"github.com/openziti/channel/v5/websockets"
"github.com/openziti/edge-api/rest_model"
nfPem "github.com/openziti/foundation/v2/pem"
"github.com/openziti/foundation/v2/util"
"github.com/openziti/foundation/v2/versions"
idlib "github.com/openziti/identity"
"github.com/openziti/identity/certtools"
edgeApis "github.com/openziti/sdk-golang/v2/edge-apis"
"github.com/openziti/sdk-golang/v2/ziti"
"github.com/openziti/sdk-golang/v2/ziti/edge"
sdkEnroll "github.com/openziti/sdk-golang/v2/ziti/enroll"
"github.com/openziti/transport/v2"
"github.com/openziti/transport/v2/tcp"
"github.com/openziti/transport/v2/tls"
"github.com/openziti/ziti/v2/common/eid"
"github.com/openziti/ziti/v2/controller"
"github.com/openziti/ziti/v2/controller/config"
"github.com/openziti/ziti/v2/controller/env"
restClientRouter "github.com/openziti/ziti/v2/controller/rest_client/router"
fabricRestModel "github.com/openziti/ziti/v2/controller/rest_model"
"github.com/openziti/ziti/v2/controller/server"
"github.com/openziti/ziti/v2/controller/xt_smartrouting"
"github.com/openziti/ziti/v2/router"
"github.com/openziti/ziti/v2/router/enroll"
routerEnv "github.com/openziti/ziti/v2/router/env"
"github.com/pkg/errors"
"github.com/sirupsen/logrus"
oidcPkg "github.com/zitadel/oidc/v3/pkg/oidc"
oauth2Pkg "golang.org/x/oauth2"
"gopkg.in/resty.v1"
)
func init() {
logOptions := pfxlog.DefaultOptions().
SetTrimPrefix("github.com/openziti/").
StartingToday()
pfxlog.GlobalInit(logrus.InfoLevel, logOptions)
pfxlog.SetFormatter(pfxlog.NewFormatter(logOptions))
_ = os.Setenv("ZITI_TRACE_ENABLED", "false")
transport.AddAddressParser(tls.AddressParser{})
transport.AddAddressParser(tcp.AddressParser{})
}
func ToPtr[T any](in T) *T {
return &in
}
// ST returns a pointer to a strfmt.Date time. A helper function
// for creating rest_model types
func ST(t time.Time) *strfmt.DateTime {
st := strfmt.DateTime(t)
return &st
}
type TestContext struct {
*CustomAssertions
ApiHost string
AdminAuthenticator *updbAuthenticator
Managers *ManagerHelpers
AdminManagementSession *session
AdminClientSession *session
RestClients *RestClients
fabricController *controller.Controller
EdgeController *server.Controller
Req *CustomAssertions
clientApiClient *resty.Client
managementApiClient *resty.Client
enabledJsonLogging bool
edgeRouterEntity *edgeRouter
transitRouterEntity *transitRouter
routers []*router.Router
peerControllers []*peerController
testing *testing.T
LogLevel string
ControllerConfig *config.Config
configSet ConfigSet
}
var defaultTestContext = newDefaultTestContext()
func newDefaultTestContext() *TestContext {
ctx := &TestContext{
AdminAuthenticator: &updbAuthenticator{
Username: eid.New(),
Password: eid.New(),
},
}
ctx.Managers = &ManagerHelpers{ctx: ctx}
return ctx
}
func NewTestContext(t *testing.T) *TestContext {
ret := &TestContext{
ApiHost: "127.0.0.1:1281",
AdminAuthenticator: &updbAuthenticator{
Username: eid.New(),
Password: eid.New(),
},
LogLevel: os.Getenv("ZITI_TEST_LOG_LEVEL"),
configSet: DefaultATS,
}
ret.Managers = &ManagerHelpers{ctx: ret}
ret.testContextChanged(t)
return ret
}
// NewTestContextWithConfigSet creates a TestContext that uses the supplied ConfigSet
// instead of DefaultATS. All other behavior, including DB path and API host, is unchanged.
func NewTestContextWithConfigSet(t *testing.T, cs ConfigSet) *TestContext {
ret := NewTestContext(t)
ret.configSet = cs
return ret
}
func (ctx *TestContext) controllerConfFile() string {
return ctx.configSet.CtrlConfig
}
func GetTestContext() *TestContext {
return defaultTestContext
}
// testContextChanged is used to update the *testing.T reference used by library
// level tests. Necessary because, using the wrong *testing.T will cause go test library
// errors.
func (ctx *TestContext) testContextChanged(t *testing.T) {
ctx.testing = t
newReq := NewCustomAssertions(t)
ctx.Req = newReq
ctx.CustomAssertions = newReq
}
// NextTest is an alias for testContextChanged and reflects the bbolt testing framework
func (ctx *TestContext) NextTest(t *testing.T) {
ctx.testContextChanged(t)
}
func (ctx *TestContext) T() *testing.T {
return ctx.testing
}
func (ctx *TestContext) NewTransport() *http.Transport {
return ctx.NewTransportWithClientCert(nil, nil)
}
func (ctx *TestContext) ClientApiUrl() *url.URL {
clientApiUrl, err := url.Parse("https://" + ctx.ApiHost + EdgeClientApiPath)
if err != nil {
panic(err)
}
return clientApiUrl
}
func (ctx *TestContext) ManagementApiUrl() *url.URL {
manApiUrl, err := url.Parse("https://" + ctx.ApiHost + EdgeManagementApiPath)
if err != nil {
panic(err)
}
return manApiUrl
}
func (ctx *TestContext) ControllerCaPool() *x509.CertPool {
return ctx.ControllerConfig.Id.CA()
}
func (ctx *TestContext) NewEdgeClientApi(totpProvider func(chan string)) *ClientHelperClient {
if totpProvider == nil {
totpProvider = func(chan string) {}
}
client := edgeApis.NewClientApiClient([]*url.URL{ctx.ClientApiUrl()}, ctx.ControllerCaPool(), totpProvider)
return &ClientHelperClient{
ClientApiClient: client,
testCtx: ctx,
}
}
func (ctx *TestContext) NewEdgeManagementApi(totpProvider func(chan string)) *ManagementHelperClient {
if totpProvider == nil {
totpProvider = func(chan string) {}
}
client := edgeApis.NewManagementApiClient([]*url.URL{ctx.ManagementApiUrl()}, ctx.ControllerCaPool(), totpProvider)
return &ManagementHelperClient{
ManagementApiClient: client,
testCtx: ctx,
}
}
// NewEdgeManagementApiWithToken returns a ManagementHelperClient pre-loaded with
// the given raw OIDC access token. No authentication step is required.
func (ctx *TestContext) NewEdgeManagementApiWithToken(accessToken string) *ManagementHelperClient {
client := ctx.NewEdgeManagementApi(nil)
var session edgeApis.ApiSession = &edgeApis.ApiSessionOidc{
OidcTokens: &oidcPkg.Tokens[*oidcPkg.IDTokenClaims]{
Token: &oauth2Pkg.Token{
AccessToken: accessToken,
},
},
}
client.ApiSession.Store(&session)
return client
}
// NewEdgeClientApiWithToken returns a ClientHelperClient pre-loaded with the
// given raw OIDC access token. No authentication step is required.
func (ctx *TestContext) NewEdgeClientApiWithToken(accessToken string) *ClientHelperClient {
client := ctx.NewEdgeClientApi(nil)
var session edgeApis.ApiSession = &edgeApis.ApiSessionOidc{
OidcTokens: &oidcPkg.Tokens[*oidcPkg.IDTokenClaims]{
Token: &oauth2Pkg.Token{
AccessToken: accessToken,
},
},
}
client.ApiSession.Store(&session)
return client
}
func (ctx *TestContext) NewTransportWithIdentity(i idlib.Identity) *http.Transport {
return &http.Transport{
Proxy: http.ProxyFromEnvironment,
DialContext: (&net.Dialer{
Timeout: 30 * time.Second,
KeepAlive: 30 * time.Second,
}).DialContext,
ForceAttemptHTTP2: true,
MaxIdleConns: 100,
IdleConnTimeout: 90 * time.Second,
TLSHandshakeTimeout: 10 * time.Second,
ExpectContinueTimeout: 1 * time.Second,
TLSClientConfig: i.ClientTLSConfig(),
}
}
func (ctx *TestContext) NewTransportWithClientCert(certs []*x509.Certificate, privateKey crypto.PrivateKey) *http.Transport {
// #nosec
tlsClientConfig := &cryptoTls.Config{
InsecureSkipVerify: true,
}
if certs != nil && privateKey != nil {
rawCerts := make([][]byte, len(certs))
for i, cert := range certs {
rawCerts[i] = cert.Raw
}
tlsClientConfig.Certificates = []cryptoTls.Certificate{
{Certificate: rawCerts, PrivateKey: privateKey, Leaf: certs[0]},
}
}
return &http.Transport{
Proxy: http.ProxyFromEnvironment,
DialContext: (&net.Dialer{
Timeout: 30 * time.Second,
KeepAlive: 30 * time.Second,
}).DialContext,
ForceAttemptHTTP2: true,
MaxIdleConns: 100,
IdleConnTimeout: 90 * time.Second,
TLSHandshakeTimeout: 10 * time.Second,
ExpectContinueTimeout: 1 * time.Second,
TLSClientConfig: tlsClientConfig,
}
}
func (ctx *TestContext) NewHttpClient(transport *http.Transport) *http.Client {
jar, err := cookiejar.New(nil)
ctx.Req.NoError(err)
return &http.Client{
Transport: transport,
CheckRedirect: nil,
Jar: jar,
Timeout: 2000 * time.Second,
}
}
func (ctx *TestContext) NewRestClientWithDefaults() *resty.Client {
return resty.NewWithClient(ctx.NewHttpClient(ctx.NewTransport()))
}
func (ctx *TestContext) NewRestClient(i idlib.Identity) *resty.Client {
return resty.NewWithClient(ctx.NewHttpClient(ctx.NewTransportWithIdentity(i)))
}
func (ctx *TestContext) DefaultClientApiClient() *resty.Client {
if ctx.clientApiClient == nil {
ctx.clientApiClient, _, _ = ctx.NewClientComponents(EdgeClientApiPath)
ctx.clientApiClient.AllowGetMethodPayload = true
}
return ctx.clientApiClient
}
func (ctx *TestContext) DefaultManagementApiClient() *resty.Client {
if ctx.managementApiClient == nil {
ctx.managementApiClient, _, _ = ctx.NewClientComponents(EdgeManagementApiPath)
ctx.managementApiClient.AllowGetMethodPayload = true
}
return ctx.managementApiClient
}
func (ctx *TestContext) NewClientComponents(apiPath string) (*resty.Client, *http.Client, *http.Transport) {
clientTransport := ctx.NewTransport()
httpClient := ctx.NewHttpClient(clientTransport)
client := resty.NewWithClient(httpClient)
apiUrl, err := url.Parse("https://" + ctx.ApiHost)
if err != nil {
panic(err)
}
apiPathUrl, err := url.Parse(apiPath)
if err != nil {
panic(err)
}
baseUrl := apiUrl.ResolveReference(apiPathUrl)
client.SetHostURL(baseUrl.String())
return client, httpClient, clientTransport
}
func (ctx *TestContext) NewWsMgmtChannel(bindHandler channel.BindHandler) (channel.Channel, error) {
log := pfxlog.Logger()
wsUrl := "wss://" + ctx.ApiHost + "/fabric/v1/ws-api"
dialer := &websocket.Dialer{
Proxy: http.ProxyFromEnvironment,
TLSClientConfig: ctx.DefaultClientApiClient().GetClient().Transport.(*http.Transport).TLSClientConfig,
HandshakeTimeout: 5 * time.Second,
}
authHeader := http.Header{}
authHeader.Set(env.ZitiSession, *ctx.AdminManagementSession.AuthResponse.Token)
conn, resp, err := dialer.Dial(wsUrl, authHeader)
if err != nil {
if resp != nil {
if body, rerr := io.ReadAll(resp.Body); rerr == nil {
log.WithError(err).Errorf("response body [%v]", string(body))
}
} else {
log.WithError(err).Error("no response from websocket dial")
}
return nil, err
}
id := &idlib.TokenId{Token: "mgmt"}
underlayFactory := websockets.NewUnderlayFactory(id, conn, nil)
ch, err := channel.NewSingleChannel("mgmt", underlayFactory, bindHandler, nil)
if err != nil {
return nil, err
}
return ch, nil
}
func (ctx *TestContext) NewClientComponentsWithClientCert(certs []*x509.Certificate, privateKey crypto.PrivateKey) (*resty.Client, *http.Client, *http.Transport) {
clientTransport := ctx.NewTransportWithClientCert(certs, privateKey)
httpClient := ctx.NewHttpClient(clientTransport)
client := resty.NewWithClient(httpClient)
client.SetHostURL("https://" + ctx.ApiHost + EdgeClientApiPath)
return client, httpClient, clientTransport
}
func (ctx *TestContext) StartServer() *ControllerHelper {
return ctx.StartServerFor("testdata/default.db", true)
}
// StartServerWithConfigModifier starts the controller and edge layer, applying
// modifier to the loaded config before the controller is created. This lets a
// single test override settings (e.g., OIDC token durations) without affecting
// other tests that call StartServer with the shared default config.
func (ctx *TestContext) StartServerWithConfigModifier(modifier func(*config.Config)) *ControllerHelper {
return ctx.startServerWith("testdata/default.db", true, modifier, false)
}
func (ctx *TestContext) StartServerFor(testDb string, clean bool) *ControllerHelper {
return ctx.startServerWith(testDb, clean, nil, false)
}
// StartServerRaft starts a controller using a raft/cluster config set (e.g. SingleRaft). It removes
// any stale raft data directory first so the node bootstraps a fresh single-node cluster. The shared
// InitializeDefaultAdmin path bootstraps raft (Dispatcher.Bootstrap) and waits for leadership before
// creating the admin, so no separate bootstrap step is needed here.
func (ctx *TestContext) StartServerRaft() *ControllerHelper {
err := os.RemoveAll(SingleRaftDataDir)
ctx.Req.NoError(err)
return ctx.startServerWith("testdata/single-raft-unused.db", true, nil, true)
}
func (ctx *TestContext) startServerWith(testDb string, clean bool, modifier func(*config.Config), initAdminAfterRun bool) *ControllerHelper {
if ctx.LogLevel != "" {
if level, err := logrus.ParseLevel(ctx.LogLevel); err == nil {
logrus.StandardLogger().SetLevel(level)
}
}
log := pfxlog.Logger()
_ = os.Mkdir("testdata", os.FileMode(0755))
if clean {
err := os.Remove(testDb)
if !os.IsNotExist(err) {
//try again after a small wait
time.Sleep(100 * time.Millisecond)
err := os.Remove(testDb)
if !os.IsNotExist(err) {
ctx.Req.NoError(err)
}
}
}
err := os.Setenv("ZITI_TEST_DB", testDb)
ctx.Req.NoError(err)
log.Info("loading config")
ctrlConfig, err := config.LoadConfig(ctx.controllerConfFile())
ctx.Req.NoError(err)
if modifier != nil {
modifier(ctrlConfig)
}
ctx.ControllerConfig = ctrlConfig
log.Info("creating fabric controller")
ctx.fabricController, err = controller.NewController(ctrlConfig, NewVersionProviderTest())
ctx.Req.NoError(err)
log.Info("creating edge controller")
ctx.EdgeController, err = server.NewController(ctx.fabricController)
ctx.Req.NoError(err)
ctx.EdgeController.Initialize()
// In raft/cluster mode the default admin is created after Run, mirroring `ziti agent cluster
// init` against a live controller: InitializeDefaultAdmin bootstraps raft and triggers self
// registration, which captures API addresses from the running xweb, so it must run once xweb is
// fully up. Non-raft controllers initialize the admin before Run as before.
if !initAdminAfterRun {
ctx.initializeDefaultAdmin()
}
ctx.EdgeController.Run()
go func() {
err = ctx.fabricController.Run()
ctx.Req.NoError(err)
}()
err = ctx.waitForRestAPIPort(time.Minute * 5)
ctx.Req.NoError(err)
if initAdminAfterRun {
ctx.initializeDefaultAdmin()
}
return &ControllerHelper{Controller: ctx.EdgeController}
}
// initializeDefaultAdmin creates the default admin identity. In raft mode this also bootstraps the
// cluster (via Dispatcher.Bootstrap) and waits for leadership before the admin is created.
func (ctx *TestContext) initializeDefaultAdmin() {
err := ctx.EdgeController.AppEnv.Managers.Identity.InitializeDefaultAdmin(ctx.AdminAuthenticator.Username, ctx.AdminAuthenticator.Password, eid.New())
if err != nil {
pfxlog.Logger().WithError(err).Warn("error during initialize admin")
}
logrus.Infof("default admin - username: %v", ctx.AdminAuthenticator.Username)
logrus.Infof("default admin - password: %v", ctx.AdminAuthenticator.Password)
}
func (ctx *TestContext) createAndEnrollEdgeRouter(tunneler bool, roleAttributes ...string) *edgeRouter {
ctx.requireCreateEdgeRouter(tunneler, roleAttributes...)
ctx.requireEnrollEdgeRouter(tunneler, ctx.edgeRouterEntity.id)
return ctx.edgeRouterEntity
}
func (ctx *TestContext) requireCreateEdgeRouter(tunneler bool, roleAttributes ...string) *edgeRouter {
// If an edge router has already been created, delete it and create a new one
if ctx.edgeRouterEntity != nil {
ctx.AdminManagementSession.requireDeleteEntity(ctx.edgeRouterEntity)
ctx.edgeRouterEntity = nil
}
_ = os.MkdirAll("testdata/edge-router", os.FileMode(0755))
if tunneler {
ctx.edgeRouterEntity = ctx.AdminManagementSession.requireNewTunnelerEnabledEdgeRouter(roleAttributes...)
} else {
ctx.edgeRouterEntity = ctx.AdminManagementSession.requireNewEdgeRouter(roleAttributes...)
}
return ctx.edgeRouterEntity
}
func (ctx *TestContext) requireEnrollEdgeRouter(tunneler bool, routerId string) {
jwt := ctx.AdminManagementSession.getEdgeRouterJwt(routerId)
configFile := ctx.configSet.EdgeRouter
if tunneler {
configFile = ctx.configSet.TunnelerRouter
}
routerConfig, err := routerEnv.LoadConfigWithOptions(configFile, false)
ctx.Req.NoError(err)
enroller := enroll.NewRestEnroller(routerConfig)
var keyAlg ziti.KeyAlgVar
_ = keyAlg.Set("RSA")
ctx.Req.NoError(enroller.Enroll([]byte(jwt), true, "", keyAlg))
}
func (ctx *TestContext) createAndEnrollTransitRouter() *transitRouter {
// If a tx router has already been created, delete it and create a new one
if ctx.transitRouterEntity != nil {
ctx.AdminManagementSession.requireDeleteEntity(ctx.transitRouterEntity)
ctx.transitRouterEntity = nil
}
_ = os.MkdirAll("testdata/transit-router", os.FileMode(0755))
ctx.transitRouterEntity = ctx.AdminManagementSession.requireNewTransitRouter()
jwt := ctx.AdminManagementSession.getTransitRouterJwt(ctx.transitRouterEntity.id)
routerConfig, err := routerEnv.LoadConfigWithOptions(ctx.configSet.TransitRouter, false)
ctx.Req.NoError(err)
enroller := enroll.NewRestEnroller(routerConfig)
var keyAlg ziti.KeyAlgVar
_ = keyAlg.Set("RSA")
ctx.Req.NoError(enroller.Enroll([]byte(jwt), true, "", keyAlg))
return ctx.transitRouterEntity
}
func (ctx *TestContext) createEnrollAndStartTransitRouter() {
ctx.createAndEnrollTransitRouter()
ctx.startTransitRouter()
}
func (ctx *TestContext) startTransitRouter() {
routerConfig, err := routerEnv.LoadConfig(ctx.configSet.TransitRouter)
ctx.Req.NoError(err)
newRouter := router.Create(routerConfig, NewVersionProviderTest())
ctx.routers = append(ctx.routers, newRouter)
ctx.Req.NoError(newRouter.Start())
}
func (ctx *TestContext) CreateEnrollAndStartTunnelerEdgeRouter(roleAttributes ...string) {
ctx.shutdownRouters()
ctx.createAndEnrollEdgeRouter(true, roleAttributes...)
ctx.startEdgeRouter(nil)
}
// CreateEnrollAndStartTunnelerEdgeRouterWithCfgTweaks creates a tunneler-enabled edge router
// and allows the caller to modify the router config before startup.
func (ctx *TestContext) CreateEnrollAndStartTunnelerEdgeRouterWithCfgTweaks(cfgTweaks func(*routerEnv.Config), roleAttributes ...string) {
ctx.shutdownRouters()
ctx.createAndEnrollEdgeRouter(true, roleAttributes...)
ctx.startEdgeRouter(cfgTweaks)
}
func (ctx *TestContext) CreateEnrollAndStartEdgeRouter(roleAttributes ...string) *EdgeRouterHelper {
ctx.shutdownRouters()
ctx.createAndEnrollEdgeRouter(false, roleAttributes...)
return ctx.startEdgeRouter(nil)
}
func (ctx *TestContext) CreateEnrollAndStartEdgeRouterWithCfgTweaks(cfgTweaks func(*routerEnv.Config), roleAttributes ...string) *EdgeRouterHelper {
ctx.shutdownRouters()
ctx.createAndEnrollEdgeRouter(false, roleAttributes...)
return ctx.startEdgeRouter(cfgTweaks)
}
func (ctx *TestContext) CreateEnrollAndStartHAEdgeRouter(roleAttributes ...string) *EdgeRouterHelper {
ctx.shutdownRouters()
ctx.createAndEnrollEdgeRouter(false, roleAttributes...)
return ctx.startEdgeRouter(nil)
}
func (ctx *TestContext) startEdgeRouter(cfgTweaks func(*routerEnv.Config)) *EdgeRouterHelper {
configFile := ctx.configSet.EdgeRouter
if ctx.edgeRouterEntity.isTunnelerEnabled {
configFile = ctx.configSet.TunnelerRouter
}
routerCfg, err := routerEnv.LoadConfig(configFile)
ctx.Req.NoError(err)
if cfgTweaks != nil {
cfgTweaks(routerCfg)
}
newRouter := router.Create(routerCfg, NewVersionProviderTest())
ctx.routers = append(ctx.routers, newRouter)
ctx.Req.NoError(newRouter.Start())
return &EdgeRouterHelper{Router: newRouter}
}
func (ctx *TestContext) EnrollIdentity(identityId string) *ziti.Config {
jwt := ctx.AdminManagementSession.getIdentityJwt(identityId)
tkn, _, err := sdkEnroll.ParseToken(jwt)
ctx.Req.NoError(err)
flags := sdkEnroll.EnrollmentFlags{
Token: tkn,
KeyAlg: "RSA",
}
if tkn.EnrollmentMethod == rest_model.EnrollmentCreateMethodUpdb {
flags.Password = uuid.NewString() + "!"
}
conf, err := sdkEnroll.Enroll(flags)
ctx.Req.NoError(err)
return conf
}
func (ctx *TestContext) waitForRestAPIPort(duration time.Duration) error {
return ctx.waitForPort(ctx.ApiHost, duration)
}
func (ctx *TestContext) waitForPort(address string, duration time.Duration) error {
now := time.Now()
endTime := now.Add(duration)
maxWait := duration
for {
conn, err := net.DialTimeout("tcp", address, maxWait)
if err == nil {
_ = conn.Close()
return nil
}
now = time.Now()
if !now.Before(endTime) {
return err
}
maxWait = endTime.Sub(now)
time.Sleep(10 * time.Millisecond)
}
}
func (ctx *TestContext) RequireAdminManagementApiLogin() {
var err error
ctx.AdminManagementSession, err = ctx.AdminAuthenticator.AuthenticateManagementApi(ctx)
ctx.Req.NoError(err)
ctx.RestClients, err = newRestClients(ctx.ApiHost, *ctx.AdminManagementSession.AuthResponse.Token)
ctx.Req.NoError(err)
}
func (ctx *TestContext) RequireAdminClientApiLogin() {
var err error
ctx.AdminClientSession, err = ctx.AdminAuthenticator.AuthenticateClientApi(ctx)
ctx.Req.NoError(err)
}
func (ctx *TestContext) Teardown() {
pfxlog.Logger().Info("tearing down test context")
ctx.shutdownRouters()
ctx.shutdownPeerControllers()
if ctx.EdgeController != nil {
ctx.EdgeController.Shutdown()
ctx.EdgeController = nil
}
if ctx.fabricController != nil {
ctx.fabricController.Shutdown()
ctx.fabricController = nil
}
}
func (ctx *TestContext) newAnonymousClientApiRequest() *resty.Request {
return ctx.DefaultClientApiClient().R().
SetHeader("content-type", "application/json")
}
func (ctx *TestContext) newAnonymousManagementApiRequest() *resty.Request {
return ctx.DefaultManagementApiClient().R().
SetHeader("content-type", "application/json")
}
func (ctx *TestContext) newRequestWithClientCert(certs []*x509.Certificate, privateKey crypto.PrivateKey) *resty.Request {
client, _, _ := ctx.NewClientComponentsWithClientCert(certs, privateKey)
return client.R().
SetHeader("content-type", "application/json")
}
// newRequestWithTlsCerts creates an anonymous resty request that presents the given TLS
// certificates for client authentication. Useful for testing cert proof-of-possession
// with OIDC bearer tokens.
func (ctx *TestContext) newRequestWithTlsCerts(tlsCerts []cryptoTls.Certificate) *resty.Request {
transport := ctx.NewTransport()
transport.TLSClientConfig.Certificates = tlsCerts
client := resty.NewWithClient(ctx.NewHttpClient(transport))
client.SetHostURL("https://" + ctx.ApiHost + EdgeClientApiPath)
return client.R().
SetHeader("content-type", "application/json")
}
func (ctx *TestContext) completeUpdbEnrollment(identityId string, password string) {
result := ctx.AdminManagementSession.requireQuery(fmt.Sprintf("identities/%v", identityId))
path := result.Search(path("data.enrollment.updb.token")...)
ctx.Req.NotNil(path)
str, ok := path.Data().(string)
ctx.Req.True(ok)
enrollBody := gabs.New()
ctx.setJsonValue(enrollBody, password, "password")
resp, err := ctx.newAnonymousClientApiRequest().
SetBody(enrollBody.String()).
Post("enroll?token=" + str)
ctx.Req.NoError(err)
ctx.logJson(resp.Body())
ctx.Req.Equal(http.StatusOK, resp.StatusCode())
}
func (ctx *TestContext) completeOttCaEnrollment(certAuth *certAuthenticator) {
trans := ctx.NewTransport()
trans.TLSClientConfig.Certificates = certAuth.TLSCertificates()
client := resty.NewWithClient(ctx.NewHttpClient(trans))
client.SetHostURL("https://" + ctx.ApiHost + EdgeClientApiPath)
resp, err := client.NewRequest().
Post("enroll?method=ca")
ctx.Req.NoError(err)
ctx.logJson(resp.Body())
ctx.Req.Equal(http.StatusOK, resp.StatusCode())
}
func (ctx *TestContext) completeCaAutoEnrollmentWithName(certAuth *certAuthenticator, name string) {
trans := ctx.NewTransport()
trans.TLSClientConfig.Certificates = certAuth.TLSCertificates()
client := resty.NewWithClient(ctx.NewHttpClient(trans))
client.SetHostURL("https://" + ctx.ApiHost + EdgeClientApiPath)
body := gabs.New()
_, _ = body.SetP(name, "name")
resp, err := client.NewRequest().
SetHeader("content-type", "application/json").
SetBody(body.String()).
Post("enroll?method=ca")
ctx.Req.NoError(err)
ctx.logJson(resp.Body())
ctx.Req.Equal(http.StatusOK, resp.StatusCode())
}
func (ctx *TestContext) completeOttEnrollment(identityId string) *certAuthenticator {
return ctx.completeOttEnrollmentAtApiHost(identityId, ctx.ApiHost)
}
// completeOttEnrollmentAtApiHost completes an identity's OTT enrollment against the client API
// at the given host, so cluster tests can enroll via a specific controller.
func (ctx *TestContext) completeOttEnrollmentAtApiHost(identityId string, apiHost string) *certAuthenticator {
result := ctx.AdminManagementSession.requireQuery(fmt.Sprintf("identities/%v", identityId))
tokenValue := result.Path("data.enrollment.ott.token")
ctx.Req.NotNil(tokenValue)
token, ok := tokenValue.Data().(string)
ctx.Req.True(ok)
privateKey, err := ecdsa.GenerateKey(elliptic.P384(), rand.Reader)
ctx.Req.NoError(err)
request, err := certtools.NewCertRequest(map[string]string{
"C": "US", "O": "NetFoundry-API-Test", "CN": identityId,
}, nil)
ctx.Req.NoError(err)
csr, err := x509.CreateCertificateRequest(rand.Reader, request, privateKey)
ctx.Req.NoError(err)
csrPem := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE REQUEST", Bytes: csr})
resp, err := ctx.NewRestClientWithDefaults().R().
SetHeader("content-type", "application/x-pem-file").
SetHeader("accept", "application/json").
SetBody(csrPem).
Post("https://" + apiHost + EdgeClientApiPath + "/enroll?token=" + token)
ctx.Req.NoError(err)
ctx.logJson(resp.Body())
ctx.Req.Equal(http.StatusOK, resp.StatusCode())
envelope := &rest_model.EnrollmentCertsEnvelope{}
err = json.Unmarshal(resp.Body(), envelope)
ctx.Req.NoError(err)
certs := nfPem.PemStringToCertificates(envelope.Data.Cert)
ctx.Req.NotEmpty(certs)
return &certAuthenticator{
certs: certs,
key: privateKey,
certPem: envelope.Data.Cert,
}
}
func (ctx *TestContext) validateDateFieldsForCreate(start time.Time, jsonEntity *gabs.Container) time.Time {
// we lose a little time resolution, so if it's in the same millisecond, it's ok
start = start.Add(-time.Millisecond)
now := time.Now().Add(time.Millisecond)
createdAt, updatedAt := ctx.getEntityDates(jsonEntity)
ctx.Req.Equal(createdAt, updatedAt)
ctx.Req.True(start.Before(createdAt) || start.Equal(createdAt), "%v should be before or equal to %v", start, createdAt)
ctx.Req.True(now.After(createdAt) || now.Equal(createdAt), "%v should be after or equal to %v", now, createdAt)
return createdAt
}
func (ctx *TestContext) newPostureCheckProcessMulti(semantic rest_model.Semantic, processes []*rest_model.ProcessMulti, roleAttributes []string) *rest_model.PostureCheckProcessMultiCreate {
check := &rest_model.PostureCheckProcessMultiCreate{
Processes: processes,
Semantic: &semantic,
}
attributes := rest_model.Attributes(roleAttributes)
check.SetRoleAttributes(&attributes)
name := uuid.New().String()
check.SetName(&name)
check.SetTypeID(rest_model.PostureCheckTypePROCESSMULTI)
return check
}
func (ctx *TestContext) newPostureCheckDomain(domains []string, roleAttributes []string) *postureCheckDomain {
return &postureCheckDomain{
postureCheck: postureCheck{
name: eid.New(),
typeId: "DOMAIN",
roleAttributes: roleAttributes,
tags: nil,
},
domains: domains,
}
}
func (ctx *TestContext) newService(roleAttributes, configs []string) *service {
return &service{
Name: eid.New(),
terminatorStrategy: xt_smartrouting.Name,
roleAttributes: roleAttributes,
configs: configs,
encryptionRequired: false,
tags: nil,
}
}
func (ctx *TestContext) newTerminator(serviceId, routerId, binding, address string) *terminator {
return &terminator{
serviceId: serviceId,
routerId: routerId,
binding: binding,
address: address,
cost: 0,
precedence: "default",
tags: nil,
}
}
func (ctx *TestContext) newConfig(configType string, data map[string]interface{}) *Config {
return &Config{
Name: eid.New(),
ConfigTypeId: configType,
Data: data,
Tags: nil,
}
}
func (ctx *TestContext) newConfigType() *configType {
return &configType{
Name: eid.New(),
Target: util.Ptr("service"),
Tags: nil,
}
}
func (ctx *TestContext) getEntityDates(jsonEntity *gabs.Container) (time.Time, time.Time) {
createdAtStr := jsonEntity.S("createdAt").Data().(string)
updatedAtStr := jsonEntity.S("updatedAt").Data().(string)
ctx.Req.NotNil(createdAtStr)
ctx.Req.NotNil(updatedAtStr)
createdAt, err := time.Parse(time.RFC3339, createdAtStr)
ctx.Req.NoError(err)
updatedAt, err := time.Parse(time.RFC3339, updatedAtStr)
ctx.Req.NoError(err)
return createdAt, updatedAt
}
func (ctx *TestContext) validateDateFieldsForUpdate(start time.Time, origCreatedAt time.Time, jsonEntity *gabs.Container) time.Time {
// we lose a little time resolution, so if it's in the same millisecond, it's ok
start = start.Add(-time.Millisecond)
now := time.Now().Add(time.Millisecond)
createdAt, updatedAt := ctx.getEntityDates(jsonEntity)
ctx.Req.Equal(origCreatedAt, createdAt)
ctx.Req.True(createdAt.Before(updatedAt))
ctx.Req.True(start.Before(updatedAt) || start.Equal(updatedAt))
ctx.Req.True(now.After(updatedAt) || now.Equal(updatedAt))
return createdAt
}
func (ctx *TestContext) validateEntity(entity entity, jsonEntity *gabs.Container) *gabs.Container {
entity.validate(ctx, jsonEntity)
return jsonEntity
}
func (ctx *TestContext) requireEntityNotEnrolled(name string, entity *gabs.Container) {
fingerprint := entity.Path("fingerprint").Data()
ctx.Req.Nil(fingerprint, "expected "+name+" with isVerified=false to have an empty fingerprint")
token, ok := entity.Path("enrollmentToken").Data().(string)
ctx.Req.True(ok, "expected "+name+" with isVerified=false to have an enrollment token, could not cast")
ctx.Req.NotEmpty(token, "expected "+name+" with isVerified=false to have an enrollment token, was empty")
jwt, ok := entity.Path("enrollmentJwt").Data().(string)
ctx.Req.True(ok, "expected "+name+" with isVerified=false to have an enrollment jwt, could not cast")
ctx.Req.NotEmpty(jwt, "expected "+name+" with isVerified=false to have an enrollment jwt, was empty")
createdAtStr, ok := entity.Path("enrollmentCreatedAt").Data().(string)
ctx.Req.True(ok, "expected "+name+" with isVerified=false to have an enrollment created at date, could not cast")
ctx.Req.NotEmpty(createdAtStr, "expected "+name+" with isVerified=false to have an enrollment created at date string, was empty")
createdAt, err := time.Parse(time.RFC3339, createdAtStr)
ctx.Req.NoError(err, "expected "+name+" with isVerified=false to have a parsable created at date time string")
ctx.Req.NotEmpty(createdAt, "expected "+name+" with isVerified=false to have an enrollment created at date, was empty")
expiresAtStr, ok := entity.Path("enrollmentExpiresAt").Data().(string)
ctx.Req.True(ok, "expected "+name+" with isVerified=false to have an enrollment expires at date, could not cast")
ctx.Req.NotEmpty(expiresAtStr, "expected "+name+" with isVerified=false to have an enrollment expires at date string, was empty")
expiresAt, err := time.Parse(time.RFC3339, expiresAtStr)
ctx.Req.NoError(err, "expected "+name+" with isVerified=false to have a parsable expires at date time string")
ctx.Req.True(ok, "expected "+name+" with isVerified=false to have an enrollment expires at date, could not cast")
ctx.Req.NotEmpty(expiresAt, "expected "+name+" with isVerified=false to have an enrollment expires at date, was empty")
ctx.Req.True(expiresAt.After(createdAt), "expected "+name+" with isVerified=false to have an enrollment expires at date after the created at date")
}
func (ctx *TestContext) requireEntityEnrolled(name string, entity *gabs.Container) {
fingerprint, ok := entity.Path("fingerprint").Data().(string)
ctx.Req.True(ok, "expected "+name+" with isVerified=true to have a fingerprint, could not cast")
ctx.Req.NotEmpty(fingerprint, "expected "+name+" with isVerified=true to have a fingerprint, was empty")
ctx.Req.False(strings.Contains(fingerprint, ":"), "fingerprint should not contain colons")
ctx.Req.False(strings.ToLower(fingerprint) != fingerprint, "fingerprint should not contain uppercase characters")
token := entity.Path("enrollmentToken").Data()
ctx.Req.Nil(token, "expected "+name+" with isVerified=true to have an nil enrollment token")
jwt := entity.Path("enrollmentJwt").Data()
ctx.Req.Nil(jwt, "expected "+name+" with isVerified=true to have an nil enrollment jwt")
createdAt := entity.Path("enrollmentCreatedAt").Data()
ctx.Req.Nil(createdAt, "expected "+name+" with isVerified=true to have an nil enrollment created at date")
expiresAt := entity.Path("enrollmentExpiresAt").Data()
ctx.Req.Nil(expiresAt, "expected "+name+" with isVerified=true to have an nil enrollment expires at date")
}
func (ctx *TestContext) WrapNetConn(conn edge.Conn, err error) *TestConn {
ctx.Req.NoError(err)
return &TestConn{
Conn: conn,
ctx: ctx,
}
}
func (ctx *TestContext) WrapConn(conn edge.Conn, err error) *TestConn {
ctx.Req.NoError(err)
return &TestConn{
Conn: conn,
ctx: ctx,
}
}
func (ctx *TestContext) shutdownRouters() {
for _, r := range ctx.routers {
ctx.Req.NoError(r.Shutdown())
}
ctx.routers = nil
}
func (ctx *TestContext) NewAdminCredentials() *edgeApis.UpdbCredentials {
creds := edgeApis.NewUpdbCredentials(ctx.AdminAuthenticator.Username, ctx.AdminAuthenticator.Password)
creds.CaPool = ctx.ControllerCaPool()
return creds
}
func (ctx *TestContext) EnrollFabricRouter(id string, name string, certFile string) {
cert, err := certtools.LoadCertFromFile(certFile)
ctx.Req.NoError(err)
fingerprint := fmt.Sprintf("%x", sha1.Sum(cert[0].Raw))
timeoutContext, cancelF := context.WithTimeout(context.Background(), 10*time.Second)
defer cancelF()
createRouterParams := &restClientRouter.CreateRouterParams{
Router: &fabricRestModel.RouterCreate{
Cost: util.Ptr(int64(0)),
Fingerprint: &fingerprint,
ID: &id,
Name: &name,
NoTraversal: util.Ptr(false),
},
Context: timeoutContext,
}
_, err = ctx.RestClients.Fabric.Router.CreateRouter(createRouterParams, nil)
if err != nil {
js, _ := json.MarshalIndent(err, "", " ")
fmt.Println(string(js))
}
ctx.Req.NoError(err)
}
func (ctx *TestContext) startFabricRouter(index uint8) *router.Router {
routerCfg, err := routerEnv.LoadConfig(ctx.configSet.FabricRouters[index-1])
ctx.Req.NoError(err)
r := router.Create(routerCfg, versions.NewDefaultVersionProvider())
ctx.Req.NoError(r.Start())
ctx.routers = append(ctx.routers, r)
return r
}
func (ctx *TestContext) waitForPortClose(address string, duration time.Duration) error {
now := time.Now()
endTime := now.Add(duration)
maxWait := duration
for {
conn, err := net.DialTimeout("tcp", address, maxWait)
if err != nil {
return nil
}
_ = conn.Close()
now = time.Now()
if !now.Before(endTime) {
return err
}
maxWait = endTime.Sub(now)
time.Sleep(10 * time.Millisecond)
}
}
type TestConn struct {
edge.Conn
ctx *TestContext
}
func (conn *TestConn) WriteString(val string, timeout time.Duration) {
conn.ctx.Req.NoError(conn.SetWriteDeadline(time.Now().Add(timeout)))
defer func() { _ = conn.SetWriteDeadline(time.Time{}) }()
buf := []byte(val)
n, err := conn.Write(buf)
conn.ctx.Req.NoError(err)
conn.ctx.Req.Equal(n, len(buf))
}
func (conn *TestConn) ReadString(maxSize int, timeout time.Duration) string {
conn.ctx.Req.NoError(conn.SetReadDeadline(time.Now().Add(timeout)))
defer func() { _ = conn.SetReadDeadline(time.Time{}) }()
buf := make([]byte, maxSize)
n, err := conn.Read(buf)
conn.ctx.Req.NoError(err, "read timeout on connId=%v", conn.Id())
return string(buf[:n])
}
func (conn *TestConn) ReadExpected(expected string, timeout time.Duration) {
val := conn.ReadString(len(expected)+1, timeout)
conn.ctx.Req.Equal(expected, val, "read failure on connId=%v", conn.Id())
}
func (conn *TestConn) RequireClose() {
conn.ctx.Req.NoError(conn.Close())
}
var testServerCounter uint64
func newTestServer(listener edge.Listener, dispatcher func(conn *testServerConn) error) *testServer {
idx := atomic.AddUint64(&testServerCounter, 1)
return &testServer{
idx: idx,
listener: listener,
errorC: make(chan error, 10),
msgCount: 0,
dispatcher: dispatcher,
waiter: &sync.WaitGroup{},
}
}
type testServer struct {
idx uint64
listener edge.Listener
errorC chan error
msgCount uint32
dispatcher func(conn *testServerConn) error
waiter *sync.WaitGroup
connIdGen uint32
}
func (server *testServer) waitForDone(ctx *TestContext, timeout time.Duration) {
select {
case err, ok := <-server.errorC:
if ok {
ctx.Req.NoError(err)
}
case <-time.After(timeout):
ctx.Req.Fail("wait for done on test server timed out")
}
}
func (server *testServer) start() {
go server.acceptLoop()
}
func (server *testServer) close() error {
return server.listener.Close()
}
func (server *testServer) acceptLoop() {
var err error
for !server.listener.IsClosed() {
var conn net.Conn
conn, err = server.listener.Accept()
if conn != nil {
server.waiter.Add(1)
connId := atomic.AddUint32(&server.connIdGen, 1)
go server.dispatch(&testServerConn{id: connId, Conn: conn, server: server})
} else {
break
}
}
// if the listener is closed, assume this error is just letting us know the listener was closed
if !server.listener.IsClosed() {
if err != nil {
server.errorC <- err
}
}
waitDone := make(chan struct{})
go func() {
server.waiter.Wait()
close(waitDone)
}()
select {
case _, ok := <-waitDone:
if !ok {
pfxlog.Logger().Debugf("all connections closed")
}
case <-time.After(10 * time.Second):
pfxlog.Logger().Warn("timed out waiting for all connections to close")
}
close(server.errorC)
pfxlog.Logger().Debugf("%v: service exiting", server.idx)
}
func (server *testServer) dispatch(conn *testServerConn) {
defer func() {
server.waiter.Done()
}()
log := pfxlog.Logger()
defer func() {
val := recover()
if val != nil {
if err, ok := val.(error); ok {
log.WithError(err).Error("panic from server.dispatch")
server.errorC <- err
}
}
}()
defer func() {
conn.RequireClose()
}()
log.Debugf("beginning dispatch to conn %v-%v", conn.server.idx, conn.id)
err := server.dispatcher(conn)
log.Debugf("finished dispatch to conn %v-%v", conn.server.idx, conn.id)
if err != nil {
log.WithError(err).Error("failure from server.dispatch")
server.errorC <- err
}
}
type testServerConn struct {
id uint32
net.Conn
server *testServer
}
func (conn *testServerConn) WriteString(val string, timeout time.Duration) {
err := conn.SetWriteDeadline(time.Now().Add(timeout))
if err != nil {
panic(err)
}
defer func() { _ = conn.SetWriteDeadline(time.Time{}) }()
buf := []byte(val)
n, err := conn.Write(buf)
if err != nil {
panic(fmt.Errorf("conn %v-%v timed out trying to write string %v (%w)", conn.server.idx, conn.id, val, err))
}
if n != len(buf) {
panic(errors.Errorf("conn %v-%v expected to write %v bytes, but only wrote %v", conn.server.idx, conn.id, len(buf), n))
}
}
func (conn *testServerConn) ReadString(maxSize int, timeout time.Duration) (string, bool) {
err := conn.SetReadDeadline(time.Now().Add(timeout))
if err != nil {
panic(err)
}
defer func() { _ = conn.SetReadDeadline(time.Time{}) }()
buf := make([]byte, maxSize)
n, err := conn.Read(buf)
if err != nil {
if err == io.EOF {
return "", true
}
panic(fmt.Errorf("conn %v-%v timed out trying to read (%w)", conn.server.idx, conn.id, err))
}
return string(buf[:n]), false
}
func (conn *testServerConn) ReadExpected(expected string, timeout time.Duration) {
val, eof := conn.ReadString(len(expected)+1, timeout)
if eof {
panic(errors.Errorf("expected to read string '%v', but got EOF", expected))
}
if val != expected {
panic(errors.Errorf("expected to read string '%v', but got '%v'", expected, val))
}
}
func (conn *testServerConn) RequireClose() {
err := conn.Close()
if err != nil {
panic(err)
}
}
type VersionProviderTest struct {
}
func (v VersionProviderTest) Branch() string {
return "local"
}
func (v VersionProviderTest) EncoderDecoder() versions.VersionEncDec {
return &versions.StdVersionEncDec
}
func (v VersionProviderTest) Version() string {
return "v0.0.0"
}
func (v VersionProviderTest) BuildDate() string {
return time.Now().String()
}
func (v VersionProviderTest) Revision() string {
return ""
}
func (v VersionProviderTest) AsVersionInfo() *versions.VersionInfo {
return &versions.VersionInfo{
Version: v.Version(),
Revision: v.Revision(),
BuildDate: v.BuildDate(),
OS: runtime.GOOS,
Arch: runtime.GOARCH,
}
}
func NewVersionProviderTest() versions.VersionProvider {
return &VersionProviderTest{}
}