Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions api/api.go
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,7 @@ func Routes() []*router.Route {
router.NewRoute("POST", "/auth/upgrade-guest-existing", handlers.UpgradeGuestExisting),
router.NewRoute("POST", "/network/auth-client", handlers.AuthNetworkClient),
router.NewRoute("POST", "/network/remove-client", handlers.RemoveNetworkClient),
router.NewRoute("POST", "/network/remove-clients", handlers.RemoveNetworkClients),
router.NewRoute("GET", "/network/clients", handlers.NetworkClients),
router.NewRoute("GET", "/network/provider-locations", handlers.NetworkGetProviderLocations),
router.NewRoute("POST", "/network/find-provider-locations", handlers.NetworkFindProviderLocations),
Expand Down
4 changes: 4 additions & 0 deletions api/handlers/network_client_handlers.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,10 @@ func RemoveNetworkClient(w http.ResponseWriter, r *http.Request) {
router.WrapWithInputRequireAuth(model.RemoveNetworkClient, w, r)
}

func RemoveNetworkClients(w http.ResponseWriter, r *http.Request) {
router.WrapWithInputRequireAuth(model.RemoveNetworkClients, w, r)
}

func RemoveNetwork(w http.ResponseWriter, r *http.Request) {
router.WrapRequireAuth(controller.NetworkRemove, w, r)
}
Expand Down
37 changes: 37 additions & 0 deletions model/network_client_model.go
Original file line number Diff line number Diff line change
Expand Up @@ -445,6 +445,43 @@ func RemoveNetworkClient(
return removeClientResult, removeClientErr
}

type RemoveNetworkClientsArgs struct {
ClientIds []server.Id `json:"client_ids"`
}

type RemoveNetworkClientsResult struct{}

func RemoveNetworkClients(
removeClients *RemoveNetworkClientsArgs,
session *session.ClientSession,
) (*RemoveNetworkClientsResult, error) {
var removeClientsResult *RemoveNetworkClientsResult
if len(removeClients.ClientIds) == 0 {
return &RemoveNetworkClientsResult{}, nil
}

server.Tx(session.Ctx, func(tx server.PgTx) {
_, err := tx.Exec(
session.Ctx,
`
UPDATE network_client
SET
active = false
WHERE
client_id = ANY($1) AND
network_id = $2
`,
removeClients.ClientIds,
session.ByJwt.NetworkId,
)
server.Raise(err)

removeClientsResult = &RemoveNetworkClientsResult{}
})

return removeClientsResult, nil
}

type NetworkClientsResult struct {
Clients []*NetworkClientInfo `json:"clients"`
}
Expand Down
81 changes: 77 additions & 4 deletions model/network_client_model_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,8 @@ import (
"github.com/go-playground/assert/v2"

"github.com/urnetwork/server"
"github.com/urnetwork/server/jwt"
"github.com/urnetwork/server/session"
)

func TestNetworkClientHandlerLifecycle(t *testing.T) {
Expand Down Expand Up @@ -138,7 +140,6 @@ func TestSetProvide(t *testing.T) {

ctx, cancel := context.WithCancel(context.Background())
defer cancel()

newSecretKeys := func() map[ProvideMode][]byte {
k := make([]byte, 32)
mathrand.Read(k)
Expand All @@ -151,7 +152,6 @@ func TestSetProvide(t *testing.T) {
secretKeys := newSecretKeys()

startTime := server.NowUtc()

changeCount, provideModes := GetProvideKeyChanges(ctx, clientId, startTime)
assert.Equal(t, changeCount, 0)
assert.Equal(t, provideModes, map[ProvideMode]bool{})
Expand Down Expand Up @@ -214,7 +214,6 @@ func TestGetProvideFallsBackToDb(t *testing.T) {
}

SetProvide(ctx, clientId, secretKeys)

// drop the cache so the reads are forced through the db fallback
server.Redis(ctx, func(r server.RedisClient) {
r.Del(ctx, provideModesKey(clientId))
Expand All @@ -238,7 +237,6 @@ func TestGetProvideModesNotSet(t *testing.T) {
ctx := context.Background()

clientId := server.NewId()

// a client that never provided returns an empty set and no error
provideModes, err := GetProvideModes(ctx, clientId)
assert.Equal(t, err, nil)
Expand Down Expand Up @@ -385,3 +383,78 @@ func TestMigrateProvideMode(t *testing.T) {
})
})
}

func TestRemoveNetworkClients(t *testing.T) {
server.DefaultTestEnv().Run(t, func(t testing.TB) {
ctx := context.Background()
networkId := server.NewId()

// Create a mock session
sess := &session.ClientSession{
Ctx: ctx,
ByJwt: &jwt.ByJwt{
NetworkId: networkId,
},
}

// Generate random IDs to test the ANY($1) binding
clientIds := []server.Id{server.NewId(), server.NewId(), server.NewId()}

args := &RemoveNetworkClientsArgs{
ClientIds: clientIds,
}

// This will panic if the driver fails to cast []server.Id to uuid[]
_, err := RemoveNetworkClients(args, sess)

// Assert that the function ran without returning an error
assert.Equal(t, err, nil)
})
}

func TestRemoveNetworkClientsOnlyRemovesOwnNetwork(t *testing.T) {
server.DefaultTestEnv().Run(t, func(t testing.TB) {
ctx := context.Background()

ourNetworkId := server.NewId()
otherNetworkId := server.NewId()

// Seed clients in two networks. Use SQL so we can explicitly set network_id
// without depending on unrelated model code paths.
ourClientId := server.NewId()
otherClientId := server.NewId()
server.Tx(ctx, func(tx server.PgTx) {
_, err := tx.Exec(ctx, `
INSERT INTO network_client (client_id, network_id, active)
VALUES
($1, $2, true),
($3, $4, true)
`, ourClientId, ourNetworkId, otherClientId, otherNetworkId)
server.Raise(err)
})

sess := &session.ClientSession{
Ctx: ctx,
ByJwt: &jwt.ByJwt{
NetworkId: ourNetworkId,
},
}

_, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ClientIds: []server.Id{ourClientId, otherClientId}}, sess)
assert.Equal(t, err, nil)

// our client is disabled
server.Tx(ctx, func(tx server.PgTx) {
var ourActive bool
err := tx.QueryRow(ctx, `SELECT active FROM network_client WHERE client_id = $1`, ourClientId).Scan(&ourActive)
assert.Equal(t, err, nil)
assert.Equal(t, ourActive, false)

// other-network client must stay active
var otherActive bool
err = tx.QueryRow(ctx, `SELECT active FROM network_client WHERE client_id = $1`, otherClientId).Scan(&otherActive)
assert.Equal(t, err, nil)
assert.Equal(t, otherActive, true)
})
})
}
1 change: 1 addition & 0 deletions test_util.go
Original file line number Diff line number Diff line change
Expand Up @@ -220,6 +220,7 @@ func (self *TestEnv) setup() func() {
OWNER=%s
ENCODING=UTF8
LOCALE='en_US.UTF-8'
TEMPLATE='template0'
`,
testPgDbName,
pg["user"],
Expand Down