diff --git a/api/api.go b/api/api.go index 63f69199..2de3ba22 100644 --- a/api/api.go +++ b/api/api.go @@ -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), diff --git a/api/handlers/network_client_handlers.go b/api/handlers/network_client_handlers.go index 69aa269d..daf96f8c 100644 --- a/api/handlers/network_client_handlers.go +++ b/api/handlers/network_client_handlers.go @@ -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) } diff --git a/model/network_client_model.go b/model/network_client_model.go index f96354f9..549e841c 100644 --- a/model/network_client_model.go +++ b/model/network_client_model.go @@ -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"` } diff --git a/model/network_client_model_test.go b/model/network_client_model_test.go index 40756e00..371950ef 100644 --- a/model/network_client_model_test.go +++ b/model/network_client_model_test.go @@ -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) { @@ -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) @@ -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{}) @@ -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)) @@ -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) @@ -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) + }) + }) +} diff --git a/test_util.go b/test_util.go index d4f79912..57ee0083 100644 --- a/test_util.go +++ b/test_util.go @@ -220,6 +220,7 @@ func (self *TestEnv) setup() func() { OWNER=%s ENCODING=UTF8 LOCALE='en_US.UTF-8' + TEMPLATE='template0' `, testPgDbName, pg["user"],