diff --git a/api/api.go b/api/api.go index fef528b8..bfce464d 100644 --- a/api/api.go +++ b/api/api.go @@ -47,10 +47,13 @@ func Routes() []*router.Route { router.NewRoute("POST", "/auth/network-delete", handlers.RemoveNetwork), router.NewRoute("POST", "/auth/code-create", handlers.AuthCodeCreate), router.NewRoute("POST", "/auth/code-login", handlers.AuthCodeLogin), - router.NewRoute("POST", "/auth/upgrade-guest", handlers.UpgradeGuest), - router.NewRoute("POST", "/auth/upgrade-guest-existing", handlers.UpgradeGuestExisting), + router.NewRoute("POST", "/auth/add-auth", handlers.AuthAdd), + router.NewRoute("POST", "/auth/remove-auth", handlers.AuthRemove), + router.NewRoute("POST", "/auth/regenerate-seedphrase", handlers.AuthRegenerateSeedphrase), + router.NewRoute("POST", "/auth/generate-seedphrase", handlers.AuthGenerateSeedphrase), 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/peers", handlers.NetworkPeers), router.NewRoute("GET", "/network/provider-locations", handlers.NetworkGetProviderLocations), @@ -132,6 +135,8 @@ func Routes() []*router.Route { router.NewRoute("GET", "/account/referral-network", handlers.GetReferralNetwork), router.NewRoute("GET", "/account/unlink-referral-network", handlers.UnlinkReferralNetwork), router.NewRoute("POST", "/account/set-referral", handlers.SetNetworkReferral), + router.NewRoute("POST", "/account/change-name", handlers.ChangeNetworkName), + router.NewRoute("POST", "/account/claim-name", handlers.ClaimNetworkName), router.NewRoute("GET", "/account/points", handlers.GetAccountPoints), router.NewRoute("GET", "/account/balance-codes", handlers.GetNetworkRedeemedBalanceCodes), router.NewRoute("POST", "/referral-code/validate", handlers.ValidateReferralCode), diff --git a/api/handlers/account_handlers.go b/api/handlers/account_handlers.go index 008e2366..e9449c28 100644 --- a/api/handlers/account_handlers.go +++ b/api/handlers/account_handlers.go @@ -39,3 +39,11 @@ func UnlinkReferralNetwork(w http.ResponseWriter, r *http.Request) { func GetNetworkRedeemedBalanceCodes(w http.ResponseWriter, r *http.Request) { router.WrapRequireAuth(controller.GetNetworkRedeemedBalanceCodes, w, r) } + +func ChangeNetworkName(w http.ResponseWriter, r *http.Request) { + router.WrapWithInputRequireAuth(controller.ChangeNetworkName, w, r) +} + +func ClaimNetworkName(w http.ResponseWriter, r *http.Request) { + router.WrapWithInputRequireAuth(controller.ClaimNetworkName, w, r) +} diff --git a/api/handlers/auth_handlers.go b/api/handlers/auth_handlers.go index e84156f0..b788f941 100644 --- a/api/handlers/auth_handlers.go +++ b/api/handlers/auth_handlers.go @@ -241,3 +241,11 @@ func AuthRefreshToken(w http.ResponseWriter, r *http.Request) { // router.WrapWithInputNoAuth(controller.AuthRefreshToken, w, r) router.WrapRequireAuth(controller.RefreshToken, w, r) } + +func AuthAdd(w http.ResponseWriter, r *http.Request) { + router.WrapWithInputRequireAuth(controller.AddAuth, w, r) +} + +func AuthRemove(w http.ResponseWriter, r *http.Request) { + router.WrapWithInputRequireAuth(controller.RemoveAuth, w, r) +} diff --git a/api/handlers/network_client_handlers.go b/api/handlers/network_client_handlers.go index 59ae5b9c..a29942af 100644 --- a/api/handlers/network_client_handlers.go +++ b/api/handlers/network_client_handlers.go @@ -38,6 +38,11 @@ func RemoveNetworkClient(w http.ResponseWriter, r *http.Request) { router.WrapWithInputRequireAuth(model.RemoveNetworkClient, w, r) } +func RemoveNetworkClients(w http.ResponseWriter, r *http.Request) { + r.Body = http.MaxBytesReader(w, r.Body, 100<<20) // 100 MB cap (allows ~2M UUIDs worst case) + router.WrapWithInputRequireAuth(model.RemoveNetworkClients, w, r) +} + func RemoveNetwork(w http.ResponseWriter, r *http.Request) { router.WrapRequireAuth(controller.NetworkRemove, w, r) } diff --git a/api/handlers/network_user_handlers.go b/api/handlers/network_user_handlers.go index e82d30ed..141979da 100644 --- a/api/handlers/network_user_handlers.go +++ b/api/handlers/network_user_handlers.go @@ -10,11 +10,3 @@ import ( func GetNetworkUser(w http.ResponseWriter, r *http.Request) { router.WrapRequireAuth(controller.GetNetworkUser, w, r) } - -func UpgradeGuest(w http.ResponseWriter, r *http.Request) { - router.WrapWithInputRequireAuth(controller.UpgradeFromGuest, w, r) -} - -func UpgradeGuestExisting(w http.ResponseWriter, r *http.Request) { - router.WrapWithInputRequireAuth(controller.UpgradeFromGuestExisting, w, r) -} diff --git a/api/handlers/seedphrase_handlers.go b/api/handlers/seedphrase_handlers.go new file mode 100644 index 00000000..71522582 --- /dev/null +++ b/api/handlers/seedphrase_handlers.go @@ -0,0 +1,16 @@ +package handlers + +import ( + "net/http" + + "github.com/urnetwork/server/controller" + "github.com/urnetwork/server/router" +) + +func AuthRegenerateSeedphrase(w http.ResponseWriter, r *http.Request) { + router.WrapWithInputRequireAuth(controller.RegenerateSeedphrase, w, r) +} + +func AuthGenerateSeedphrase(w http.ResponseWriter, r *http.Request) { + router.WrapWithInputRequireAuth(controller.GenerateSeedphrase, w, r) +} diff --git a/api/spec_conformance_test.go b/api/spec_conformance_test.go index fd79a904..37a24267 100644 --- a/api/spec_conformance_test.go +++ b/api/spec_conformance_test.go @@ -78,8 +78,10 @@ func registry() []specEndpoint { {"POST", "/auth/network-delete", nil, rt(controller.NetworkRemoveResult{})}, {"POST", "/auth/code-create", rt(model.AuthCodeCreateArgs{}), rt(model.AuthCodeCreateResult{})}, {"POST", "/auth/code-login", rt(model.AuthCodeLoginArgs{}), rt(model.AuthCodeLoginResult{})}, - {"POST", "/auth/upgrade-guest", rt(model.UpgradeGuestArgs{}), rt(model.UpgradeGuestResult{})}, - {"POST", "/auth/upgrade-guest-existing", rt(model.UpgradeGuestExistingArgs{}), rt(model.UpgradeGuestExistingResult{})}, + {"POST", "/auth/add-auth", rt(controller.AddAuthArgs{}), rt(controller.AddAuthResult{})}, + {"POST", "/auth/remove-auth", rt(controller.RemoveAuthArgs{}), rt(controller.RemoveAuthResult{})}, + {"POST", "/auth/regenerate-seedphrase", rt(controller.RegenerateSeedphraseArgs{}), rt(controller.RegenerateSeedphraseResult{})}, + {"POST", "/auth/generate-seedphrase", rt(controller.GenerateSeedphraseArgs{}), rt(controller.GenerateSeedphraseResult{})}, {"POST", "/network/auth-client", rt(model.AuthNetworkClientArgs{}), rt(model.AuthNetworkClientResult{})}, {"POST", "/network/remove-client", rt(model.RemoveNetworkClientArgs{}), rt(model.RemoveNetworkClientResult{})}, @@ -146,6 +148,8 @@ func registry() []specEndpoint { {"POST", "/account/wallets/verify-seeker", rt(controller.VerifySeekerNftHolderArgs{}), rt(controller.VerifySeekerNftHolderResult{})}, {"GET", "/account/balance-codes", nil, rt(controller.GetNetworkRedeemedBalanceCodesResult{})}, {"GET", "/account/referral-code", nil, rt(controller.NetworkReferralResult{})}, + {"POST", "/account/change-name", rt(controller.ChangeNetworkNameArgs{}), rt(controller.ChangeNetworkNameResult{})}, + {"POST", "/account/claim-name", rt(controller.ClaimNetworkNameArgs{}), rt(controller.ClaimNetworkNameResult{})}, {"POST", "/referral-code/validate", rt(controller.ValidateReferralCodeArgs{}), rt(controller.ValidateNetworkReferralCodeResult{})}, {"GET", "/account/referral-network", nil, rt(controller.GetNetworkReferralResult{})}, {"GET", "/account/unlink-referral-network", nil, rt(controller.UnlinkReferralNetworkResult{})}, diff --git a/controller/account_controller.go b/controller/account_controller.go new file mode 100644 index 00000000..de405694 --- /dev/null +++ b/controller/account_controller.go @@ -0,0 +1,259 @@ +package controller + +import ( + "context" + "fmt" + "time" + + "github.com/urnetwork/server" + "github.com/urnetwork/server/model" + "github.com/urnetwork/server/session" +) + +const networkNameReclaimCooldown = 24 * time.Hour + +type ChangeNetworkNameArgs struct { + NetworkName string `json:"network_name"` + NewName string `json:"new_name"` +} + +type ChangeNetworkNameError struct { + Message string `json:"message"` +} + +type ChangeNetworkNameResult struct { + NetworkName string `json:"network_name"` + Error *ChangeNetworkNameError `json:"error,omitempty"` +} + +type ClaimNetworkNameArgs = ChangeNetworkNameArgs +type ClaimNetworkNameError = ChangeNetworkNameError +type ClaimNetworkNameResult = ChangeNetworkNameResult + +func ChangeNetworkName( + args ChangeNetworkNameArgs, + session *session.ClientSession, +) (*ChangeNetworkNameResult, error) { + return changeNetworkName(args, session, true) +} + +func ClaimNetworkName( + args ClaimNetworkNameArgs, + session *session.ClientSession, +) (*ClaimNetworkNameResult, error) { + return changeNetworkName(args, session, false) +} + +func changeNetworkName( + args ChangeNetworkNameArgs, + session *session.ClientSession, + reclaimCooldown bool, +) (*ChangeNetworkNameResult, error) { + // Seedphrase users must have email/phone, SSO, or a wallet bound to + // claim/change name — a signed wallet challenge is itself proof the + // user controls that identity, same as a verified email or SSO login. + if err := requireVerifiedIdentityBound(session.Ctx, session.ByJwt.UserId); err != nil { + return &ChangeNetworkNameResult{ + Error: &ChangeNetworkNameError{ + Message: err.Error(), + }, + }, nil + } + + // Accept either network_name or new_name field + name := args.NetworkName + if name == "" { + name = args.NewName + } + normalizedName, err := model.ValidateNetworkName(name) + if err != nil { + return &ChangeNetworkNameResult{ + Error: &ChangeNetworkNameError{ + Message: err.Error(), + }, + }, nil + } + + available, err := isNetworkNameAvailableForUser(session.Ctx, normalizedName, session.ByJwt.UserId) + if err != nil { + return nil, err + } + if !available { + return &ChangeNetworkNameResult{ + Error: &ChangeNetworkNameError{ + Message: "Network name not available.", + }, + }, nil + } + + // Claim (first-time name set, reachable once per account) and change + // (renaming an existing name) get separate daily budgets: bundling them + // would let a normal onboarding claim consume half of a new user's + // same-day rename budget before they've even picked a real name. + rateLimitAction := model.AccountActionClaimNetworkName + rateLimitDailyLimit := model.AccountActionClaimNetworkNameDailyLimit + if reclaimCooldown { + rateLimitAction = model.AccountActionChangeNetworkName + rateLimitDailyLimit = model.AccountActionChangeNetworkNameDailyLimit + } + if err := model.CheckAndRecordAccountActionRateLimit( + session.Ctx, + session.ByJwt.UserId, + rateLimitAction, + rateLimitDailyLimit, + model.AccountActionDailyWindow, + ); err != nil { + return &ChangeNetworkNameResult{ + Error: &ChangeNetworkNameError{ + Message: err.Error(), + }, + }, nil + } + + var oldName *string + server.Db(session.Ctx, func(conn server.PgConn) { + result, err := conn.Query( + session.Ctx, + ` + SELECT network_name FROM network + WHERE admin_user_id = $1 + `, + session.ByJwt.UserId, + ) + server.WithPgResult(result, err, func() { + if result.Next() { + var name string + server.Raise(result.Scan(&name)) + oldName = &name + } + }) + }) + + server.Tx(session.Ctx, func(tx server.PgTx) { + if reclaimCooldown && oldName != nil && *oldName != normalizedName { + coolDownUntil := server.NowUtc().Add(networkNameReclaimCooldown) + server.RaisePgResult(tx.Exec( + session.Ctx, + ` + INSERT INTO network_name_reclaim (old_name, cool_down_until) + VALUES ($1, $2) + ON CONFLICT (old_name) DO UPDATE + SET cool_down_until = EXCLUDED.cool_down_until + `, + *oldName, + coolDownUntil, + )) + } + + server.RaisePgResult(tx.Exec( + session.Ctx, + ` + UPDATE network + SET network_name = $2 + WHERE admin_user_id = $1 + `, + session.ByJwt.UserId, + normalizedName, + )) + }) + + return &ChangeNetworkNameResult{ + NetworkName: normalizedName, + }, nil +} + +func isNetworkNameAvailableForUser( + ctx context.Context, + name string, + userId server.Id, +) (bool, error) { + available := true + reasonErr := error(nil) + + server.Db(ctx, func(conn server.PgConn) { + // check if the name is taken by another network + result, err := conn.Query( + ctx, + ` + SELECT admin_user_id FROM network + WHERE network_name = $1 + `, + name, + ) + server.WithPgResult(result, err, func() { + if result.Next() { + var owner server.Id + server.Raise(result.Scan(&owner)) + if owner != userId { + available = false + } + } + }) + + if !available { + return + } + + // check if the name is in cooldown (network_name_reclaim) + result, err = conn.Query( + ctx, + ` + SELECT cool_down_until FROM network_name_reclaim + WHERE old_name = $1 + `, + name, + ) + server.WithPgResult(result, err, func() { + if result.Next() { + var coolDownUntil time.Time + server.Raise(result.Scan(&coolDownUntil)) + if server.NowUtc().Before(coolDownUntil) { + available = false + } + } + }) + }) + + if reasonErr != nil { + return false, fmt.Errorf("failed to check network name availability: %w", reasonErr) + } + + return available, nil +} + +// requireVerifiedIdentityBound checks that the user has at least one verified +// email/phone password auth, SSO auth, or wallet bound. Seedphrase-only users +// can't claim/change names — they need a verifiable identity method. +func requireVerifiedIdentityBound(ctx context.Context, userId server.Id) error { + var hasBoundAuth bool + + server.Db(ctx, func(conn server.PgConn) { + result, err := conn.Query( + ctx, + ` + SELECT EXISTS ( + SELECT 1 FROM network_user_auth_password + WHERE user_id = $1 AND verified = true + UNION ALL + SELECT 1 FROM network_user_auth_sso + WHERE user_id = $1 + UNION ALL + SELECT 1 FROM network_user_auth_wallet + WHERE user_id = $1 + ) + `, + userId, + ) + server.WithPgResult(result, err, func() { + if result.Next() { + server.Raise(result.Scan(&hasBoundAuth)) + } + }) + }) + + if !hasBoundAuth { + return fmt.Errorf("You must verify an email, bind a social login, or connect a wallet before changing your network name.") + } + + return nil +} diff --git a/controller/auth_controller.go b/controller/auth_controller.go index 9f7c6d63..b963aded 100644 --- a/controller/auth_controller.go +++ b/controller/auth_controller.go @@ -297,11 +297,33 @@ func RefreshToken(session *session.ClientSession) (*RefreshTokenResult, error) { &networkId, ) + // Re-read the current network name from the DB rather than carrying + // forward session.ByJwt.NetworkName -- that's just whatever name was + // baked into the JWT being refreshed, so a rename via + // change-name/claim-name would never show up in the client until it + // happened to get a truly fresh JWT (e.g. by logging out and back in). + networkName := session.ByJwt.NetworkName + server.Db(session.Ctx, func(conn server.PgConn) { + result, err := conn.Query( + session.Ctx, + ` + SELECT network_name FROM network + WHERE admin_user_id = $1 + `, + session.ByJwt.UserId, + ) + server.WithPgResult(result, err, func() { + if result.Next() { + server.Raise(result.Scan(&networkName)) + } + }) + }) + byJwt := jwt.NewByJwt( networkId, session.ByJwt.UserId, - session.ByJwt.NetworkName, - session.ByJwt.GuestMode, + networkName, + false, isPro, ) @@ -309,3 +331,34 @@ func RefreshToken(session *session.ClientSession) (*RefreshTokenResult, error) { ByJwt: byJwt.Client(*session.ByJwt.DeviceId, *session.ByJwt.ClientId).Sign(), }, nil } + +type AddAuthArgs = model.AddAuthMethod +type AddAuthResult = model.AddAuthMethodResult + +func AddAuth(args AddAuthArgs, session *session.ClientSession) (*AddAuthResult, error) { + return model.AddAuth(args, session) +} + +type RemoveAuthArgs struct { + AuthType string `json:"auth_type"` +} + +type RemoveAuthResult struct { + Error *RemoveAuthError `json:"error,omitempty"` +} + +type RemoveAuthError struct { + Message string `json:"message"` +} + +func RemoveAuth(args RemoveAuthArgs, session *session.ClientSession) (*RemoveAuthResult, error) { + err := model.RemoveAuth(session.Ctx, session.ByJwt.UserId, args.AuthType) + if err != nil { + return &RemoveAuthResult{ + Error: &RemoveAuthError{ + Message: err.Error(), + }, + }, nil + } + return &RemoveAuthResult{}, nil +} diff --git a/controller/network_controller.go b/controller/network_controller.go index 911bc0df..be6e481d 100644 --- a/controller/network_controller.go +++ b/controller/network_controller.go @@ -131,72 +131,6 @@ func UpdateNetworkName( return &UpdateNetworkNameResult{}, nil } -/** - * Upgrades a guest to a new account - */ -func UpgradeFromGuest( - upgradeGuest model.UpgradeGuestArgs, - session *session.ClientSession, -) (*model.UpgradeGuestResult, error) { - - result, err := model.UpgradeGuest( - upgradeGuest, - session, - ) - - if err != nil { - return nil, err - } - - // if verification required, send it - if result.VerificationRequired != nil { - verifySend := AuthVerifySendArgs{ - UserAuth: result.VerificationRequired.UserAuth, - UseNumeric: true, - } - AuthVerifySend(verifySend, session) - } else { - - if result.UserAuth != nil { - awsMessageSender := GetAWSMessageSender() - awsMessageSender.SendAccountMessageTemplate( - *result.UserAuth, - &NetworkWelcomeTemplate{}, - ) - } - - } - - return result, nil - -} - -/** - * Upgrades a guest to an existing account - */ -func UpgradeFromGuestExisting( - upgradeGuest model.UpgradeGuestExistingArgs, - session *session.ClientSession, -) (*model.UpgradeGuestExistingResult, error) { - - result, err := model.UpgradeFromGuestExisting( - upgradeGuest, - session, - ) - - // if verification required, send it - if result.VerificationRequired != nil { - verifySend := AuthVerifySendArgs{ - UserAuth: result.VerificationRequired.UserAuth, - UseNumeric: true, - } - AuthVerifySend(verifySend, session) - } - - return result, err - -} - type NetworkRemoveResultError struct { Message string `json:"message"` } diff --git a/controller/network_controller_test.go b/controller/network_controller_test.go index 2378e669..16cd61aa 100644 --- a/controller/network_controller_test.go +++ b/controller/network_controller_test.go @@ -41,7 +41,6 @@ func TestNetworkCreate(t *testing.T) { Password: &password, NetworkName: "foobar", Terms: true, - GuestMode: false, ReferralCode: &referralCode.ReferralCode, } result, err := NetworkCreate(networkCreate, session) @@ -102,7 +101,6 @@ func TestNetworkCreateWithProfanity(t *testing.T) { Password: &password, NetworkName: "shitty", // must be at least 6 characters Terms: true, - GuestMode: false, ReferralCode: &referralCode, } result, err := NetworkCreate(networkCreate, session) @@ -213,7 +211,6 @@ func TestNetworkCreateWithBalanceCodeSuccess(t *testing.T) { Password: &password, NetworkName: "foobar", Terms: true, - GuestMode: false, BalanceCode: &balanceCode.Secret, } diff --git a/controller/seedphrase_controller.go b/controller/seedphrase_controller.go new file mode 100644 index 00000000..7ffc2523 --- /dev/null +++ b/controller/seedphrase_controller.go @@ -0,0 +1,114 @@ +package controller + +import ( + "github.com/urnetwork/server/model" + "github.com/urnetwork/server/session" +) + +type RegenerateSeedphraseArgs struct { +} + +type RegenerateSeedphraseResult struct { + Seedphrase string `json:"seedphrase"` + Error *RegenerateSeedphraseError `json:"error,omitempty"` +} + +type RegenerateSeedphraseError struct { + Message string `json:"message"` +} + +type GenerateSeedphraseArgs struct { +} + +type GenerateSeedphraseResult struct { + Seedphrase string `json:"seedphrase"` + Error *GenerateSeedphraseError `json:"error,omitempty"` +} + +type GenerateSeedphraseError struct { + Message string `json:"message"` +} + +func RegenerateSeedphrase( + args RegenerateSeedphraseArgs, + session *session.ClientSession, +) (*RegenerateSeedphraseResult, error) { + hasSeedphrase, err := model.HasSeedphraseAuth(session.Ctx, session.ByJwt.UserId) + if err != nil { + return nil, err + } + if !hasSeedphrase { + return &RegenerateSeedphraseResult{ + Error: &RegenerateSeedphraseError{ + Message: "No seedphrase auth found.", + }, + }, nil + } + + if err := model.CheckAccountActionRateLimit( + session.Ctx, + session.ByJwt.UserId, + model.AccountActionRegenerateSeedphrase, + model.AccountActionRegenerateSeedphraseDailyLimit, + model.AccountActionDailyWindow, + ); err != nil { + return &RegenerateSeedphraseResult{ + Error: &RegenerateSeedphraseError{ + Message: err.Error(), + }, + }, nil + } + + seedphrase, err := model.RegenerateSeedphrase(session.Ctx, session.ByJwt.UserId) + if err != nil { + return nil, err + } + + model.RecordAccountActionAttempt(session.Ctx, session.ByJwt.UserId, model.AccountActionRegenerateSeedphrase) + + return &RegenerateSeedphraseResult{ + Seedphrase: seedphrase, + }, nil +} + +func GenerateSeedphrase( + args GenerateSeedphraseArgs, + session *session.ClientSession, +) (*GenerateSeedphraseResult, error) { + hasSeedphrase, err := model.HasSeedphraseAuth(session.Ctx, session.ByJwt.UserId) + if err != nil { + return nil, err + } + if hasSeedphrase { + return &GenerateSeedphraseResult{ + Error: &GenerateSeedphraseError{ + Message: "Seedphrase already exists, use regenerate instead.", + }, + }, nil + } + + if err := model.CheckAccountActionRateLimit( + session.Ctx, + session.ByJwt.UserId, + model.AccountActionGenerateSeedphrase, + model.AccountActionGenerateSeedphraseDailyLimit, + model.AccountActionDailyWindow, + ); err != nil { + return &GenerateSeedphraseResult{ + Error: &GenerateSeedphraseError{ + Message: err.Error(), + }, + }, nil + } + + seedphrase, err := model.GenerateSeedphrase(session.Ctx, session.ByJwt.UserId) + if err != nil { + return nil, err + } + + model.RecordAccountActionAttempt(session.Ctx, session.ByJwt.UserId, model.AccountActionGenerateSeedphrase) + + return &GenerateSeedphraseResult{ + Seedphrase: seedphrase, + }, nil +} diff --git a/db_migrations.go b/db_migrations.go index 063f091e..90293f3a 100644 --- a/db_migrations.go +++ b/db_migrations.go @@ -4332,6 +4332,55 @@ var migrations = []any{ ON transfer_contract (destination_id) WHERE open `), + // seedphrase auth + newSqlMigration(` + CREATE TABLE IF NOT EXISTS network_user_auth_seedphrase ( + user_id uuid NOT NULL PRIMARY KEY, + seedphrase_lookup bytea NOT NULL, + seedphrase_hash bytea NOT NULL, + seedphrase_salt bytea NOT NULL, + create_time timestamp NOT NULL DEFAULT now() + ) + `), + newSqlMigration(` + CREATE UNIQUE INDEX IF NOT EXISTS network_user_auth_seedphrase_lookup + ON network_user_auth_seedphrase (seedphrase_lookup) + `), + + // network name reclaim (1-day cooldown) + newSqlMigration(` + CREATE TABLE IF NOT EXISTS network_name_reclaim ( + old_name varchar(256) NOT NULL PRIMARY KEY, + cool_down_until timestamp NOT NULL + ) + `), + + // network create rate limit (5 per IP per day) + newSqlMigration(` + CREATE TABLE IF NOT EXISTS network_create_attempt ( + network_create_attempt_id uuid NOT NULL PRIMARY KEY, + client_address_hash bytea NOT NULL, + create_time timestamp NOT NULL DEFAULT now() + ) + `), + newSqlMigration(` + CREATE INDEX IF NOT EXISTS network_create_attempt_hash_time + ON network_create_attempt (client_address_hash, create_time) + `), + + // Index for bulk client deactivation via `client_id = ANY($1) AND network_id = $2`. + // Both the synchronous batch path (RemoveNetworkClientsBatch, ≤10k ids) and the + // background task path (RemoveNetworkClientsTask, 200k ids per invocation) use + // this predicate. Without this index the query planner must scan for each batch; + // at 200k IDs across 20 batches that can cause long lock times on large tables. + // On the large existing network_client table this must be built manually with + // CREATE INDEX CONCURRENTLY out of band — the IF NOT EXISTS gate makes this + // migration a no-op once it is pre-created. + newSqlMigration(` + CREATE INDEX IF NOT EXISTS network_client_network_id_client_id + ON network_client (network_id, client_id) + `), + // Defuse the 2026-07-17 planner-stats landmine (monitor/SIGNALS.md 2.3/5.8): at // steady state only ~10-50k of 530M rows have open=true (~6e-5), so the // default ANALYZE sample (300 x 100 = 30k rows) can miss every one. @@ -4385,4 +4434,56 @@ var migrations = []any{ CREATE INDEX IF NOT EXISTS network_client_top_level_contract_time ON network_client (contract_time) WHERE (active = true AND source_client_id IS NULL AND contract_time IS NOT NULL) `), + // account-based (not IP-based) daily rate limits on sensitive account + // actions: add/remove auth method, change/claim network name, and + // generate/regenerate seedphrase. One shared table, keyed by (user_id, + // action), so each action gets its own independent daily counter. + newSqlMigration(` + CREATE TABLE IF NOT EXISTS network_user_action_attempt ( + network_user_action_attempt_id uuid NOT NULL PRIMARY KEY, + user_id uuid NOT NULL, + action varchar(64) NOT NULL, + create_time timestamp NOT NULL DEFAULT now() + ) + `), + newSqlMigration(` + CREATE INDEX IF NOT EXISTS network_user_action_attempt_user_action_time + ON network_user_action_attempt (user_id, action, create_time) + `), + + // a single, deployment-wide ledger for the bulk client removal API + // (RemoveNetworkClients): each admitted request records how many client + // ids it covered, and CheckAndRecordBulkClientRemovalQuota sums this + // column over the trailing hour to enforce MaxBulkClientRemovalsPerHour. + // network_id is stored for observability only -- the limit itself is + // global, not scoped per network. + newSqlMigration(` + CREATE TABLE IF NOT EXISTS bulk_client_removal_quota ( + bulk_client_removal_quota_id uuid NOT NULL PRIMARY KEY, + network_id uuid NOT NULL, + client_count int NOT NULL, + create_time timestamp NOT NULL DEFAULT now() + ) + `), + newSqlMigration(` + CREATE INDEX IF NOT EXISTS bulk_client_removal_quota_create_time + ON bulk_client_removal_quota (create_time) + `), + + // bulk_client_removal_quota moves from a rolling trailing-hour window to + // fixed, non-overlapping hourly buckets (UTC): ReserveBulkClientRemovalSlot + // reserves a row against a specific bucket_start (the current hour, or a + // future one if the current hour is full), instead of every row implicitly + // counting against whatever "now" happens to be when queried. create_time + // remains as an audit field (when the reservation was made), separate from + // bucket_start (which hour it counts toward). No rows exist yet in + // practice, so the NOT NULL default here is never relied on for real data. + newSqlMigration(` + ALTER TABLE bulk_client_removal_quota + ADD COLUMN IF NOT EXISTS bucket_start timestamp NOT NULL DEFAULT now() + `), + newSqlMigration(` + CREATE INDEX IF NOT EXISTS bulk_client_removal_quota_bucket_start + ON bulk_client_removal_quota (bucket_start) + `), } diff --git a/go.mod b/go.mod index e4bddacf..ac1afaca 100644 --- a/go.mod +++ b/go.mod @@ -11,6 +11,7 @@ require ( github.com/ethereum/go-ethereum v1.16.7 github.com/gagliardetto/solana-go v1.21.0 github.com/go-jose/go-jose/v3 v3.0.4 + github.com/go-playground/assert/v2 v2.2.0 github.com/golang-jwt/jwt/v5 v5.3.1 github.com/gorilla/websocket v1.5.3 github.com/jackc/pgerrcode v0.0.0-20250907135507-afb5586c32a6 diff --git a/go.sum b/go.sum index c4bf3bcd..6c249f76 100644 --- a/go.sum +++ b/go.sum @@ -49,8 +49,8 @@ github.com/coreos/go-semver v0.3.1 h1:yi21YpKnrx1gt5R+la8n5WgS0kCrsPp33dmEyHReZr github.com/coreos/go-semver v0.3.1/go.mod h1:irMmmIw/7yzSRPWryHsK7EYSg09caPQL03VsM8rvUec= github.com/cosmos/go-bip39 v1.0.0 h1:pcomnQdrdH22njcAatO0yWojsUnCO3y2tNoV1cb6hHY= github.com/cosmos/go-bip39 v1.0.0/go.mod h1:RNJv0H/pOIVgxw6KS7QeX2a0Uo0aKUlfhZ4xuwvCdJw= -github.com/cpuguy83/go-md2man/v2 v2.0.5 h1:ZtcqGrnekaHpVLArFSe4HK5DoKx1T0rq2DwVB0alcyc= -github.com/cpuguy83/go-md2man/v2 v2.0.5/go.mod h1:tgQtvFlXSQOSOSIRvRPT7W67SCa46tRHOmNcaadrF8o= +github.com/cpuguy83/go-md2man/v2 v2.0.7 h1:zbFlGlXEAKlwXpmvle3d8Oe3YnkKIK4xSRTd3sHPnBo= +github.com/cpuguy83/go-md2man/v2 v2.0.7/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g= github.com/crate-crypto/go-eth-kzg v1.4.0 h1:WzDGjHk4gFg6YzV0rJOAsTK4z3Qkz5jd4RE3DAvPFkg= github.com/crate-crypto/go-eth-kzg v1.4.0/go.mod h1:J9/u5sWfznSObptgfa92Jq8rTswn6ahQWEuiLHOjCUI= github.com/crate-crypto/go-ipa v0.0.0-20240724233137-53bbb0ceb27a h1:W8mUrRp6NOVl3J+MYp5kPMoUZPp7aOYHtaua31lwRHg= @@ -108,6 +108,8 @@ github.com/go-jose/go-jose/v3 v3.0.4/go.mod h1:5b+7YgP7ZICgJDBdfjZaIt+H/9L9T/YQr github.com/go-ole/go-ole v1.2.5/go.mod h1:pprOEPIfldk/42T2oK7lQ4v4JSDwmV0As9GaiUsvbm0= github.com/go-ole/go-ole v1.3.0 h1:Dt6ye7+vXGIKZ7Xtk4s6/xVdGDQynvom7xCFEdWr6uE= github.com/go-ole/go-ole v1.3.0/go.mod h1:5LS6F96DhAwUc7C+1HLexzMXY1xGRSryjyPPKW6zv78= +github.com/go-playground/assert/v2 v2.2.0 h1:JvknZsQTYeFEAhQwI4qEt9cyV5ONwRHC+lYKSsYSR8s= +github.com/go-playground/assert/v2 v2.2.0/go.mod h1:VDjEfimB/XKnb+ZQfWdccd7VUvScMdVu0Titje2rxJ4= github.com/goccy/go-json v0.10.6 h1:p8HrPJzOakx/mn/bQtjgNjdTcN+/S6FcG2CTtQOrHVU= github.com/goccy/go-json v0.10.6/go.mod h1:oq7eo15ShAhp70Anwd5lgX2pLfOS3QCiwU/PULtXL6M= github.com/gofrs/flock v0.13.0 h1:95JolYOvGMqeH31+FC7D2+uULf6mG61mEZ/A8dRYMzw= @@ -366,8 +368,8 @@ github.com/tyler-smith/go-bip39 v1.1.0 h1:5eUemwrMargf3BSLRRCalXT93Ns6pQJIjYQN2n github.com/tyler-smith/go-bip39 v1.1.0/go.mod h1:gUYDtqQw1JS3ZJ8UWVcGTGqqr6YIN3CWg+kkNaLt55U= github.com/ulikunitz/xz v0.5.15 h1:9DNdB5s+SgV3bQ2ApL10xRc35ck0DuIX/isZvIk+ubY= github.com/ulikunitz/xz v0.5.15/go.mod h1:nbz6k7qbPmH4IRqmfOplQw/tblSgqTqBwxkY0oWt/14= -github.com/urfave/cli/v2 v2.27.5 h1:WoHEJLdsXr6dDWoJgMq/CboDmyY/8HMMH1fTECbih+w= -github.com/urfave/cli/v2 v2.27.5/go.mod h1:3Sevf16NykTbInEnD0yKkjDAeZDS0A6bzhBH5hrMvTQ= +github.com/urfave/cli/v2 v2.27.7 h1:bH59vdhbjLv3LAvIu6gd0usJHgoTTPhCFib8qqOwXYU= +github.com/urfave/cli/v2 v2.27.7/go.mod h1:CyNAG/xg+iAOg0N4MPGZqVmv2rCoP267496AOXUZjA4= github.com/wlynxg/anet v0.0.5 h1:J3VJGi1gvo0JwZ/P1/Yc/8p63SoW98B5dHkYDmpgvvU= github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguHxoA= github.com/xrash/smetrics v0.0.0-20240521201337-686a1a2994c1 h1:gEOO8jv9F4OT7lGCjxCBTO/36wtF6j2nSip77qHd4x4= diff --git a/jwt/by_jwt.go b/jwt/by_jwt.go index c7ca552c..c0b620e2 100644 --- a/jwt/by_jwt.go +++ b/jwt/by_jwt.go @@ -160,8 +160,9 @@ type ByJwt struct { CreateTime time.Time `json:"create_time,omitempty"` DeviceId *server.Id `json:"device_id,omitempty"` ClientId *server.Id `json:"client_id,omitempty"` - GuestMode bool `json:"guest_mode,omitempty"` - Pro bool `json:"pro,omitempty"` + // Deprecated: always false for new tokens. Field kept for backward compat with existing guest JWTs. + GuestMode bool `json:"guest_mode,omitempty"` + Pro bool `json:"pro,omitempty"` // identity roles and principal, assigned at client or auth code creation. // The values have no meaning to the network. Roles []string `json:"roles,omitempty"` diff --git a/model/account_action_rate_limit.go b/model/account_action_rate_limit.go new file mode 100644 index 00000000..9931a4cb --- /dev/null +++ b/model/account_action_rate_limit.go @@ -0,0 +1,180 @@ +package model + +import ( + "context" + "fmt" + "time" + + "github.com/urnetwork/server" +) + +// Account-based (not IP-based) daily rate limits on sensitive account-mutation +// actions. Each action has its own independent daily counter, keyed by userId. + +const ( + AccountActionAddAuth = "add_auth" + AccountActionRemoveAuth = "remove_auth" + AccountActionClaimNetworkName = "claim_network_name" + AccountActionChangeNetworkName = "change_network_name" + AccountActionGenerateSeedphrase = "generate_seedphrase" + AccountActionRegenerateSeedphrase = "regenerate_seedphrase" +) + +const ( + AccountActionAddAuthDailyLimit = 5 + AccountActionRemoveAuthDailyLimit = 5 + AccountActionClaimNetworkNameDailyLimit = 2 + AccountActionChangeNetworkNameDailyLimit = 2 + AccountActionGenerateSeedphraseDailyLimit = 5 + AccountActionRegenerateSeedphraseDailyLimit = 5 +) + +const AccountActionDailyWindow = 24 * time.Hour + +// actionDisplayNames gives each action a human-readable name for the +// rate-limit error message, so a client can tell which of the several +// daily limits it hit instead of seeing one generic message for all of them. +var actionDisplayNames = map[string]string{ + AccountActionAddAuth: "adding a sign-in method", + AccountActionRemoveAuth: "removing a sign-in method", + AccountActionClaimNetworkName: "setting your network name", + AccountActionChangeNetworkName: "changing your network name", + AccountActionGenerateSeedphrase: "generating a seedphrase", + AccountActionRegenerateSeedphrase: "regenerating a seedphrase", +} + +func maxAccountActionAttemptsError(action string) error { + name, ok := actionDisplayNames[action] + if !ok { + name = "this action" + } + return fmt.Errorf("You have reached the maximum number of attempts for %s today. Please try again later.", name) +} + +func countRecentAccountActionAttempts( + ctx context.Context, + tx server.PgTx, + userId server.Id, + action string, + window time.Duration, +) int { + var count int + // the cutoff is computed here in Go, not via `now() - INTERVAL ...` in + // SQL: create_time is a naive `timestamp` column holding UTC wall-clock + // values (via server.NowUtc()), and `now()` returns `timestamptz`. + // Comparing the two forces Postgres to cast one side using the + // session's TimeZone setting -- on a non-UTC session that silently + // shifts the effective window by the zone offset. Passing an explicit + // UTC cutoff avoids any zone-dependent cast. + cutoff := server.NowUtc().Add(-window) + result, err := tx.Query( + ctx, + ` + SELECT COUNT(*) FROM network_user_action_attempt + WHERE user_id = $1 AND action = $2 + AND $3 <= create_time + `, + userId, action, cutoff, + ) + server.WithPgResult(result, err, func() { + if result.Next() { + server.Raise(result.Scan(&count)) + } + }) + return count +} + +// CheckAccountActionRateLimit reports whether userId is still under `limit` +// attempts of `action` in the trailing `window`. It does not record +// anything — pair with RecordAccountActionAttempt after the action actually +// succeeds, so failed/rejected attempts don't burn the caller's budget. +func CheckAccountActionRateLimit( + ctx context.Context, + userId server.Id, + action string, + limit int, + window time.Duration, +) error { + var count int + cutoff := server.NowUtc().Add(-window) + server.Db(ctx, func(conn server.PgConn) { + result, err := conn.Query( + ctx, + ` + SELECT COUNT(*) FROM network_user_action_attempt + WHERE user_id = $1 AND action = $2 + AND $3 <= create_time + `, + userId, action, cutoff, + ) + server.WithPgResult(result, err, func() { + if result.Next() { + server.Raise(result.Scan(&count)) + } + }) + }) + if limit <= count { + return maxAccountActionAttemptsError(action) + } + return nil +} + +// RecordAccountActionAttempt records one attempt of `action` by userId. +// Call this only once the action it is guarding has actually succeeded. +func RecordAccountActionAttempt(ctx context.Context, userId server.Id, action string) { + server.Tx(ctx, func(tx server.PgTx) { + server.RaisePgResult(tx.Exec( + ctx, + ` + INSERT INTO network_user_action_attempt (network_user_action_attempt_id, user_id, action, create_time) + VALUES ($1, $2, $3, $4) + `, + server.NewId(), userId, action, server.NowUtc(), + )) + }) +} + +// CheckAndRecordAccountActionRateLimit atomically checks and records an +// attempt in one transaction. Use this for strict, attempt-counted actions +// (e.g. a name change) where every real attempt -- not just successes -- +// should count against the budget. Call it only after cheap validation +// (format, availability) has already passed, so malformed requests don't +// burn budget. +func CheckAndRecordAccountActionRateLimit( + ctx context.Context, + userId server.Id, + action string, + limit int, + window time.Duration, +) error { + var count int + server.Tx(ctx, func(tx server.PgTx) { + count = countRecentAccountActionAttempts(ctx, tx, userId, action, window) + if limit <= count { + return + } + server.RaisePgResult(tx.Exec( + ctx, + ` + INSERT INTO network_user_action_attempt (network_user_action_attempt_id, user_id, action, create_time) + VALUES ($1, $2, $3, $4) + `, + server.NewId(), userId, action, server.NowUtc(), + )) + }) + if limit <= count { + return maxAccountActionAttemptsError(action) + } + return nil +} + +// RemoveExpiredAccountActionAttempts cleans up attempts older than minTime. +func RemoveExpiredAccountActionAttempts(ctx context.Context, minTime time.Time) { + server.MaintenanceTx(ctx, func(tx server.PgTx) { + server.RaisePgResult(tx.Exec( + ctx, + `DELETE FROM network_user_action_attempt WHERE create_time < $1`, + minTime.UTC(), + )) + }) +} diff --git a/model/account_action_rate_limit_test.go b/model/account_action_rate_limit_test.go new file mode 100644 index 00000000..1c754ecc --- /dev/null +++ b/model/account_action_rate_limit_test.go @@ -0,0 +1,175 @@ +package model + +import ( + "context" + "testing" + "time" + + "github.com/go-playground/assert/v2" + "github.com/urnetwork/server" +) + +func TestCheckAccountActionRateLimitAllowsUpToLimitThenBlocks(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + userId := server.NewId() + + const limit = 5 + for i := 0; i < limit; i++ { + err := CheckAccountActionRateLimit(ctx, userId, "test_action", limit, AccountActionDailyWindow) + assert.Equal(t, err, nil) + RecordAccountActionAttempt(ctx, userId, "test_action") + } + + // the 6th check, after 5 recorded successes, must be blocked + err := CheckAccountActionRateLimit(ctx, userId, "test_action", limit, AccountActionDailyWindow) + assert.NotEqual(t, err, nil) + }) +} + +func TestCheckAccountActionRateLimitIsPerActionAndPerUser(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + userId := server.NewId() + otherUserId := server.NewId() + + const limit = 1 + err := CheckAccountActionRateLimit(ctx, userId, "action_a", limit, AccountActionDailyWindow) + assert.Equal(t, err, nil) + RecordAccountActionAttempt(ctx, userId, "action_a") + + // same user, different action: independent counter + err = CheckAccountActionRateLimit(ctx, userId, "action_b", limit, AccountActionDailyWindow) + assert.Equal(t, err, nil) + + // different user, same action: independent counter + err = CheckAccountActionRateLimit(ctx, otherUserId, "action_a", limit, AccountActionDailyWindow) + assert.Equal(t, err, nil) + + // same user, same action, already at limit: blocked + err = CheckAccountActionRateLimit(ctx, userId, "action_a", limit, AccountActionDailyWindow) + assert.NotEqual(t, err, nil) + }) +} + +func TestCheckAccountActionRateLimitDoesNotCountUnrecordedChecks(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + userId := server.NewId() + + const limit = 1 + + // repeatedly checking (without recording) never burns budget -- this + // is the success-only counting semantics: a caller that checks, then + // fails validation/the underlying action, and never calls + // RecordAccountActionAttempt, must not have spent the user's budget. + for i := 0; i < 10; i++ { + err := CheckAccountActionRateLimit(ctx, userId, "test_action", limit, AccountActionDailyWindow) + assert.Equal(t, err, nil) + } + }) +} + +func TestCheckAndRecordAccountActionRateLimitAllowsUpToLimitThenBlocks(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + userId := server.NewId() + + const limit = 2 + for i := 0; i < limit; i++ { + err := CheckAndRecordAccountActionRateLimit(ctx, userId, AccountActionChangeNetworkName, limit, AccountActionDailyWindow) + assert.Equal(t, err, nil) + } + + // the 3rd attempt, already at limit, must be blocked -- and blocked + // attempts must not themselves record (else the block would never + // clear even after the window passes) + err := CheckAndRecordAccountActionRateLimit(ctx, userId, AccountActionChangeNetworkName, limit, AccountActionDailyWindow) + assert.NotEqual(t, err, nil) + }) +} + +func TestAccountActionRateLimitWindowExpires(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + userId := server.NewId() + + const limit = 1 + err := CheckAccountActionRateLimit(ctx, userId, "test_action", limit, AccountActionDailyWindow) + assert.Equal(t, err, nil) + RecordAccountActionAttempt(ctx, userId, "test_action") + + err = CheckAccountActionRateLimit(ctx, userId, "test_action", limit, AccountActionDailyWindow) + assert.NotEqual(t, err, nil) + + // with a window shorter than the time that has already elapsed since + // recording (any positive duration, since the insert above already + // took non-zero time), the prior attempt falls outside the window and + // no longer counts + err = CheckAccountActionRateLimit(ctx, userId, "test_action", limit, time.Nanosecond) + assert.Equal(t, err, nil) + }) +} + +func TestAccountActionConstantsAreDistinctAndHaveDisplayNames(t *testing.T) { + actions := []string{ + AccountActionAddAuth, + AccountActionRemoveAuth, + AccountActionClaimNetworkName, + AccountActionChangeNetworkName, + AccountActionGenerateSeedphrase, + AccountActionRegenerateSeedphrase, + } + + seen := map[string]bool{} + for _, action := range actions { + // every real action constant must have a distinct string value -- + // two constants colliding would silently merge their rate limits + assert.Equal(t, seen[action], false) + seen[action] = true + + // every real action constant must produce a specific, non-generic + // error message -- the fallback "this action" wording should only + // ever be reached for a caller-supplied action string that isn't + // one of the real constants, never for a real one + err := maxAccountActionAttemptsError(action) + assert.NotEqual(t, err.Error(), maxAccountActionAttemptsError("not_a_real_action").Error()) + } +} + +func TestClaimAndChangeNetworkNameHaveIndependentCounters(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + userId := server.NewId() + + // exhaust the claim budget + for i := 0; i < AccountActionClaimNetworkNameDailyLimit; i++ { + err := CheckAndRecordAccountActionRateLimit(ctx, userId, AccountActionClaimNetworkName, AccountActionClaimNetworkNameDailyLimit, AccountActionDailyWindow) + assert.Equal(t, err, nil) + } + err := CheckAndRecordAccountActionRateLimit(ctx, userId, AccountActionClaimNetworkName, AccountActionClaimNetworkNameDailyLimit, AccountActionDailyWindow) + assert.NotEqual(t, err, nil) + + // a subsequent rename (a real, separate action) must still be + // allowed -- claiming a name during onboarding must not consume the + // same-day rename budget + err = CheckAndRecordAccountActionRateLimit(ctx, userId, AccountActionChangeNetworkName, AccountActionChangeNetworkNameDailyLimit, AccountActionDailyWindow) + assert.Equal(t, err, nil) + }) +} + +func TestRemoveExpiredAccountActionAttempts(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + userId := server.NewId() + + RecordAccountActionAttempt(ctx, userId, "test_action") + + // cutoff in the future: the just-recorded attempt is older than it, + // so it gets removed, freeing up budget immediately + RemoveExpiredAccountActionAttempts(ctx, server.NowUtc().Add(time.Second)) + + err := CheckAccountActionRateLimit(ctx, userId, "test_action", 1, AccountActionDailyWindow) + assert.Equal(t, err, nil) + }) +} diff --git a/model/account_model.go b/model/account_model.go index 8d3be07f..e5ae48c0 100644 --- a/model/account_model.go +++ b/model/account_model.go @@ -222,6 +222,15 @@ func RemoveNetwork( `, userId, ) + + // (cascade) delete network_user_auth_seedphrase + batch.Queue( + ` + DELETE FROM network_user_auth_seedphrase + WHERE user_id = $1 + `, + userId, + ) } }) diff --git a/model/auth_model.go b/model/auth_model.go index 905e0c94..9769fc92 100644 --- a/model/auth_model.go +++ b/model/auth_model.go @@ -33,12 +33,13 @@ const VerifyCodeTimeout = 4 * time.Hour type AuthType = string const ( - AuthTypePassword AuthType = "password" - AuthTypeApple AuthType = "apple" - AuthTypeGoogle AuthType = "google" - AuthTypeBringYour AuthType = "bringyour" - AuthTypeGuest AuthType = "guest" - AuthTypeSolana AuthType = "solana" + AuthTypePassword AuthType = "password" + AuthTypeApple AuthType = "apple" + AuthTypeGoogle AuthType = "google" + AuthTypeBringYour AuthType = "bringyour" + AuthTypeGuest AuthType = "guest" + AuthTypeSolana AuthType = "solana" + AuthTypeSeedphrase AuthType = "seedphrase" ) type WalletAuthArgs struct { @@ -57,6 +58,7 @@ type AuthLoginArgs struct { AuthJwtType *string `json:"auth_jwt_type,omitempty"` AuthJwt *string `json:"auth_jwt,omitempty"` WalletAuth *WalletAuthArgs `json:"wallet_auth,omitempty"` + Seedphrase *string `json:"seedphrase,omitempty"` } type AuthLoginResult struct { @@ -151,14 +153,40 @@ func AuthLogin( } } else if login.WalletAuth != nil { - return handleLoginWallet( + result, err := handleLoginWallet( login.WalletAuth, session.Ctx, ) - + // Only wallet logins that resolve to an existing user (a JWT-bearing + // result) count as a successful attempt -- a "new wallet user" + // result still requires registration, matching how the SSO branch + // above only marks success once it returns a Network result. + if err == nil && result != nil && result.Network != nil { + SetUserAuthAttemptSuccess(session.Ctx, userAuthAttemptId, true) + } + return result, err + } else if login.Seedphrase != nil && *login.Seedphrase != "" { + result, err := LoginWithSeedphrase(session.Ctx, *login.Seedphrase) + if err != nil { + return &AuthLoginResult{ + Error: &AuthLoginResultError{ + Message: err.Error(), + }, + }, nil + } + SetUserAuthAttemptSuccess(session.Ctx, userAuthAttemptId, true) + return &AuthLoginResult{ + Network: &AuthLoginResultNetwork{ + ByJwt: result.ByJwt, + }, + }, nil } - return nil, errors.New("invalid login") + return &AuthLoginResult{ + Error: &AuthLoginResultError{ + Message: "Invalid login credentials.", + }, + }, nil } /** diff --git a/model/auth_model_identity.go b/model/auth_model_identity.go index 8496d828..f30a6e48 100644 --- a/model/auth_model_identity.go +++ b/model/auth_model_identity.go @@ -132,3 +132,19 @@ func createResetCode() string { } return hex.EncodeToString(resetCode) } + +func computeSeedphraseHash(seedphrase []byte, salt []byte) []byte { + pepperedSeedphrase := []byte{} + pepperedSeedphrase = append(pepperedSeedphrase, passwordPepper()...) + pepperedSeedphrase = append(pepperedSeedphrase, seedphrase...) + return argon2.Key(pepperedSeedphrase, salt, 3, 32*1024, 4, 32) +} + +func createSeedphraseSalt() []byte { + salt := make([]byte, 32) + _, err := rand.Read(salt) + if err != nil { + panic(err) + } + return salt +} diff --git a/model/bulk_client_removal_rate_limit.go b/model/bulk_client_removal_rate_limit.go new file mode 100644 index 00000000..8b53cf58 --- /dev/null +++ b/model/bulk_client_removal_rate_limit.go @@ -0,0 +1,147 @@ +package model + +import ( + "context" + "fmt" + "time" + + "github.com/urnetwork/server" +) + +// A deployment-wide budget on the bulk removal API (RemoveNetworkClients), +// counted across all networks and accounts together -- this does not +// distinguish who is calling it, only how much total bulk-removal work has +// been admitted. It guards against many networks each triggering large +// bulk-deletes at once placing an unbounded aggregate load on the database +// in a short window. +// +// The budget is split into fixed, non-overlapping hourly buckets (UTC), +// rather than a rolling trailing window: a request is reserved against the +// earliest bucket -- starting with the current hour -- that has room for +// it. If the current hour is full, the request is queued for the next hour +// instead of rejected, and so on, up to MaxBulkClientRemovalLookaheadBuckets +// hours out. Only once no bucket in that whole lookahead window has room -- +// meaning the deployment-wide daily budget (MaxBulkClientRemovalsPerDay) is +// genuinely exhausted -- is the request rejected and shown to the caller. +// Fixed buckets (rather than a rolling window) are what make "queue for the +// next hour" well-defined: there needs to be a discrete boundary to wait +// for. +// +// This intentionally does NOT apply to the single-client removal API +// (RemoveNetworkClient) -- that path is a single-row update with no +// comparable blast radius, and callers should not have their normal, +// incremental client cleanup throttled by bulk-API abuse elsewhere. +const BulkClientRemovalBucketDuration = time.Hour +const MaxBulkClientRemovalsPerBucket = 1000000 +const MaxBulkClientRemovalLookaheadBuckets = 24 +const MaxBulkClientRemovalsPerDay = MaxBulkClientRemovalLookaheadBuckets * MaxBulkClientRemovalsPerBucket + +func maxBulkClientRemovalsError() error { + return fmt.Errorf( + "The bulk client removal limit (%d per hour, %d per day) has been reached for this deployment. Please try again later.", + MaxBulkClientRemovalsPerBucket, + MaxBulkClientRemovalsPerDay, + ) +} + +// bulkClientRemovalBucketStart truncates t to the start of its UTC hourly +// bucket. Truncate operates on the absolute duration since the zero time, +// so this only aligns to true UTC hour boundaries when t is already UTC. +func bulkClientRemovalBucketStart(t time.Time) time.Time { + return t.UTC().Truncate(BulkClientRemovalBucketDuration) +} + +// ReserveBulkClientRemovalSlot finds the earliest hourly bucket -- starting +// with the current hour -- that has room for `count` more removals, and +// reserves it there. It returns the reservation's id (for +// CancelBulkClientRemovalReservation, if the caller can't ultimately use the +// slot it just reserved -- e.g. a duplicate in-flight request for the same +// network) and the bucket's start time, which the caller should use as the +// task's RunAt when the bucket isn't the current hour. +// +// If no bucket within MaxBulkClientRemovalLookaheadBuckets hours has room, +// nothing is reserved and an error is returned -- this is the deployment's +// daily cap. A single request can never itself be too large to fit an empty +// bucket (RemoveNetworkClients already caps a request at +// MaxRemoveNetworkClientsCount, which equals MaxBulkClientRemovalsPerBucket), +// so this only happens under genuine contention from other reservations. +func ReserveBulkClientRemovalSlot( + ctx context.Context, + networkId server.Id, + count int, +) (reservationId server.Id, bucketStart time.Time, err error) { + windowStart := bulkClientRemovalBucketStart(server.NowUtc()) + windowEnd := windowStart.Add(MaxBulkClientRemovalLookaheadBuckets * BulkClientRemovalBucketDuration) + + server.Tx(ctx, func(tx server.PgTx) { + usedByBucket := map[time.Time]int{} + result, qerr := tx.Query( + ctx, + ` + SELECT bucket_start, SUM(client_count) FROM bulk_client_removal_quota + WHERE $1 <= bucket_start AND bucket_start < $2 + GROUP BY bucket_start + `, + windowStart, windowEnd, + ) + server.WithPgResult(result, qerr, func() { + for result.Next() { + var bucket time.Time + var used int + server.Raise(result.Scan(&bucket, &used)) + usedByBucket[bucket] = used + } + }) + + for i := 0; i < MaxBulkClientRemovalLookaheadBuckets; i++ { + candidate := windowStart.Add(time.Duration(i) * BulkClientRemovalBucketDuration) + if usedByBucket[candidate]+count <= MaxBulkClientRemovalsPerBucket { + id := server.NewId() + server.RaisePgResult(tx.Exec( + ctx, + ` + INSERT INTO bulk_client_removal_quota (bulk_client_removal_quota_id, network_id, client_count, bucket_start, create_time) + VALUES ($1, $2, $3, $4, $5) + `, + id, networkId, count, candidate, server.NowUtc(), + )) + reservationId = id + bucketStart = candidate + return + } + } + err = maxBulkClientRemovalsError() + }) + return reservationId, bucketStart, err +} + +// CancelBulkClientRemovalReservation releases a reservation made by +// ReserveBulkClientRemovalSlot that the caller ultimately couldn't use -- +// e.g. the network turned out to already have a bulk-delete run in +// progress, discovered only after the slot was reserved (the target bucket +// has to be known before scheduling the task, so reservation necessarily +// happens before that check). +func CancelBulkClientRemovalReservation(ctx context.Context, reservationId server.Id) { + server.Tx(ctx, func(tx server.PgTx) { + server.RaisePgResult(tx.Exec( + ctx, + `DELETE FROM bulk_client_removal_quota WHERE bulk_client_removal_quota_id = $1`, + reservationId, + )) + }) +} + +// RemoveExpiredBulkClientRemovalQuota deletes quota ledger rows whose bucket +// is entirely in the past relative to minBucketStart. Callers should pass a +// cutoff safely behind the lookahead window (e.g. now minus a couple of +// days), since a bucket up to MaxBulkClientRemovalLookaheadBuckets hours in +// the future can still be actively reserved against. +func RemoveExpiredBulkClientRemovalQuota(ctx context.Context, minBucketStart time.Time) { + server.MaintenanceTx(ctx, func(tx server.PgTx) { + server.RaisePgResult(tx.Exec( + ctx, + `DELETE FROM bulk_client_removal_quota WHERE bucket_start < $1`, + minBucketStart.UTC(), + )) + }) +} diff --git a/model/bulk_client_removal_rate_limit_test.go b/model/bulk_client_removal_rate_limit_test.go new file mode 100644 index 00000000..f21ccdbc --- /dev/null +++ b/model/bulk_client_removal_rate_limit_test.go @@ -0,0 +1,187 @@ +package model + +import ( + "context" + "testing" + "time" + + "github.com/go-playground/assert/v2" + + "github.com/urnetwork/server" + "github.com/urnetwork/server/jwt" + "github.com/urnetwork/server/session" +) + +// The budget is global, not per-network: two different networks draw from +// the same shared per-bucket ceiling. +func TestReserveBulkClientRemovalSlotFillsCurrentBucketThenQueuesNextHour(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkIdA := server.NewId() + networkIdB := server.NewId() + + currentBucket := bulkClientRemovalBucketStart(server.NowUtc()) + + _, bucketStart, err := ReserveBulkClientRemovalSlot(ctx, networkIdA, MaxBulkClientRemovalsPerBucket-1) + assert.Equal(t, err, nil) + assert.Equal(t, bucketStart, currentBucket) + + // 1 more from a different network fits exactly at the current + // bucket's ceiling + _, bucketStart, err = ReserveBulkClientRemovalSlot(ctx, networkIdB, 1) + assert.Equal(t, err, nil) + assert.Equal(t, bucketStart, currentBucket) + + // the current bucket is now fully spent; the next reservation must + // be queued for the next hour instead of being rejected + _, bucketStart, err = ReserveBulkClientRemovalSlot(ctx, networkIdA, 1) + assert.Equal(t, err, nil) + assert.Equal(t, bucketStart, currentBucket.Add(BulkClientRemovalBucketDuration)) + }) +} + +// Once every bucket in the whole lookahead window (the deployment's daily +// budget) is spent, a new reservation must be rejected outright, not queued +// indefinitely -- this is the hard daily cap. +func TestReserveBulkClientRemovalSlotRejectsWhenDailyBudgetExhausted(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + + for i := 0; i < MaxBulkClientRemovalLookaheadBuckets; i++ { + _, _, err := ReserveBulkClientRemovalSlot(ctx, server.NewId(), MaxBulkClientRemovalsPerBucket) + assert.Equal(t, err, nil) + } + + // every bucket in the lookahead window is now full + _, _, err := ReserveBulkClientRemovalSlot(ctx, server.NewId(), 1) + assert.NotEqual(t, err, nil) + }) +} + +// A cancelled reservation must free its slot back up -- e.g. for a request +// that turns out to be rejected for an unrelated reason (AlreadyInProgress) +// after the slot was already reserved. +func TestCancelBulkClientRemovalReservationFreesSlot(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkId := server.NewId() + + currentBucket := bulkClientRemovalBucketStart(server.NowUtc()) + + reservationId, bucketStart, err := ReserveBulkClientRemovalSlot(ctx, networkId, MaxBulkClientRemovalsPerBucket) + assert.Equal(t, err, nil) + assert.Equal(t, bucketStart, currentBucket) + + CancelBulkClientRemovalReservation(ctx, reservationId) + + // the current bucket's full ceiling must be available again + _, bucketStart, err = ReserveBulkClientRemovalSlot(ctx, networkId, MaxBulkClientRemovalsPerBucket) + assert.Equal(t, err, nil) + assert.Equal(t, bucketStart, currentBucket) + }) +} + +// RemoveExpiredBulkClientRemovalQuota must only remove buckets safely +// before the cutoff, leaving buckets at or after it untouched. +func TestRemoveExpiredBulkClientRemovalQuota(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + currentBucket := bulkClientRemovalBucketStart(server.NowUtc()) + oldBucket := currentBucket.Add(-72 * time.Hour) + + server.Tx(ctx, func(tx server.PgTx) { + server.RaisePgResult(tx.Exec( + ctx, + ` + INSERT INTO bulk_client_removal_quota (bulk_client_removal_quota_id, network_id, client_count, bucket_start, create_time) + VALUES ($1, $2, $3, $4, $5) + `, + server.NewId(), server.NewId(), 5, oldBucket, server.NowUtc(), + )) + }) + + _, _, err := ReserveBulkClientRemovalSlot(ctx, server.NewId(), 7) + assert.Equal(t, err, nil) + + RemoveExpiredBulkClientRemovalQuota(ctx, currentBucket.Add(-48*time.Hour)) + + var total int + server.Db(ctx, func(conn server.PgConn) { + result, qerr := conn.Query(ctx, `SELECT COALESCE(SUM(client_count), 0) FROM bulk_client_removal_quota`) + server.WithPgResult(result, qerr, func() { + if result.Next() { + server.Raise(result.Scan(&total)) + } + }) + }) + // the old (72h back) row must be gone; the current-bucket row (7) + // from ReserveBulkClientRemovalSlot must remain + assert.Equal(t, total, 7) + }) +} + +// RemoveNetworkClients must actually surface the daily-cap rejection as an +// error to the caller, not just leave the reservation mechanism correct in +// isolation. Filling every bucket in the lookahead window directly (cheap) +// avoids needing to construct 24 million-id requests just to exercise this +// wiring. +func TestRemoveNetworkClientsSurfacesDailyBudgetRejection(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + for i := 0; i < MaxBulkClientRemovalLookaheadBuckets; i++ { + _, _, err := ReserveBulkClientRemovalSlot(ctx, server.NewId(), MaxBulkClientRemovalsPerBucket) + assert.Equal(t, err, nil) + } + + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{NetworkId: server.NewId()}, + } + + _, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: []server.Id{server.NewId(), server.NewId()}, + }, sess) + assert.NotEqual(t, err, nil) + }) +} + +// When the current hour's budget is already spent, RemoveNetworkClients +// must queue the request for the next hour -- even a small request that +// would normally run synchronously -- rather than reject it. This is the +// behavior the user asked for: wait for the next hour instead of failing. +func TestRemoveNetworkClientsQueuesSmallRequestWhenCurrentHourIsFull(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + currentBucket := bulkClientRemovalBucketStart(server.NowUtc()) + + _, _, err := ReserveBulkClientRemovalSlot(ctx, server.NewId(), MaxBulkClientRemovalsPerBucket) + assert.Equal(t, err, nil) + + networkId := server.NewId() + deviceId := server.NewId() + clientId := server.NewId() + Testing_CreateDevice(ctx, networkId, deviceId, clientId, "test", "test") + + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{NetworkId: networkId}, + } + + result, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: []server.Id{clientId}, + }, sess) + assert.Equal(t, err, nil) + assert.Equal(t, result.Scheduled, true) + assert.NotEqual(t, result.ScheduledFor, nil) + + // deferred to the next hour, jittered somewhere within it -- not + // exactly the hour boundary, so check the range rather than an + // exact timestamp + nextBucket := currentBucket.Add(BulkClientRemovalBucketDuration) + assert.Equal(t, !result.ScheduledFor.Before(nextBucket), true) + assert.Equal(t, result.ScheduledFor.Before(nextBucket.Add(BulkClientRemovalBucketDuration)), true) + + // not applied yet -- it's queued for the next hour, not run now + assert.NotEqual(t, GetNetworkClient(ctx, clientId), nil) + }) +} diff --git a/model/device_association_model_test.go b/model/device_association_model_test.go index 186a00eb..476dc7c2 100644 --- a/model/device_association_model_test.go +++ b/model/device_association_model_test.go @@ -6,7 +6,7 @@ import ( "context" "testing" - "github.com/urnetwork/connect" + "github.com/go-playground/assert/v2" "github.com/urnetwork/server" "github.com/urnetwork/server/jwt" @@ -40,8 +40,8 @@ func TestDeviceAdopt(t *testing.T) { Testing_CreateDevice(ctx, networkIdA, deviceIdA, clientIdA, "devicea", "speca") clientsResult0, err := GetNetworkClients(clientSessionA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(clientsResult0.Clients), 1) + assert.Equal(t, err, nil) + assert.Equal(t, len(clientsResult0.Clients), 1) result1, err := DeviceCreateAdoptCode( &DeviceCreateAdoptCodeArgs{ @@ -50,10 +50,10 @@ func TestDeviceAdopt(t *testing.T) { }, clientSessionNoAuth, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, result1.Error, nil) - connect.AssertNotEqual(t, result1.AdoptCode, "") - connect.AssertNotEqual(t, result1.AdoptSecret, "") + assert.Equal(t, err, nil) + assert.Equal(t, result1.Error, nil) + assert.NotEqual(t, result1.AdoptCode, "") + assert.NotEqual(t, result1.AdoptSecret, "") qrResult0, err := DeviceAdoptCodeQR( &DeviceAdoptCodeQRArgs{ @@ -61,8 +61,8 @@ func TestDeviceAdopt(t *testing.T) { }, clientSessionNoAuth, ) - connect.AssertEqual(t, err, nil) - connect.AssertNotEqual(t, qrResult0.PngBytes, nil) + assert.Equal(t, err, nil) + assert.NotEqual(t, qrResult0.PngBytes, nil) result2, err := DeviceAdd( &DeviceAddArgs{ @@ -70,8 +70,8 @@ func TestDeviceAdopt(t *testing.T) { }, clientSessionA, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, result2.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, result2.Error, nil) result3, err := DeviceAdoptStatus( &DeviceAdoptStatusArgs{ @@ -79,15 +79,15 @@ func TestDeviceAdopt(t *testing.T) { }, clientSessionNoAuth, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, result3.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, result3.Error, nil) // at this point there should be adopt association associationResult0, err := DeviceAssociations(clientSessionA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(associationResult0.PendingAdoptionDevices), 1) - connect.AssertEqual(t, len(associationResult0.IncomingSharedDevices), 0) - connect.AssertEqual(t, len(associationResult0.OutgoingSharedDevices), 0) + assert.Equal(t, err, nil) + assert.Equal(t, len(associationResult0.PendingAdoptionDevices), 1) + assert.Equal(t, len(associationResult0.IncomingSharedDevices), 0) + assert.Equal(t, len(associationResult0.OutgoingSharedDevices), 0) result4, err := DeviceConfirmAdopt( &DeviceConfirmAdoptArgs{ @@ -97,27 +97,27 @@ func TestDeviceAdopt(t *testing.T) { }, clientSessionNoAuth, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, result4.Error, nil) - // connect.AssertEqual(t, result4.AssociatedNetworkName, "a") - connect.AssertNotEqual(t, result4.ByClientJwt, "") + assert.Equal(t, err, nil) + assert.Equal(t, result4.Error, nil) + // assert.Equal(t, result4.AssociatedNetworkName, "a") + assert.NotEqual(t, result4.ByClientJwt, "") byJwt, err := jwt.ParseByJwt(ctx, result4.ByClientJwt) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, byJwt.NetworkId, networkIdA) - connect.AssertEqual(t, byJwt.NetworkName, "a") - connect.AssertEqual(t, byJwt.UserId, userIdA) - connect.AssertEqual(t, byJwt.GuestMode, false) + assert.Equal(t, err, nil) + assert.Equal(t, byJwt.NetworkId, networkIdA) + assert.Equal(t, byJwt.NetworkName, "a") + assert.Equal(t, byJwt.UserId, userIdA) + assert.Equal(t, byJwt.GuestMode, false) // at this point there should be no adopt association associationResult1, err := DeviceAssociations(clientSessionA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(associationResult1.PendingAdoptionDevices), 0) - connect.AssertEqual(t, len(associationResult1.IncomingSharedDevices), 0) - connect.AssertEqual(t, len(associationResult1.OutgoingSharedDevices), 0) + assert.Equal(t, err, nil) + assert.Equal(t, len(associationResult1.PendingAdoptionDevices), 0) + assert.Equal(t, len(associationResult1.IncomingSharedDevices), 0) + assert.Equal(t, len(associationResult1.OutgoingSharedDevices), 0) clientResult1, err := GetNetworkClients(clientSessionA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(clientResult1.Clients), 2) + assert.Equal(t, err, nil) + assert.Equal(t, len(clientResult1.Clients), 2) }) } @@ -163,30 +163,30 @@ func TestDeviceConfirmAdoptWrongSecretRejected(t *testing.T) { &DeviceCreateAdoptCodeArgs{DeviceName: "devicec", DeviceSpec: "specc"}, noAuthSession, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, created.Error, nil) - connect.AssertNotEqual(t, created.AdoptCode, "") - connect.AssertNotEqual(t, created.AdoptSecret, "") + assert.Equal(t, err, nil) + assert.Equal(t, created.Error, nil) + assert.NotEqual(t, created.AdoptCode, "") + assert.NotEqual(t, created.AdoptSecret, "") // Step 3: the victim adopts the (public) code. added, err := DeviceAdd(&DeviceAddArgs{Code: created.AdoptCode}, victimSession) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, added.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, added.Error, nil) // Step 4: unauthenticated adopt-status leaks the victim's network name. status, err := DeviceAdoptStatus( &DeviceAdoptStatusArgs{AdoptCode: created.AdoptCode}, noAuthSession, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, status.Error, nil) - connect.AssertEqual(t, status.AssociatedNetworkName, "a") + assert.Equal(t, err, nil) + assert.Equal(t, status.Error, nil) + assert.Equal(t, status.AssociatedNetworkName, "a") // Step 5: confirm with the leaked network name but a DELIBERATELY WRONG // secret. A legitimate device proves possession of adopt_secret here; an // attacker who only eavesdropped the code does not have it. wrongSecret := "00000000000000000000000000000000deadbeefdeadbeef" - connect.AssertNotEqual(t, wrongSecret, created.AdoptSecret) + assert.NotEqual(t, wrongSecret, created.AdoptSecret) confirmed, err := DeviceConfirmAdopt( &DeviceConfirmAdoptArgs{ @@ -196,7 +196,7 @@ func TestDeviceConfirmAdoptWrongSecretRejected(t *testing.T) { }, noAuthSession, ) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) // SECURE expectation: a wrong secret must not yield a client credential. if confirmed != nil && confirmed.ByClientJwt != "" { @@ -204,17 +204,17 @@ func TestDeviceConfirmAdoptWrongSecretRejected(t *testing.T) { if parseErr == nil { t.Fatalf("SECURITY (VDP1): /device/confirm-adopt minted a client JWT with a WRONG "+ "adopt_secret — authentication bypass. Minted credential networkId=%s networkName=%q "+ - "userId=%s guestMode=%v. DeviceConfirmAdopt never references confirmAdopt.AdoptSecret; "+ + "userId=%s. DeviceConfirmAdopt never references confirmAdopt.AdoptSecret; "+ "add `AND adopt_secret = $N` (constant-time) to the device_adopt UPDATE, as "+ "DeviceRemoveAdoptCode already does.", - byJwt.NetworkId, byJwt.NetworkName, byJwt.UserId, byJwt.GuestMode) + byJwt.NetworkId, byJwt.NetworkName, byJwt.UserId) } t.Fatalf("SECURITY (VDP1): /device/confirm-adopt returned a by_client_jwt with a WRONG " + "adopt_secret — authentication bypass.") } // After the fix, confirm affects 0 rows and returns the generic error. - connect.AssertNotEqual(t, confirmed.Error, nil) + assert.NotEqual(t, confirmed.Error, nil) }) } @@ -248,8 +248,8 @@ func TestDeviceAdoptPartialOfferRemove(t *testing.T) { Testing_CreateDevice(ctx, networkIdA, deviceIdA, clientIdA, "devicea", "speca") clientsResult0, err := GetNetworkClients(clientSessionA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(clientsResult0.Clients), 1) + assert.Equal(t, err, nil) + assert.Equal(t, len(clientsResult0.Clients), 1) result1, err := DeviceCreateAdoptCode( &DeviceCreateAdoptCodeArgs{ @@ -258,10 +258,10 @@ func TestDeviceAdoptPartialOfferRemove(t *testing.T) { }, clientSessionNoAuth, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, result1.Error, nil) - connect.AssertNotEqual(t, result1.AdoptCode, "") - connect.AssertNotEqual(t, result1.AdoptSecret, "") + assert.Equal(t, err, nil) + assert.Equal(t, result1.Error, nil) + assert.NotEqual(t, result1.AdoptCode, "") + assert.NotEqual(t, result1.AdoptSecret, "") result2, err := DeviceAdd( &DeviceAddArgs{ @@ -269,8 +269,8 @@ func TestDeviceAdoptPartialOfferRemove(t *testing.T) { }, clientSessionA, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, result2.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, result2.Error, nil) result3, err := DeviceAdoptStatus( &DeviceAdoptStatusArgs{ @@ -278,15 +278,15 @@ func TestDeviceAdoptPartialOfferRemove(t *testing.T) { }, clientSessionNoAuth, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, result3.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, result3.Error, nil) // at this point there should be adopt association associationResult0, err := DeviceAssociations(clientSessionA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(associationResult0.PendingAdoptionDevices), 1) - connect.AssertEqual(t, len(associationResult0.IncomingSharedDevices), 0) - connect.AssertEqual(t, len(associationResult0.OutgoingSharedDevices), 0) + assert.Equal(t, err, nil) + assert.Equal(t, len(associationResult0.PendingAdoptionDevices), 1) + assert.Equal(t, len(associationResult0.IncomingSharedDevices), 0) + assert.Equal(t, len(associationResult0.OutgoingSharedDevices), 0) removeResult0, err := DeviceRemoveAdoptCode( &DeviceRemoveAdoptCodeArgs{ @@ -295,14 +295,14 @@ func TestDeviceAdoptPartialOfferRemove(t *testing.T) { }, clientSessionNoAuth, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, removeResult0.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, removeResult0.Error, nil) associationResult1, err := DeviceAssociations(clientSessionA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(associationResult1.PendingAdoptionDevices), 0) - connect.AssertEqual(t, len(associationResult1.IncomingSharedDevices), 0) - connect.AssertEqual(t, len(associationResult1.OutgoingSharedDevices), 0) + assert.Equal(t, err, nil) + assert.Equal(t, len(associationResult1.PendingAdoptionDevices), 0) + assert.Equal(t, len(associationResult1.IncomingSharedDevices), 0) + assert.Equal(t, len(associationResult1.OutgoingSharedDevices), 0) }) } @@ -337,8 +337,8 @@ func TestDeviceAdoptPartialOwnerRemove(t *testing.T) { Testing_CreateDevice(ctx, networkIdA, deviceIdA, clientIdA, "devicea", "speca") clientsResult0, err := GetNetworkClients(clientSessionA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(clientsResult0.Clients), 1) + assert.Equal(t, err, nil) + assert.Equal(t, len(clientsResult0.Clients), 1) result1, err := DeviceCreateAdoptCode( &DeviceCreateAdoptCodeArgs{ @@ -347,10 +347,10 @@ func TestDeviceAdoptPartialOwnerRemove(t *testing.T) { }, clientSessionNoAuth, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, result1.Error, nil) - connect.AssertNotEqual(t, result1.AdoptCode, "") - connect.AssertNotEqual(t, result1.AdoptSecret, "") + assert.Equal(t, err, nil) + assert.Equal(t, result1.Error, nil) + assert.NotEqual(t, result1.AdoptCode, "") + assert.NotEqual(t, result1.AdoptSecret, "") result2, err := DeviceAdd( &DeviceAddArgs{ @@ -358,8 +358,8 @@ func TestDeviceAdoptPartialOwnerRemove(t *testing.T) { }, clientSessionA, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, result2.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, result2.Error, nil) result3, err := DeviceAdoptStatus( &DeviceAdoptStatusArgs{ @@ -367,15 +367,15 @@ func TestDeviceAdoptPartialOwnerRemove(t *testing.T) { }, clientSessionNoAuth, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, result3.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, result3.Error, nil) // at this point there should be adopt association associationResult0, err := DeviceAssociations(clientSessionA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(associationResult0.PendingAdoptionDevices), 1) - connect.AssertEqual(t, len(associationResult0.IncomingSharedDevices), 0) - connect.AssertEqual(t, len(associationResult0.OutgoingSharedDevices), 0) + assert.Equal(t, err, nil) + assert.Equal(t, len(associationResult0.PendingAdoptionDevices), 1) + assert.Equal(t, len(associationResult0.IncomingSharedDevices), 0) + assert.Equal(t, len(associationResult0.OutgoingSharedDevices), 0) removeResult0, err := DeviceRemoveAssociation( &DeviceRemoveAssociationArgs{ @@ -383,14 +383,14 @@ func TestDeviceAdoptPartialOwnerRemove(t *testing.T) { }, clientSessionA, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, removeResult0.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, removeResult0.Error, nil) associationResult1, err := DeviceAssociations(clientSessionA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(associationResult1.PendingAdoptionDevices), 0) - connect.AssertEqual(t, len(associationResult1.IncomingSharedDevices), 0) - connect.AssertEqual(t, len(associationResult1.OutgoingSharedDevices), 0) + assert.Equal(t, err, nil) + assert.Equal(t, len(associationResult1.PendingAdoptionDevices), 0) + assert.Equal(t, len(associationResult1.IncomingSharedDevices), 0) + assert.Equal(t, len(associationResult1.OutgoingSharedDevices), 0) }) } @@ -430,12 +430,12 @@ func TestDeviceShare(t *testing.T) { Testing_CreateDevice(ctx, networkIdB, deviceIdB, clientIdB, "deviceb", "specb") clientsResult0, err := GetNetworkClients(clientSessionA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(clientsResult0.Clients), 1) + assert.Equal(t, err, nil) + assert.Equal(t, len(clientsResult0.Clients), 1) clientsResult1, err := GetNetworkClients(clientSessionB) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(clientsResult1.Clients), 1) + assert.Equal(t, err, nil) + assert.Equal(t, len(clientsResult1.Clients), 1) result1, err := DeviceCreateShareCode( &DeviceCreateShareCodeArgs{ @@ -444,18 +444,18 @@ func TestDeviceShare(t *testing.T) { }, clientSessionA, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, result1.Error, nil) - connect.AssertNotEqual(t, result1.ShareCode, "") + assert.Equal(t, err, nil) + assert.Equal(t, result1.Error, nil) + assert.NotEqual(t, result1.ShareCode, "") associationResult0, err := DeviceAssociations(clientSessionA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(associationResult0.PendingAdoptionDevices), 0) - connect.AssertEqual(t, len(associationResult0.IncomingSharedDevices), 0) - connect.AssertEqual(t, len(associationResult0.OutgoingSharedDevices), 1) - connect.AssertEqual(t, associationResult0.OutgoingSharedDevices[0].Pending, true) - connect.AssertEqual(t, associationResult0.OutgoingSharedDevices[0].NetworkName, "") - connect.AssertEqual(t, associationResult0.OutgoingSharedDevices[0].DeviceName, "devicea") + assert.Equal(t, err, nil) + assert.Equal(t, len(associationResult0.PendingAdoptionDevices), 0) + assert.Equal(t, len(associationResult0.IncomingSharedDevices), 0) + assert.Equal(t, len(associationResult0.OutgoingSharedDevices), 1) + assert.Equal(t, associationResult0.OutgoingSharedDevices[0].Pending, true) + assert.Equal(t, associationResult0.OutgoingSharedDevices[0].NetworkName, "") + assert.Equal(t, associationResult0.OutgoingSharedDevices[0].DeviceName, "devicea") shareStatus1, err := DeviceShareStatus( &DeviceShareStatusArgs{ @@ -463,8 +463,8 @@ func TestDeviceShare(t *testing.T) { }, clientSessionA, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, shareStatus1.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, shareStatus1.Error, nil) qrResult0, err := DeviceShareCodeQR( &DeviceShareCodeQRArgs{ @@ -472,8 +472,8 @@ func TestDeviceShare(t *testing.T) { }, clientSessionA, ) - connect.AssertEqual(t, err, nil) - connect.AssertNotEqual(t, qrResult0.PngBytes, nil) + assert.Equal(t, err, nil) + assert.NotEqual(t, qrResult0.PngBytes, nil) result2, err := DeviceAdd( &DeviceAddArgs{ @@ -481,8 +481,8 @@ func TestDeviceShare(t *testing.T) { }, clientSessionB, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, result2.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, result2.Error, nil) result3, err := DeviceShareStatus( &DeviceShareStatusArgs{ @@ -490,28 +490,28 @@ func TestDeviceShare(t *testing.T) { }, clientSessionA, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, result3.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, result3.Error, nil) // at this point there should be share association pending associationResult1, err := DeviceAssociations(clientSessionA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(associationResult1.PendingAdoptionDevices), 0) - connect.AssertEqual(t, len(associationResult1.IncomingSharedDevices), 0) - connect.AssertEqual(t, len(associationResult1.OutgoingSharedDevices), 1) - connect.AssertEqual(t, associationResult1.OutgoingSharedDevices[0].Pending, true) - connect.AssertEqual(t, associationResult1.OutgoingSharedDevices[0].NetworkName, "b") - connect.AssertEqual(t, associationResult1.OutgoingSharedDevices[0].DeviceName, "devicea") + assert.Equal(t, err, nil) + assert.Equal(t, len(associationResult1.PendingAdoptionDevices), 0) + assert.Equal(t, len(associationResult1.IncomingSharedDevices), 0) + assert.Equal(t, len(associationResult1.OutgoingSharedDevices), 1) + assert.Equal(t, associationResult1.OutgoingSharedDevices[0].Pending, true) + assert.Equal(t, associationResult1.OutgoingSharedDevices[0].NetworkName, "b") + assert.Equal(t, associationResult1.OutgoingSharedDevices[0].DeviceName, "devicea") associationResult2, err := DeviceAssociations(clientSessionB) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(associationResult2.PendingAdoptionDevices), 0) - connect.AssertEqual(t, len(associationResult2.IncomingSharedDevices), 1) - connect.AssertEqual(t, len(associationResult2.OutgoingSharedDevices), 0) - connect.AssertEqual(t, associationResult2.IncomingSharedDevices[0].Pending, true) - connect.AssertEqual(t, associationResult2.IncomingSharedDevices[0].NetworkName, "a") - connect.AssertEqual(t, associationResult2.IncomingSharedDevices[0].DeviceName, "devicea") + assert.Equal(t, err, nil) + assert.Equal(t, len(associationResult2.PendingAdoptionDevices), 0) + assert.Equal(t, len(associationResult2.IncomingSharedDevices), 1) + assert.Equal(t, len(associationResult2.OutgoingSharedDevices), 0) + assert.Equal(t, associationResult2.IncomingSharedDevices[0].Pending, true) + assert.Equal(t, associationResult2.IncomingSharedDevices[0].NetworkName, "a") + assert.Equal(t, associationResult2.IncomingSharedDevices[0].DeviceName, "devicea") result4, err := DeviceConfirmShare( &DeviceConfirmShareArgs{ @@ -520,28 +520,28 @@ func TestDeviceShare(t *testing.T) { }, clientSessionA, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, result4.Error, nil) - connect.AssertEqual(t, result4.AssociatedNetworkName, "b") + assert.Equal(t, err, nil) + assert.Equal(t, result4.Error, nil) + assert.Equal(t, result4.AssociatedNetworkName, "b") // at this point there should be a share association not pending, as outgoing for A, incoming for B associationResult3, err := DeviceAssociations(clientSessionA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(associationResult3.PendingAdoptionDevices), 0) - connect.AssertEqual(t, len(associationResult3.IncomingSharedDevices), 0) - connect.AssertEqual(t, len(associationResult3.OutgoingSharedDevices), 1) - connect.AssertEqual(t, associationResult3.OutgoingSharedDevices[0].Pending, false) - connect.AssertEqual(t, associationResult3.OutgoingSharedDevices[0].NetworkName, "b") - connect.AssertEqual(t, associationResult3.OutgoingSharedDevices[0].DeviceName, "devicea") + assert.Equal(t, err, nil) + assert.Equal(t, len(associationResult3.PendingAdoptionDevices), 0) + assert.Equal(t, len(associationResult3.IncomingSharedDevices), 0) + assert.Equal(t, len(associationResult3.OutgoingSharedDevices), 1) + assert.Equal(t, associationResult3.OutgoingSharedDevices[0].Pending, false) + assert.Equal(t, associationResult3.OutgoingSharedDevices[0].NetworkName, "b") + assert.Equal(t, associationResult3.OutgoingSharedDevices[0].DeviceName, "devicea") associationResult4, err := DeviceAssociations(clientSessionB) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(associationResult4.PendingAdoptionDevices), 0) - connect.AssertEqual(t, len(associationResult4.IncomingSharedDevices), 1) - connect.AssertEqual(t, len(associationResult4.OutgoingSharedDevices), 0) - connect.AssertEqual(t, associationResult4.IncomingSharedDevices[0].Pending, false) - connect.AssertEqual(t, associationResult4.IncomingSharedDevices[0].NetworkName, "a") - connect.AssertEqual(t, associationResult4.IncomingSharedDevices[0].DeviceName, "devicea") + assert.Equal(t, err, nil) + assert.Equal(t, len(associationResult4.PendingAdoptionDevices), 0) + assert.Equal(t, len(associationResult4.IncomingSharedDevices), 1) + assert.Equal(t, len(associationResult4.OutgoingSharedDevices), 0) + assert.Equal(t, associationResult4.IncomingSharedDevices[0].Pending, false) + assert.Equal(t, associationResult4.IncomingSharedDevices[0].NetworkName, "a") + assert.Equal(t, associationResult4.IncomingSharedDevices[0].DeviceName, "devicea") setNameResult1, err := DeviceSetAssociationName( &DeviceSetAssociationNameArgs{ @@ -550,8 +550,8 @@ func TestDeviceShare(t *testing.T) { }, clientSessionA, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, setNameResult1.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, setNameResult1.Error, nil) setNameResult2, err := DeviceSetAssociationName( &DeviceSetAssociationNameArgs{ @@ -560,34 +560,34 @@ func TestDeviceShare(t *testing.T) { }, clientSessionB, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, setNameResult2.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, setNameResult2.Error, nil) associationResult5, err := DeviceAssociations(clientSessionA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(associationResult5.PendingAdoptionDevices), 0) - connect.AssertEqual(t, len(associationResult5.IncomingSharedDevices), 0) - connect.AssertEqual(t, len(associationResult5.OutgoingSharedDevices), 1) - connect.AssertEqual(t, associationResult5.OutgoingSharedDevices[0].Pending, false) - connect.AssertEqual(t, associationResult5.OutgoingSharedDevices[0].NetworkName, "b") - connect.AssertEqual(t, associationResult5.OutgoingSharedDevices[0].DeviceName, "That device I shared") + assert.Equal(t, err, nil) + assert.Equal(t, len(associationResult5.PendingAdoptionDevices), 0) + assert.Equal(t, len(associationResult5.IncomingSharedDevices), 0) + assert.Equal(t, len(associationResult5.OutgoingSharedDevices), 1) + assert.Equal(t, associationResult5.OutgoingSharedDevices[0].Pending, false) + assert.Equal(t, associationResult5.OutgoingSharedDevices[0].NetworkName, "b") + assert.Equal(t, associationResult5.OutgoingSharedDevices[0].DeviceName, "That device I shared") associationResult6, err := DeviceAssociations(clientSessionB) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(associationResult6.PendingAdoptionDevices), 0) - connect.AssertEqual(t, len(associationResult6.IncomingSharedDevices), 1) - connect.AssertEqual(t, len(associationResult6.OutgoingSharedDevices), 0) - connect.AssertEqual(t, associationResult6.IncomingSharedDevices[0].Pending, false) - connect.AssertEqual(t, associationResult6.IncomingSharedDevices[0].NetworkName, "a") - connect.AssertEqual(t, associationResult6.IncomingSharedDevices[0].DeviceName, "My new device from friend") + assert.Equal(t, err, nil) + assert.Equal(t, len(associationResult6.PendingAdoptionDevices), 0) + assert.Equal(t, len(associationResult6.IncomingSharedDevices), 1) + assert.Equal(t, len(associationResult6.OutgoingSharedDevices), 0) + assert.Equal(t, associationResult6.IncomingSharedDevices[0].Pending, false) + assert.Equal(t, associationResult6.IncomingSharedDevices[0].NetworkName, "a") + assert.Equal(t, associationResult6.IncomingSharedDevices[0].DeviceName, "My new device from friend") clientsResult2, err := GetNetworkClients(clientSessionA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(clientsResult2.Clients), 1) + assert.Equal(t, err, nil) + assert.Equal(t, len(clientsResult2.Clients), 1) clientsResult3, err := GetNetworkClients(clientSessionB) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(clientsResult3.Clients), 1) + assert.Equal(t, err, nil) + assert.Equal(t, len(clientsResult3.Clients), 1) }) } @@ -619,11 +619,11 @@ func TestDeviceConfirmAdoptWrongNetworkNameRejected(t *testing.T) { created, err := DeviceCreateAdoptCode( &DeviceCreateAdoptCodeArgs{DeviceName: "d", DeviceSpec: "s"}, noAuth) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) added, err := DeviceAdd(&DeviceAddArgs{Code: created.AdoptCode}, victim) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, added.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, added.Error, nil) confirmed, err := DeviceConfirmAdopt( &DeviceConfirmAdoptArgs{ @@ -631,9 +631,9 @@ func TestDeviceConfirmAdoptWrongNetworkNameRejected(t *testing.T) { AdoptSecret: created.AdoptSecret, AssociatedNetworkName: "some-other-network", }, noAuth) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, confirmed.ByClientJwt, "") - connect.AssertNotEqual(t, confirmed.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, confirmed.ByClientJwt, "") + assert.NotEqual(t, confirmed.Error, nil) }) } @@ -646,7 +646,7 @@ func TestDeviceConfirmAdoptBeforeAdoptRejected(t *testing.T) { created, err := DeviceCreateAdoptCode( &DeviceCreateAdoptCodeArgs{DeviceName: "d", DeviceSpec: "s"}, noAuth) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) confirmed, err := DeviceConfirmAdopt( &DeviceConfirmAdoptArgs{ @@ -654,9 +654,9 @@ func TestDeviceConfirmAdoptBeforeAdoptRejected(t *testing.T) { AdoptSecret: created.AdoptSecret, AssociatedNetworkName: "a", }, noAuth) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, confirmed.ByClientJwt, "") - connect.AssertNotEqual(t, confirmed.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, confirmed.ByClientJwt, "") + assert.NotEqual(t, confirmed.Error, nil) }) } @@ -669,10 +669,10 @@ func TestDeviceConfirmAdoptReplayRejected(t *testing.T) { created, err := DeviceCreateAdoptCode( &DeviceCreateAdoptCodeArgs{DeviceName: "d", DeviceSpec: "s"}, noAuth) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) _, err = DeviceAdd(&DeviceAddArgs{Code: created.AdoptCode}, victim) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) args := &DeviceConfirmAdoptArgs{ AdoptCode: created.AdoptCode, @@ -681,20 +681,20 @@ func TestDeviceConfirmAdoptReplayRejected(t *testing.T) { } first, err := DeviceConfirmAdopt(args, noAuth) - connect.AssertEqual(t, err, nil) - connect.AssertNotEqual(t, first.ByClientJwt, "") + assert.Equal(t, err, nil) + assert.NotEqual(t, first.ByClientJwt, "") clientsAfterFirst, err := GetNetworkClients(victim) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) second, err := DeviceConfirmAdopt(args, noAuth) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, second.ByClientJwt, "") - connect.AssertNotEqual(t, second.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, second.ByClientJwt, "") + assert.NotEqual(t, second.Error, nil) clientsAfterSecond, err := GetNetworkClients(victim) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(clientsAfterSecond.Clients), len(clientsAfterFirst.Clients)) + assert.Equal(t, err, nil) + assert.Equal(t, len(clientsAfterSecond.Clients), len(clientsAfterFirst.Clients)) }) } @@ -708,16 +708,16 @@ func TestDeviceAdoptCannotHijackAfterAdopt(t *testing.T) { created, err := DeviceCreateAdoptCode( &DeviceCreateAdoptCodeArgs{DeviceName: "d", DeviceSpec: "s"}, noAuth) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) addedA, err := DeviceAdd(&DeviceAddArgs{Code: created.AdoptCode}, netA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, addedA.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, addedA.Error, nil) // B tries to adopt the same, already-owned code. addedB, err := DeviceAdd(&DeviceAddArgs{Code: created.AdoptCode}, netB) - connect.AssertEqual(t, err, nil) - connect.AssertNotEqual(t, addedB.Error, nil) + assert.Equal(t, err, nil) + assert.NotEqual(t, addedB.Error, nil) // The device confirms for A (the true owner); the JWT is bound to A, not B. confirmed, err := DeviceConfirmAdopt( @@ -726,11 +726,11 @@ func TestDeviceAdoptCannotHijackAfterAdopt(t *testing.T) { AdoptSecret: created.AdoptSecret, AssociatedNetworkName: "a", }, noAuth) - connect.AssertEqual(t, err, nil) - connect.AssertNotEqual(t, confirmed.ByClientJwt, "") + assert.Equal(t, err, nil) + assert.NotEqual(t, confirmed.ByClientJwt, "") byJwt, err := jwt.ParseByJwt(ctx, confirmed.ByClientJwt) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, byJwt.NetworkId, networkIdA) + assert.Equal(t, err, nil) + assert.Equal(t, byJwt.NetworkId, networkIdA) }) } @@ -743,27 +743,27 @@ func TestDeviceRemoveAdoptCodeWrongSecretRejected(t *testing.T) { created, err := DeviceCreateAdoptCode( &DeviceCreateAdoptCodeArgs{DeviceName: "d", DeviceSpec: "s"}, noAuth) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) _, err = DeviceAdd(&DeviceAddArgs{Code: created.AdoptCode}, netA) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) // Wrong secret: must not remove. wrong, err := DeviceRemoveAdoptCode( &DeviceRemoveAdoptCodeArgs{AdoptCode: created.AdoptCode, AdoptSecret: "wrong-secret"}, noAuth) - connect.AssertEqual(t, err, nil) - connect.AssertNotEqual(t, wrong.Error, nil) + assert.Equal(t, err, nil) + assert.NotEqual(t, wrong.Error, nil) // Still present: A still sees a pending adoption. assoc, err := DeviceAssociations(netA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(assoc.PendingAdoptionDevices), 1) + assert.Equal(t, err, nil) + assert.Equal(t, len(assoc.PendingAdoptionDevices), 1) // Correct secret: removes. ok, err := DeviceRemoveAdoptCode( &DeviceRemoveAdoptCodeArgs{AdoptCode: created.AdoptCode, AdoptSecret: created.AdoptSecret}, noAuth) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, ok.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, ok.Error, nil) }) } @@ -781,9 +781,9 @@ func TestDeviceCreateShareCodeForeignClientRejected(t *testing.T) { // A tries to share B's client. result, err := DeviceCreateShareCode( &DeviceCreateShareCodeArgs{ClientId: clientIdB, DeviceName: "stolen"}, netA) - connect.AssertEqual(t, err, nil) - connect.AssertNotEqual(t, result.Error, nil) - connect.AssertEqual(t, result.ShareCode, "") + assert.Equal(t, err, nil) + assert.NotEqual(t, result.Error, nil) + assert.Equal(t, result.ShareCode, "") }) } @@ -799,12 +799,12 @@ func TestDeviceAddOwnShareCodeRejected(t *testing.T) { share, err := DeviceCreateShareCode( &DeviceCreateShareCodeArgs{ClientId: clientIdA, DeviceName: "devicea"}, netA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, share.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, share.Error, nil) added, err := DeviceAdd(&DeviceAddArgs{Code: share.ShareCode}, netA) - connect.AssertEqual(t, err, nil) - connect.AssertNotEqual(t, added.Error, nil) + assert.Equal(t, err, nil) + assert.NotEqual(t, added.Error, nil) }) } @@ -825,17 +825,17 @@ func TestDeviceConfirmShareByUnrelatedNetworkRejected(t *testing.T) { share, err := DeviceCreateShareCode( &DeviceCreateShareCodeArgs{ClientId: clientIdA, DeviceName: "devicea"}, netA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, share.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, share.Error, nil) added, err := DeviceAdd(&DeviceAddArgs{Code: share.ShareCode}, netB) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, added.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, added.Error, nil) // C, knowing only the share code and the guest's network name, confirms. confirmed, err := DeviceConfirmShare( &DeviceConfirmShareArgs{ShareCode: share.ShareCode, AssociatedNetworkName: "b"}, netC) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) if confirmed != nil && confirmed.Error == nil { t.Fatalf("SECURITY (adjacent): unrelated network C (neither source nor guest) confirmed a " + "device share via DeviceConfirmShare, which never checks clientSession.ByJwt.NetworkId. " + @@ -854,20 +854,20 @@ func TestDeviceRemoveAssociationByUnrelatedNetworkRejected(t *testing.T) { created, err := DeviceCreateAdoptCode( &DeviceCreateAdoptCodeArgs{DeviceName: "d", DeviceSpec: "s"}, noAuth) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) _, err = DeviceAdd(&DeviceAddArgs{Code: created.AdoptCode}, netA) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) // C (unrelated) tries to remove A's pending adoption. removed, err := DeviceRemoveAssociation( &DeviceRemoveAssociationArgs{Code: created.AdoptCode}, netC) - connect.AssertEqual(t, err, nil) - connect.AssertNotEqual(t, removed.Error, nil) + assert.Equal(t, err, nil) + assert.NotEqual(t, removed.Error, nil) // A still has the association. assoc, err := DeviceAssociations(netA) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, len(assoc.PendingAdoptionDevices), 1) + assert.Equal(t, err, nil) + assert.Equal(t, len(assoc.PendingAdoptionDevices), 1) }) } @@ -881,13 +881,13 @@ func TestDeviceSetAssociationNameByUnrelatedNetworkRejected(t *testing.T) { created, err := DeviceCreateAdoptCode( &DeviceCreateAdoptCodeArgs{DeviceName: "d", DeviceSpec: "s"}, noAuth) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) _, err = DeviceAdd(&DeviceAddArgs{Code: created.AdoptCode}, netA) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) result, err := DeviceSetAssociationName( &DeviceSetAssociationNameArgs{Code: created.AdoptCode, DeviceName: "hacked"}, netC) - connect.AssertEqual(t, err, nil) - connect.AssertNotEqual(t, result.Error, nil) + assert.Equal(t, err, nil) + assert.NotEqual(t, result.Error, nil) }) } diff --git a/model/guest_account_regression_test.go b/model/guest_account_regression_test.go new file mode 100644 index 00000000..5186d59e --- /dev/null +++ b/model/guest_account_regression_test.go @@ -0,0 +1,117 @@ +package model + +import ( + "context" + "testing" + + "github.com/go-playground/assert/v2" + "github.com/urnetwork/server" + "github.com/urnetwork/server/jwt" + "github.com/urnetwork/server/session" +) + +// Testing_CreateLegacyGuestNetwork inserts a network_user/network pair shaped +// exactly like the pre-removal networkCreateGuest function used to create: +// auth_type='guest', no user_auth, no password, no rows in any auth table. +// Guest signup is retired, but rows shaped like this still exist in +// production and this fork's code must not mistreat them. +func Testing_CreateLegacyGuestNetwork(ctx context.Context, networkId server.Id, userId server.Id) { + server.Tx(ctx, func(tx server.PgTx) { + server.RaisePgResult(tx.Exec( + ctx, + `INSERT INTO network_user (user_id, user_name, auth_type) VALUES ($1, $2, $3)`, + userId, "guest", "guest", + )) + server.RaisePgResult(tx.Exec( + ctx, + `INSERT INTO network (network_id, network_name, admin_user_id) VALUES ($1, $2, $3)`, + networkId, "g"+networkId.String(), userId, + )) + }) +} + +func TestHasAnyAuthMethodReflectsRealAuthMethods(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkId := server.NewId() + userId := server.NewId() + + Testing_CreateLegacyGuestNetwork(ctx, networkId, userId) + if has := HasAnyAuthMethod(ctx, userId); has { + t.Fatalf("expected fresh legacy guest to have no auth methods, HasAnyAuthMethod=true") + } + + userAuth := "guest-upgrade-test@example.com" + addResult, err := AddAuth(AddAuthMethod{ + UserAuth: &userAuth, + Password: strPtr("SomeValidPassword123!"), + }, session.Testing_CreateClientSession(ctx, &jwt.ByJwt{ + NetworkId: networkId, UserId: userId, NetworkName: "g" + networkId.String(), GuestMode: true, + })) + if err != nil { + t.Fatalf("AddAuth returned error: %v", err) + } + if addResult.Error != nil { + t.Fatalf("AddAuth returned soft error: %s", addResult.Error.Message) + } + + if has := HasAnyAuthMethod(ctx, userId); !has { + t.Fatalf("expected HasAnyAuthMethod=true after AddAuth succeeded, got false") + } + }) +} + +func TestValidateClientIdentityArgsBlocksBareGuestButAllowsUpgradedAccount(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkId := server.NewId() + userId := server.NewId() + + Testing_CreateLegacyGuestNetwork(ctx, networkId, userId) + + // a network-level (no ClientId) guest session, as a legacy guest's + // still-valid pre-refresh JWT would present -- must still be blocked + // from assigning explicit roles/principal + guestSession := session.Testing_CreateClientSession(ctx, &jwt.ByJwt{ + NetworkId: networkId, UserId: userId, NetworkName: "g" + networkId.String(), GuestMode: true, + }) + result, err := AuthNetworkClient(&AuthNetworkClientArgs{ + Description: "d", + DeviceSpec: "s", + Roles: []string{"admin"}, + Principal: "attacker-controlled-principal", + }, guestSession) + assert.Equal(t, err, nil) + if result.Error == nil { + t.Fatalf("expected a bare guest account to be rejected when assigning explicit roles/principal, got success") + } + + // the same account, after adding a real auth method (simulating a + // guest that upgraded via AddAuth), must now be allowed -- even + // though nothing here re-signed its JWT (GuestMode is still stale + // true), and even though network_user.auth_type is still 'guest' + // (AddAuth never updates it) + userAuth := "guest-upgrade-test-2@example.com" + _, err = AddAuth(AddAuthMethod{ + UserAuth: &userAuth, + Password: strPtr("SomeValidPassword123!"), + }, guestSession) + assert.Equal(t, err, nil) + + result, err = AuthNetworkClient(&AuthNetworkClientArgs{ + Description: "d2", + DeviceSpec: "s2", + Roles: []string{"admin"}, + Principal: "a-legitimate-principal", + }, guestSession) + assert.Equal(t, err, nil) + if result.Error != nil { + t.Fatalf("expected an upgraded (formerly-guest) account to be allowed to assign explicit roles/principal, got error: %s", result.Error.Message) + } + clientByJwt, err := jwt.ParseByJwt(ctx, *result.ByClientJwt) + assert.Equal(t, err, nil) + assert.Equal(t, clientByJwt.Principal, "a-legitimate-principal") + }) +} + +func strPtr(s string) *string { return &s } diff --git a/model/network_client_model.go b/model/network_client_model.go index 976b1d69..8ea0edab 100644 --- a/model/network_client_model.go +++ b/model/network_client_model.go @@ -6,6 +6,7 @@ import ( "encoding/json" // "crypto/sha256" "errors" + mathrand "math/rand" // "net" "net/netip" // "regexp" @@ -28,6 +29,7 @@ import ( "github.com/urnetwork/server" "github.com/urnetwork/server/jwt" "github.com/urnetwork/server/session" + "github.com/urnetwork/server/task" // "github.com/urnetwork/server/ulid" "github.com/urnetwork/connect" ) @@ -189,15 +191,24 @@ type ProxyAuthResult struct { } // validateClientIdentityArgs resolves the roles and principal a new client or -// auth code is created with. Explicit values require a network-level non-guest -// session; when omitted, the session's own roles and principal are inherited. +// auth code is created with. Explicit values require a network-level session +// on an account with at least one real auth method; when omitted, the +// session's own roles and principal are inherited. +// +// This checks HasAnyAuthMethod rather than session.ByJwt.GuestMode: guest +// signup is retired, but legacy guest accounts (auth_type = 'guest', zero +// rows in any auth table) still exist and their JWTs still carry a stale +// GuestMode claim that RefreshToken unconditionally zeroes out on the very +// next refresh regardless of real account state. A live auth-method check +// is the only signal that's both accurate for a still-bare guest and +// self-healing the moment that guest adds a real method via AddAuth. func validateClientIdentityArgs( roles []string, principal string, session *session.ClientSession, ) (resolvedRoles []string, resolvedPrincipal string, message string) { if 0 < len(roles) || principal != "" { - if session.ByJwt.GuestMode || session.ByJwt.ClientId != nil { + if session.ByJwt.ClientId != nil || !HasAnyAuthMethod(session.Ctx, session.ByJwt.UserId) { message = "Roles and principal can only be assigned by a network session." return } @@ -739,6 +750,352 @@ func RemoveNetworkClient( return removeClientResult, removeClientErr } +// matches the batch size `RemoveDisconnectedNetworkClients` already uses for +// bounded maintenance sweeps of this same table (see `markTopLevelBatchCount` +// above): large enough to make a real dent per transaction, small enough that +// no single transaction runs long or holds locks for long. +const RemoveNetworkClientsBatchCount = 10000 + +func removeNetworkClientsBatchExec(ctx context.Context, tx server.PgTx, clientIds []server.Id, networkId server.Id) { + _, err := tx.Exec( + ctx, + ` + UPDATE network_client + SET + active = false, + deactivate_time = $3 + WHERE + client_id = ANY($1) AND + network_id = $2 + `, + clientIds, + networkId, + server.NowUtc(), + ) + server.Raise(err) +} + +type RemoveNetworkClientsBatchArgs struct { + ClientIds []server.Id `json:"client_ids"` +} + +type RemoveNetworkClientsBatchResult struct{} + +// RemoveNetworkClientsBatch deactivates up to RemoveNetworkClientsBatchCount +// clients in a single transaction on the regular request pool. Callers with +// more ids than that should go through RemoveNetworkClients, which routes +// large requests to the background task instead of calling this directly. +// TxReadCommitted matches RemoveDisconnectedNetworkClients's isolation level +// for this same table, avoiding serialization-failure panics under +// concurrent overlapping updates that the pool's default isolation would risk. +func RemoveNetworkClientsBatch( + removeClients *RemoveNetworkClientsBatchArgs, + session *session.ClientSession, +) (*RemoveNetworkClientsBatchResult, error) { + server.Tx(session.Ctx, func(tx server.PgTx) { + removeNetworkClientsBatchExec(session.Ctx, tx, removeClients.ClientIds, session.ByJwt.NetworkId) + }, server.TxReadCommitted) + return &RemoveNetworkClientsBatchResult{}, nil +} + +// outer sanity bound on a single request's payload, not a realistic +// operating limit (the known real-world case is ~400k-1M): guards against an +// unbounded/malformed request forcing an arbitrarily large synchronous +// json.Marshal + single-row task args_json write in RemoveNetworkClients, and +// against an unbounded in-memory slice being held for the life of a task. +// At 1M ids (~38 bytes per id as JSON ≈ 38MB in pending_task.args_json) this +// is already well into the large-payload range; the previous 2M cap was +// lowered because ~75MB per row was too costly for the task queue. +const MaxRemoveNetworkClientsCount = 1000000 + +type RemoveNetworkClientsArgs struct { + ClientIds []server.Id `json:"client_ids"` +} + +type RemoveNetworkClientsResult struct { + // true if the request was handed off to the background task instead of + // being applied synchronously; deactivation is not yet guaranteed + // complete when this is true + Scheduled bool `json:"scheduled,omitempty"` + // set whenever Scheduled is true: the start of the hourly bucket this + // request was reserved against. Equal to (or a few moments before) the + // current time when there was room in the current hour; a future time + // when the deployment-wide hourly budget was full and this request was + // queued for a later hour instead of being rejected. + ScheduledFor *time.Time `json:"scheduled_for,omitempty"` + // true if a background bulk-delete run for this network was already in + // progress, so this request's ids were NOT scheduled; the caller should + // wait for the in-progress run to finish and retry + AlreadyInProgress bool `json:"already_in_progress,omitempty"` + // true if the deployment-wide cap on concurrent background bulk-delete + // runs (MaxConcurrentBulkClientRemovalRuns) was already reached, so this + // request's ids were NOT scheduled; the caller should retry later + TooManyConcurrentRuns bool `json:"too_many_concurrent_runs,omitempty"` +} + +// MaxConcurrentBulkClientRemovalRuns caps how many networks may have a +// background bulk-delete run (RemoveNetworkClientsTask) in flight or queued +// at the same time, across the whole deployment. Each network's own run is +// already serialized by runNetworkClientsTaskKey; this adds a deployment-wide +// ceiling so many networks triggering large bulk-deletes at once can't all +// load the task queue/database simultaneously. This is a starting value, not +// a measured hard limit -- tune it against real background-task headroom. +const MaxConcurrentBulkClientRemovalRuns = 10 + +// runNetworkClientsTaskKey is the run_once key scoping "one background +// bulk-delete run per network at a time". It's shared between the initial +// schedule in RemoveNetworkClients and the continuation reschedule in +// RemoveNetworkClientsTaskPost, so the key stays held for the full duration +// of a (possibly multi-invocation) run, not just its first invocation. +func runNetworkClientsTaskKey(networkId server.Id) *task.RunOnceOption { + return task.RunOnce("model.RemoveNetworkClientsTask", networkId) +} + +// RemoveNetworkClients deactivates the given clients, scoped to the caller's +// network. A request is applied synchronously, in one transaction, only if +// it's small (at most RemoveNetworkClientsBatchCount) AND the deployment-wide +// hourly budget has room for it right now. Otherwise it's handed off to +// RemoveNetworkClientsTask, a background task: either because it's large +// enough to need batching regardless, or because the current hour's budget +// is full and it has to wait for a later hour's -- see +// ReserveBulkClientRemovalSlot. Either way, a single call can clear a +// network with hundreds of thousands (or millions) of offline clients +// without holding a long-running transaction on the request path or forcing +// the caller to chunk the request themselves. +func RemoveNetworkClients( + removeClients *RemoveNetworkClientsArgs, + session *session.ClientSession, +) (*RemoveNetworkClientsResult, error) { + if len(removeClients.ClientIds) == 0 { + return &RemoveNetworkClientsResult{}, nil + } + + // deduplicate client ids in the input — duplicates in the same request + // would otherwise cause repeated updates on the same row in one batch + clientIds := make([]server.Id, 0, len(removeClients.ClientIds)) + seen := make(map[server.Id]struct{}, len(removeClients.ClientIds)) + for _, id := range removeClients.ClientIds { + if _, ok := seen[id]; !ok { + seen[id] = struct{}{} + clientIds = append(clientIds, id) + } + } + + if MaxRemoveNetworkClientsCount < len(clientIds) { + return nil, fmt.Errorf("Too many client ids (max %d).", MaxRemoveNetworkClientsCount) + } + + // reserve this request's slot in the earliest hourly bucket with room -- + // the current hour, or a later one if the current hour's budget is + // already spent. The bucket has to be known before the request can be + // admitted (either executed now or scheduled for later), so this always + // happens before that admission step, not after -- unlike a simple + // charge, a reservation that turns out to be unusable (concurrency cap + // or an existing in-progress run, below) has to be explicitly released + // rather than just not made. + reservationId, bucketStart, err := ReserveBulkClientRemovalSlot(session.Ctx, session.ByJwt.NetworkId, len(clientIds)) + if err != nil { + return nil, err + } + + currentBucket := bulkClientRemovalBucketStart(server.NowUtc()) + if bucketStart.Equal(currentBucket) && len(clientIds) <= RemoveNetworkClientsBatchCount { + // small enough, and the current hour has room: run synchronously, + // same as always -- this path predates the reservation system and + // was never gated by the per-network run_once key or the + // concurrency cap below either. + _, err := RemoveNetworkClientsBatch(&RemoveNetworkClientsBatchArgs{ + ClientIds: clientIds, + }, session) + if err != nil { + CancelBulkClientRemovalReservation(session.Ctx, reservationId) + return nil, err + } + return &RemoveNetworkClientsResult{}, nil + } + + // everything else goes through the background task -- either the + // request is large enough to need batching regardless of timing, or its + // reserved slot landed in a future hour and there's no way to make an + // HTTP caller wait that long synchronously. + // + // the deployment-wide concurrency cap is checked here, not before the + // reservation above -- it only needs to gate work that will actually + // touch the task system, not a small request that fit the current hour + // and ran synchronously above. CountAvailableByFunctionName only counts + // tasks whose run_at has already passed, so requests merely queued for + // a future hour don't themselves count against this cap; only + // actually-running (or runnable-now) work does. This count includes + // this network's own run if one is already in flight -- in that rare + // case, a network that is itself occupying one of the counted slots + // gets this generic "too many concurrent runs" signal here instead of + // the more specific AlreadyInProgress a few lines down. Both are the + // same instruction to the caller (retry later), so this is a deliberate + // simplification rather than tracking per-network exclusions. + if concurrentRuns := task.CountAvailableByFunctionName(session.Ctx, RemoveNetworkClientsTask); MaxConcurrentBulkClientRemovalRuns <= concurrentRuns { + CancelBulkClientRemovalReservation(session.Ctx, reservationId) + return &RemoveNetworkClientsResult{TooManyConcurrentRuns: true}, nil + } + + // a request queued for a future hour gets its actual run time spread + // randomly across that hour, rather than firing at exactly the hour's + // top. CountAvailableByFunctionName above only ever sees currently- + // runnable work, so it can't tell how large the backlog already queued + // for a given future hour is; without this jitter, every request a + // sustained period of heavy bulk-delete traffic pushed into the same + // future bucket would become eligible at the identical instant, + // defeating the concurrency cap's purpose the moment that hour arrives. + // A request landing in the current hour still runs as close to + // immediately as before -- only deferred requests get spread out. + runAt := bucketStart + if currentBucket.Before(bucketStart) { + runAt = bucketStart.Add(time.Duration(mathrand.Int63n(int64(BulkClientRemovalBucketDuration)))) + } + + // one background bulk-delete run per network at a time. + // ScheduleTaskInTxIfAbsent (unlike plain ScheduleTask+RunOnce) makes the + // "only if not already pending" check atomic with the insert -- a single + // `INSERT ... ON CONFLICT (run_once_key) DO NOTHING`, reporting whether + // the row was actually inserted. This closes the race a naive + // check-then-act would have: with check-then-act, two near-simultaneous + // requests for the same network could both pass the check, and the + // second's schedule call would then hit ON CONFLICT DO UPDATE -- which + // merges only timing/priority into the existing row, NOT args_json -- + // silently dropping the second call's client_ids while still reporting + // success. The atomic insert-or-detect-conflict here means a duplicate + // is always rejected outright, never silently swallowed. + scheduled, _ := task.ScheduleTaskIfAbsent( + RemoveNetworkClientsTask, + &RemoveNetworkClientsTaskArgs{ + ClientIds: clientIds, + }, + session, + runNetworkClientsTaskKey(session.ByJwt.NetworkId), + // bulk cleanup must never compete with revenue/critical-path tasks + // (payouts, contract close) under multi-tenant load + task.Priority(task.TaskPrioritySlowest), + task.MaxTime(30*time.Minute), + task.RunAt(runAt), + ) + if !scheduled { + // the reservation was made under this network's name, but it turns + // out there's already a run in progress -- release the slot so it + // doesn't sit charged against nothing. + CancelBulkClientRemovalReservation(session.Ctx, reservationId) + return &RemoveNetworkClientsResult{AlreadyInProgress: true}, nil + } + + scheduledFor := runAt + return &RemoveNetworkClientsResult{Scheduled: true, ScheduledFor: &scheduledFor}, nil +} + +// number of batches processed per RemoveNetworkClientsTask invocation before +// it self-reschedules the remainder (see RemoveNetworkClientsTaskPost), +// instead of looping over the entire list in one invocation. This bounds +// three things: (1) a single invocation's duration stays well inside +// MaxTime regardless of total list size, (2) if an invocation is retried +// after a crash/timeout, it only re-attempts this bounded chunk, not the +// entire original list, so a run makes durable checkpointed progress instead +// of restarting from the beginning every retry, and (3) a persistently +// failing chunk only blocks its own continuation, not the network's ability +// to ever complete a bulk delete. +const RemoveNetworkClientsTaskBatchLimit = 20 + +type RemoveNetworkClientsTaskArgs struct { + ClientIds []server.Id `json:"client_ids"` +} + +type RemoveNetworkClientsTaskResult struct { + // ids not yet processed by this invocation; if non-empty, + // RemoveNetworkClientsTaskPost reschedules them as a new task under the + // same run_once key + RemainingClientIds []server.Id `json:"remaining_client_ids,omitempty"` +} + +// RemoveNetworkClientsTask is the background counterpart to +// RemoveNetworkClients for requests larger than RemoveNetworkClientsBatchCount. +// It runs on the taskworker service (not the live request pool) and processes +// up to RemoveNetworkClientsTaskBatchLimit batches of +// RemoveNetworkClientsBatchCount ids, each its own short server.MaintenanceTx +// transaction at TxReadCommitted, mirroring the batching and isolation level +// RemoveDisconnectedNetworkClients already uses for this table. The update is +// idempotent, so a retried or re-run invocation is safe. +func RemoveNetworkClientsTask( + removeClients *RemoveNetworkClientsTaskArgs, + session *session.ClientSession, +) (*RemoveNetworkClientsTaskResult, error) { + if session.ByJwt == nil { + // unreachable in the normal flow (always scheduled from an + // authenticated handler with ByJwt set), but this task can outlive + // the request that scheduled it -- fail loudly rather than panic on + // a nil dereference if a pending_task row is ever reconstructed + // with an empty/corrupted client_by_jwt_json + return nil, fmt.Errorf("Missing network for bulk client removal.") + } + + clientIds := removeClients.ClientIds + batches := 0 + for 0 < len(clientIds) && batches < RemoveNetworkClientsTaskBatchLimit { + batchCount := RemoveNetworkClientsBatchCount + if len(clientIds) < batchCount { + batchCount = len(clientIds) + } + batch := clientIds[:batchCount] + clientIds = clientIds[batchCount:] + + server.MaintenanceTx(session.Ctx, func(tx server.PgTx) { + removeNetworkClientsBatchExec(session.Ctx, tx, batch, session.ByJwt.NetworkId) + }, server.TxReadCommitted) + batches += 1 + } + + return &RemoveNetworkClientsTaskResult{ + RemainingClientIds: clientIds, + }, nil +} + +// RemoveNetworkClientsTaskPost reschedules any remainder left by +// RemoveNetworkClientsTask under the SAME run_once key as the original +// request, so the "one run per network" guarantee holds for the full +// duration of a multi-invocation run and not just its first invocation -- +// otherwise the key would free up as soon as the first chunk finished, even +// though most of the list might still be unprocessed, and a second +// concurrent request would incorrectly be allowed through. +// +// This reschedules with plain ScheduleTaskInTx (ON CONFLICT DO UPDATE), not +// ScheduleTaskInTxIfAbsent (ON CONFLICT DO NOTHING): the finishing task's own +// pending_task row is deleted as part of this same evaluation transaction, +// so the run_once key is guaranteed free at the point this INSERT runs -- +// there's no possible conflict with a legitimate concurrent request here, +// only with the row this exact chain just vacated. +func RemoveNetworkClientsTaskPost( + removeClients *RemoveNetworkClientsTaskArgs, + result *RemoveNetworkClientsTaskResult, + session *session.ClientSession, + tx server.PgTx, +) error { + if session.ByJwt == nil { + // unreachable in the normal flow, but same rationale as the guard in + // RemoveNetworkClientsTask: if a corrupted/empty client_by_jwt_json + // row reaches the post phase, fail loudly rather than nil-deref. + return fmt.Errorf("Missing network for bulk client removal post.") + } + if 0 < len(result.RemainingClientIds) { + task.ScheduleTaskInTx( + tx, + RemoveNetworkClientsTask, + &RemoveNetworkClientsTaskArgs{ + ClientIds: result.RemainingClientIds, + }, + session, + runNetworkClientsTaskKey(session.ByJwt.NetworkId), + task.Priority(task.TaskPrioritySlowest), + task.MaxTime(30*time.Minute), + ) + } + return 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 d730c921..27dfca9b 100644 --- a/model/network_client_model_test.go +++ b/model/network_client_model_test.go @@ -8,11 +8,13 @@ import ( "testing" "time" + "github.com/go-playground/assert/v2" "github.com/urnetwork/connect" "github.com/urnetwork/server" "github.com/urnetwork/server/jwt" "github.com/urnetwork/server/session" + "github.com/urnetwork/server/task" ) func TestNetworkClientHandlerLifecycle(t *testing.T) { @@ -341,6 +343,1018 @@ func TestSetProvide(t *testing.T) { }) } +// RemoveNetworkClients must accept a uuid[]-bound array of ids that don't +// exist yet (nothing to match), and must not error or panic on the +// []server.Id -> uuid[] cast. +func TestRemoveNetworkClientsUUIDArrayBinding(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkId := server.NewId() + + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{ + NetworkId: networkId, + }, + } + + // random ids with no matching rows + clientIds := []server.Id{server.NewId(), server.NewId(), server.NewId()} + + args := &RemoveNetworkClientsArgs{ + ClientIds: clientIds, + } + + result, err := RemoveNetworkClients(args, sess) + assert.Equal(t, err, nil) + assert.NotEqual(t, result, nil) + }) +} + +// The empty-ids case must be a no-op that returns cleanly without opening a +// transaction (an empty `ANY($1)` array matches nothing, but there's no +// reason to pay for the round trip). +func TestRemoveNetworkClientsEmptyIdsNoop(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkId := server.NewId() + deviceId := server.NewId() + clientId := server.NewId() + + Testing_CreateDevice(ctx, networkId, deviceId, clientId, "test", "test") + + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{ + NetworkId: networkId, + }, + } + + result, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: []server.Id{}, + }, sess) + assert.Equal(t, err, nil) + assert.NotEqual(t, result, nil) + + // the existing client must be untouched + assert.NotEqual(t, GetNetworkClient(ctx, clientId), nil) + }) +} + +// The bulk path must actually deactivate the targeted clients (mirroring the +// single-client RemoveNetworkClient behavior) while leaving other clients on +// the same network alone. +func TestRemoveNetworkClientsDeactivatesTargetedClients(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkId := server.NewId() + + deviceIdA := server.NewId() + clientIdA := server.NewId() + Testing_CreateDevice(ctx, networkId, deviceIdA, clientIdA, "test-a", "test") + + deviceIdB := server.NewId() + clientIdB := server.NewId() + Testing_CreateDevice(ctx, networkId, deviceIdB, clientIdB, "test-b", "test") + + deviceIdC := server.NewId() + clientIdC := server.NewId() + Testing_CreateDevice(ctx, networkId, deviceIdC, clientIdC, "test-c", "test") + + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{ + NetworkId: networkId, + }, + } + + beforeCall := server.NowUtc() + _, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: []server.Id{clientIdA, clientIdB}, + }, sess) + assert.Equal(t, err, nil) + + // targeted clients are deactivated + assert.Equal(t, GetNetworkClient(ctx, clientIdA), nil) + assert.Equal(t, GetNetworkClient(ctx, clientIdB), nil) + + // the untargeted client on the same network is untouched + assert.NotEqual(t, GetNetworkClient(ctx, clientIdC), nil) + + // deactivate_time must be stamped, matching the single-client path, + // so the reap job (`COALESCE(deactivate_time, create_time)`) applies + // the grace period from the actual deactivation instead of falling + // back to create_time + var deactivateTime *time.Time + server.Db(ctx, func(conn server.PgConn) { + result, err := conn.Query( + ctx, + `SELECT deactivate_time FROM network_client WHERE client_id = $1`, + clientIdA, + ) + server.WithPgResult(result, err, func() { + if result.Next() { + server.Raise(result.Scan(&deactivateTime)) + } + }) + }) + if deactivateTime == nil { + t.Fatal("deactivate_time was not set") + } + if deactivateTime.Before(beforeCall) { + t.Fatal("deactivate_time predates the call") + } + }) +} + +// A request at or under the batch size must be applied synchronously: the +// caller gets a definite "done" (Scheduled == false) and the clients are +// already deactivated by the time the call returns. +func TestRemoveNetworkClientsSmallRequestIsSynchronous(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkId := server.NewId() + + deviceId := server.NewId() + clientId := server.NewId() + Testing_CreateDevice(ctx, networkId, deviceId, clientId, "test", "test") + + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{ + NetworkId: networkId, + }, + } + + result, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: []server.Id{clientId}, + }, sess) + assert.Equal(t, err, nil) + assert.Equal(t, result.Scheduled, false) + + // already deactivated, not just enqueued + assert.Equal(t, GetNetworkClient(ctx, clientId), nil) + }) +} + +// A request over the batch size must be handed off to the background task +// instead of run inline: the caller gets Scheduled == true and the clients +// are NOT yet deactivated when the call returns (only the task is enqueued; +// RemoveNetworkClientsTask is what actually applies it, tested separately +// below). +func TestRemoveNetworkClientsLargeRequestIsScheduled(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkId := server.NewId() + + deviceId := server.NewId() + clientId := server.NewId() + Testing_CreateDevice(ctx, networkId, deviceId, clientId, "test", "test") + + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{ + NetworkId: networkId, + }, + } + + clientIds := make([]server.Id, RemoveNetworkClientsBatchCount+1) + clientIds[0] = clientId + for i := 1; i < len(clientIds); i++ { + clientIds[i] = server.NewId() + } + + result, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: clientIds, + }, sess) + assert.Equal(t, err, nil) + assert.Equal(t, result.Scheduled, true) + assert.Equal(t, result.AlreadyInProgress, false) + + // not yet applied: this call only enqueued the background task + assert.NotEqual(t, GetNetworkClient(ctx, clientId), nil) + }) +} + +// A second large request for the same network, made while the first is +// still pending, must be rejected outright (AlreadyInProgress == true, not +// Scheduled) rather than silently merged into the first task -- merging +// would drop the second call's client_ids, since ScheduleTask's run_once +// conflict path only updates timing/priority, not args. +func TestRemoveNetworkClientsRejectsDuplicateInProgress(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkId := server.NewId() + + firstDeviceId := server.NewId() + firstClientId := server.NewId() + Testing_CreateDevice(ctx, networkId, firstDeviceId, firstClientId, "first", "test") + + secondDeviceId := server.NewId() + secondClientId := server.NewId() + Testing_CreateDevice(ctx, networkId, secondDeviceId, secondClientId, "second", "test") + + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{ + NetworkId: networkId, + }, + } + + largeClientIds := func(first server.Id) []server.Id { + clientIds := make([]server.Id, RemoveNetworkClientsBatchCount+1) + clientIds[0] = first + for i := 1; i < len(clientIds); i++ { + clientIds[i] = server.NewId() + } + return clientIds + } + + firstResult, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: largeClientIds(firstClientId), + }, sess) + assert.Equal(t, err, nil) + assert.Equal(t, firstResult.Scheduled, true) + assert.Equal(t, firstResult.AlreadyInProgress, false) + + // a second large request for the same network, while the first is + // still pending (unclaimed), must be rejected, not scheduled + secondResult, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: largeClientIds(secondClientId), + }, sess) + assert.Equal(t, err, nil) + assert.Equal(t, secondResult.Scheduled, false) + assert.Equal(t, secondResult.AlreadyInProgress, true) + + // neither client is deactivated yet: the first task hasn't run + // (only enqueued), and the second call was never scheduled at all + assert.NotEqual(t, GetNetworkClient(ctx, firstClientId), nil) + assert.NotEqual(t, GetNetworkClient(ctx, secondClientId), nil) + + // the rejected (AlreadyInProgress) second call must not have burned + // any of the current bucket's budget -- only the admitted first + // call's count should have been charged. Probe by reserving exactly + // the remainder of the current bucket's ceiling: this only lands in + // the CURRENT bucket if nothing beyond the first call's charge was + // recorded there. + _, bucketStart, err := ReserveBulkClientRemovalSlot( + ctx, + server.NewId(), + MaxBulkClientRemovalsPerBucket-len(largeClientIds(firstClientId)), + ) + assert.Equal(t, err, nil) + assert.Equal(t, bucketStart, bulkClientRemovalBucketStart(server.NowUtc())) + }) +} + +// A large request for a DIFFERENT network must not be blocked by an +// in-progress run on another network: the run_once key is scoped per network. +func TestRemoveNetworkClientsInProgressIsScopedPerNetwork(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkIdA := server.NewId() + networkIdB := server.NewId() + + deviceIdA := server.NewId() + clientIdA := server.NewId() + Testing_CreateDevice(ctx, networkIdA, deviceIdA, clientIdA, "a", "test") + + deviceIdB := server.NewId() + clientIdB := server.NewId() + Testing_CreateDevice(ctx, networkIdB, deviceIdB, clientIdB, "b", "test") + + largeClientIds := func(first server.Id) []server.Id { + clientIds := make([]server.Id, RemoveNetworkClientsBatchCount+1) + clientIds[0] = first + for i := 1; i < len(clientIds); i++ { + clientIds[i] = server.NewId() + } + return clientIds + } + + sessA := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{NetworkId: networkIdA}, + } + sessB := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{NetworkId: networkIdB}, + } + + resultA, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: largeClientIds(clientIdA), + }, sessA) + assert.Equal(t, err, nil) + assert.Equal(t, resultA.Scheduled, true) + + resultB, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: largeClientIds(clientIdB), + }, sessB) + assert.Equal(t, err, nil) + assert.Equal(t, resultB.Scheduled, true) + assert.Equal(t, resultB.AlreadyInProgress, false) + }) +} + +// A request over the outer sanity cap must be rejected outright, before any +// scheduling or database work -- this is a payload-size guard, not a +// realistic operating limit (known real-world usage is far below it). +func TestRemoveNetworkClientsRejectsOversizedRequest(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkId := server.NewId() + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{ + NetworkId: networkId, + }, + } + + // distinct ids are required here: RemoveNetworkClients deduplicates + // before checking the length cap, so identical (e.g. zero-value) + // ids would collapse to a single entry and never exercise the cap. + clientIds := make([]server.Id, MaxRemoveNetworkClientsCount+1) + for i := range clientIds { + clientIds[i] = server.NewId() + } + + _, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: clientIds, + }, sess) + assert.NotEqual(t, err, nil) + }) +} + +// End-to-end through a real task.TaskWorker (not calling RemoveNetworkClientsTask +// directly): a request large enough to require more than one invocation +// (RemoveNetworkClientsTaskBatchLimit batches) must fully deactivate every +// client across the self-rescheduled chain; a duplicate large request for the +// same network must be rejected while any part of the chain is still +// in-flight; and once the whole chain completes, the run_once key must be +// free again for a new request. This is the scenario the AlreadyInProgress +// flag depends on, and it isn't exercised by calling RemoveNetworkClientsTask +// directly (that only tests one invocation's batching, not the pending_task +// lifecycle across a real claim/run/post/reschedule cycle). +func TestRemoveNetworkClientsTaskLifecycleThroughRealWorker(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkId := server.NewId() + + // exceed one invocation's batch limit so the task must self-reschedule + // at least once + targetCount := RemoveNetworkClientsTaskBatchLimit*RemoveNetworkClientsBatchCount + 5000 + var clientIds []server.Id + server.Db(ctx, func(conn server.PgConn) { + result, err := conn.Query( + ctx, + ` + INSERT INTO network_client (client_id, network_id, active, create_time, auth_time) + SELECT gen_random_uuid(), $1, true, now(), now() + FROM generate_series(1, $2) g + RETURNING client_id + `, + networkId, + targetCount, + ) + server.WithPgResult(result, err, func() { + for result.Next() { + var clientId server.Id + server.Raise(result.Scan(&clientId)) + clientIds = append(clientIds, clientId) + } + }) + }) + + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{ + NetworkId: networkId, + }, + } + + result, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: clientIds, + }, sess) + assert.Equal(t, err, nil) + assert.Equal(t, result.Scheduled, true) + + // while the run is in flight, a second large request for the same + // network must be rejected, not merged or double-scheduled + dupResult, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: clientIds[:RemoveNetworkClientsBatchCount+1], + }, sess) + assert.Equal(t, err, nil) + assert.Equal(t, dupResult.Scheduled, false) + assert.Equal(t, dupResult.AlreadyInProgress, true) + + taskWorker := task.NewTaskWorkerWithDefaults(ctx) + taskWorker.AddTargets(task.NewTaskTargetWithPost(RemoveNetworkClientsTask, RemoveNetworkClientsTaskPost)) + + // drive the worker until no pending work remains for this run; bounded + // so a stuck run fails the test instead of hanging it + for i := 0; i < 50; i++ { + finishedTaskIds, rescheduledTaskIds, postRescheduledTaskIds, err := taskWorker.EvalTasks(10) + assert.Equal(t, err, nil) + if len(finishedTaskIds)+len(rescheduledTaskIds)+len(postRescheduledTaskIds) == 0 && + len(task.ListPendingTasks(ctx)) == 0 { + break + } + } + if 0 < len(task.ListPendingTasks(ctx)) { + t.Fatal("run did not complete within the bounded number of eval passes") + } + + var remainingActive int + server.Db(ctx, func(conn server.PgConn) { + result, err := conn.Query( + ctx, + `SELECT COUNT(*) FROM network_client WHERE network_id = $1 AND active = true`, + networkId, + ) + server.WithPgResult(result, err, func() { + if result.Next() { + server.Raise(result.Scan(&remainingActive)) + } + }) + }) + assert.Equal(t, remainingActive, 0) + + // deactivate_time must be stamped across the whole async chain, not + // just the first invocation -- this is the only test exercising + // multiple RemoveNetworkClientsTask invocations, so it's the one + // that would catch a continuation batch losing the stamp + var missingDeactivateTime int + server.Db(ctx, func(conn server.PgConn) { + result, err := conn.Query( + ctx, + `SELECT COUNT(*) FROM network_client WHERE network_id = $1 AND active = false AND deactivate_time IS NULL`, + networkId, + ) + server.WithPgResult(result, err, func() { + if result.Next() { + server.Raise(result.Scan(&missingDeactivateTime)) + } + }) + }) + assert.Equal(t, missingDeactivateTime, 0) + + // the run_once key must be free again now that the full chain + // finished, so a new large request for this network succeeds + afterResult, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: clientIds[:RemoveNetworkClientsBatchCount+1], + }, sess) + assert.Equal(t, err, nil) + assert.Equal(t, afterResult.Scheduled, true) + }) +} + +// RemoveNetworkClientsTask (the background counterpart RemoveNetworkClients +// hands large requests off to) must deactivate every targeted client across +// multiple internal batches, and must not touch a client on a different +// network that happens to fall in the same id range. Called directly here +// (not through the task queue) to test its batching logic deterministically +// without a running taskworker. +func TestRemoveNetworkClientsTaskSpansMultipleBatches(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkId := server.NewId() + otherNetworkId := server.NewId() + + // seed more than 2x the batch size for the target network, plus one + // client on a different network, via a bulk insert (matches the + // generate_series seeding pattern used elsewhere in this package) + targetCount := 2*RemoveNetworkClientsBatchCount + 5000 + var clientIds []server.Id + server.Db(ctx, func(conn server.PgConn) { + result, err := conn.Query( + ctx, + ` + INSERT INTO network_client (client_id, network_id, active, create_time, auth_time) + SELECT gen_random_uuid(), $1, true, now(), now() + FROM generate_series(1, $2) g + RETURNING client_id + `, + networkId, + targetCount, + ) + server.WithPgResult(result, err, func() { + for result.Next() { + var clientId server.Id + server.Raise(result.Scan(&clientId)) + clientIds = append(clientIds, clientId) + } + }) + }) + + otherDeviceId := server.NewId() + otherClientId := server.NewId() + Testing_CreateDevice(ctx, otherNetworkId, otherDeviceId, otherClientId, "other", "test") + + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{ + NetworkId: networkId, + }, + } + + _, err := RemoveNetworkClientsTask(&RemoveNetworkClientsTaskArgs{ + ClientIds: clientIds, + }, sess) + assert.Equal(t, err, nil) + + var remainingActive int + server.Db(ctx, func(conn server.PgConn) { + result, err := conn.Query( + ctx, + `SELECT COUNT(*) FROM network_client WHERE network_id = $1 AND active = true`, + networkId, + ) + server.WithPgResult(result, err, func() { + if result.Next() { + server.Raise(result.Scan(&remainingActive)) + } + }) + }) + assert.Equal(t, remainingActive, 0) + + // a client on a different network is untouched + assert.NotEqual(t, GetNetworkClient(ctx, otherClientId), nil) + }) +} + +// If a task worker retries RemoveNetworkClientsTask (e.g. after a crash or a +// MaxTime timeout mid-run), re-running it with the same args must be a safe +// no-op the second time: `active = false` on an already-inactive row doesn't +// error, and deactivate_time just advances rather than corrupting state. +func TestRemoveNetworkClientsTaskIsIdempotentOnRetry(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkId := server.NewId() + + deviceId := server.NewId() + clientId := server.NewId() + Testing_CreateDevice(ctx, networkId, deviceId, clientId, "test", "test") + + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{ + NetworkId: networkId, + }, + } + + args := &RemoveNetworkClientsTaskArgs{ + ClientIds: []server.Id{clientId}, + } + + _, err := RemoveNetworkClientsTask(args, sess) + assert.Equal(t, err, nil) + assert.Equal(t, GetNetworkClient(ctx, clientId), nil) + + // simulate a retry of the same task with the same args + _, err = RemoveNetworkClientsTask(args, sess) + assert.Equal(t, err, nil) + assert.Equal(t, GetNetworkClient(ctx, clientId), nil) + }) +} + +// RemoveNetworkClientsTask must fail loudly (not panic on a nil dereference +// inside a MaintenanceTx) if it's ever run with a session that has no +// ByJwt -- unreachable in the normal flow (always scheduled from an +// authenticated handler), but the task can outlive the request that +// scheduled it, so this guards against a corrupted/reconstructed session. +func TestRemoveNetworkClientsTaskRejectsNilByJwt(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: nil, + } + + _, err := RemoveNetworkClientsTask(&RemoveNetworkClientsTaskArgs{ + ClientIds: []server.Id{server.NewId()}, + }, sess) + assert.NotEqual(t, err, nil) + }) +} + +// A request of exactly RemoveNetworkClientsBatchCount ids must still take +// the synchronous path (the boundary is "<=", not "<"). +func TestRemoveNetworkClientsExactlyAtBatchCountIsSynchronous(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkId := server.NewId() + + deviceId := server.NewId() + clientId := server.NewId() + Testing_CreateDevice(ctx, networkId, deviceId, clientId, "test", "test") + + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{ + NetworkId: networkId, + }, + } + + clientIds := make([]server.Id, RemoveNetworkClientsBatchCount) + clientIds[0] = clientId + for i := 1; i < len(clientIds); i++ { + clientIds[i] = server.NewId() + } + + result, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: clientIds, + }, sess) + assert.Equal(t, err, nil) + assert.Equal(t, result.Scheduled, false) + + // already applied synchronously, not enqueued + assert.Equal(t, GetNetworkClient(ctx, clientId), nil) + }) +} + +// A request of exactly MaxRemoveNetworkClientsCount ids must be accepted +// (the boundary is "> max is rejected", not ">="). This only exercises +// scheduling (one INSERT), not actually running the task, since processing +// 2M ids isn't practical inside a unit test. +func TestRemoveNetworkClientsExactlyAtCapIsAccepted(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkId := server.NewId() + + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{ + NetworkId: networkId, + }, + } + + // distinct ids are required here too: identical ids would dedup + // below the cap and this test would no longer be "exactly at cap". + clientIds := make([]server.Id, MaxRemoveNetworkClientsCount) + for i := range clientIds { + clientIds[i] = server.NewId() + } + + result, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: clientIds, + }, sess) + assert.Equal(t, err, nil) + assert.Equal(t, result.Scheduled, true) + }) +} + +// The deployment-wide concurrency cap must reject a new network's async +// bulk-delete once MaxConcurrentBulkClientRemovalRuns other networks already +// have one pending. Occupying the cap via direct task scheduling (rather +// than MaxConcurrentBulkClientRemovalRuns full RemoveNetworkClients calls) +// keeps the test to the mechanism under test -- the pending_task count -- and +// avoids needing real per-network client rows for slots that are never +// actually drained here. +func TestRemoveNetworkClientsRejectsWhenConcurrencyCapReached(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + + for i := 0; i < MaxConcurrentBulkClientRemovalRuns; i++ { + occupyingNetworkId := server.NewId() + occupyingSess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{NetworkId: occupyingNetworkId}, + } + scheduled, _ := task.ScheduleTaskIfAbsent( + RemoveNetworkClientsTask, + &RemoveNetworkClientsTaskArgs{ClientIds: []server.Id{server.NewId()}}, + occupyingSess, + runNetworkClientsTaskKey(occupyingNetworkId), + ) + assert.Equal(t, scheduled, true) + } + + networkId := server.NewId() + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{NetworkId: networkId}, + } + + clientIds := make([]server.Id, RemoveNetworkClientsBatchCount+1) + for i := range clientIds { + clientIds[i] = server.NewId() + } + + result, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: clientIds, + }, sess) + assert.Equal(t, err, nil) + assert.Equal(t, result.Scheduled, false) + assert.Equal(t, result.TooManyConcurrentRuns, true) + + // a TooManyConcurrentRuns rejection must not have burned any budget + // for the rejected request's id count -- the reservation made just + // before the concurrency check must have been released, so the + // current bucket's whole ceiling is still available. + _, bucketStart, err := ReserveBulkClientRemovalSlot(ctx, server.NewId(), MaxBulkClientRemovalsPerBucket) + assert.Equal(t, err, nil) + assert.Equal(t, bucketStart, bulkClientRemovalBucketStart(server.NowUtc())) + }) +} + +// The concurrency cap must only gate work that actually goes through the +// background task system -- a small request that fits the current hour's +// budget runs synchronously and was never subject to it, even when the cap +// is fully occupied by unrelated networks' background runs. +func TestRemoveNetworkClientsSyncPathBypassesConcurrencyCap(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + + for i := 0; i < MaxConcurrentBulkClientRemovalRuns; i++ { + occupyingNetworkId := server.NewId() + occupyingSess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{NetworkId: occupyingNetworkId}, + } + scheduled, _ := task.ScheduleTaskIfAbsent( + RemoveNetworkClientsTask, + &RemoveNetworkClientsTaskArgs{ClientIds: []server.Id{server.NewId()}}, + occupyingSess, + runNetworkClientsTaskKey(occupyingNetworkId), + ) + assert.Equal(t, scheduled, true) + } + + networkId := server.NewId() + deviceId := server.NewId() + clientId := server.NewId() + Testing_CreateDevice(ctx, networkId, deviceId, clientId, "test", "test") + + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{NetworkId: networkId}, + } + + result, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: []server.Id{clientId}, + }, sess) + assert.Equal(t, err, nil) + assert.Equal(t, result.TooManyConcurrentRuns, false) + assert.Equal(t, GetNetworkClient(ctx, clientId), nil) + }) +} + +// A request large enough to need the background task, whose reserved slot +// is the current hour (room available now, not deferred), must report +// ScheduledFor as the current bucket -- not just the deferred path. +func TestRemoveNetworkClientsSetsScheduledForOnImmediateAsyncRun(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkId := server.NewId() + currentBucket := bulkClientRemovalBucketStart(server.NowUtc()) + + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{NetworkId: networkId}, + } + + clientIds := make([]server.Id, RemoveNetworkClientsBatchCount+1) + for i := range clientIds { + clientIds[i] = server.NewId() + } + + result, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: clientIds, + }, sess) + assert.Equal(t, err, nil) + assert.Equal(t, result.Scheduled, true) + assert.NotEqual(t, result.ScheduledFor, nil) + assert.Equal(t, *result.ScheduledFor, currentBucket) + }) +} + +// The per-network single-flight lock must stay held for the whole +// queued-and-waiting duration of a deferred request, not just while it's +// actually running -- a second request for the same network arriving before +// the first's deferred bucket time is reached must be AlreadyInProgress. +// Unlike TestRemoveNetworkClientsCancelsReservationOnAlreadyInProgress +// (which fakes an in-progress run via a directly-scheduled run_at=now +// task), this drives the actual deferral path end to end: the first +// request's task genuinely has a future run_at, so +// CountAvailableByFunctionName correctly does not count it as "in flight", +// yet the run_once key it holds must still block a second request for the +// same network. +func TestRemoveNetworkClientsLocksNetworkWhileDeferredRequestIsQueued(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + + // fill the current bucket so the next reservation defers to the + // next hour + _, _, err := ReserveBulkClientRemovalSlot(ctx, server.NewId(), MaxBulkClientRemovalsPerBucket) + assert.Equal(t, err, nil) + + networkId := server.NewId() + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{NetworkId: networkId}, + } + + clientIds := make([]server.Id, RemoveNetworkClientsBatchCount+1) + for i := range clientIds { + clientIds[i] = server.NewId() + } + + firstResult, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: clientIds, + }, sess) + assert.Equal(t, err, nil) + assert.Equal(t, firstResult.Scheduled, true) + assert.NotEqual(t, firstResult.ScheduledFor, nil) + assert.NotEqual(t, *firstResult.ScheduledFor, bulkClientRemovalBucketStart(server.NowUtc())) + + // the deferred task isn't runnable yet, so it doesn't count against + // the concurrency cap -- but it must still hold the run_once key + secondResult, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: []server.Id{server.NewId()}, + }, sess) + assert.Equal(t, err, nil) + assert.Equal(t, secondResult.AlreadyInProgress, true) + }) +} + +// Deferred requests must actually be spread across their target hour, not +// merely eligible to be (which the range check in +// TestRemoveNetworkClientsQueuesSmallRequestWhenCurrentHourIsFull would pass +// even if the jitter were removed and every request landed on the exact +// hour boundary). Several distinct networks, all deferred to the same next +// hour, must not all land on the identical instant. +func TestRemoveNetworkClientsSpreadsDeferredRequestsAcrossTheHour(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + + // fill the current bucket so every request below defers to the + // next hour + _, _, err := ReserveBulkClientRemovalSlot(ctx, server.NewId(), MaxBulkClientRemovalsPerBucket) + assert.Equal(t, err, nil) + + scheduledFors := map[time.Time]bool{} + for i := 0; i < 5; i++ { + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{NetworkId: server.NewId()}, + } + clientIds := make([]server.Id, RemoveNetworkClientsBatchCount+1) + for j := range clientIds { + clientIds[j] = server.NewId() + } + + result, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: clientIds, + }, sess) + assert.Equal(t, err, nil) + assert.Equal(t, result.Scheduled, true) + scheduledFors[*result.ScheduledFor] = true + } + + // with 5 independently jittered samples across a full hour, landing + // on the exact same instant every time is not a realistic outcome + // of a correctly-jittered scheduler -- this fails if jitter is + // removed (every sample would land on the same hour boundary) + assert.Equal(t, 1 < len(scheduledFors), true) + }) +} + +// A quota rejection on the async path must cancel the task it just +// scheduled (the bucket has to be known before scheduling, since it becomes +// the task's RunAt, so reservation always happens before the run_once +// admission check, not after). Without the cancellation, the reservation +// would sit charged against a bucket for a request that never actually ran, +// eating into the shared budget for nothing. +func TestRemoveNetworkClientsCancelsReservationOnAlreadyInProgress(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkId := server.NewId() + currentBucket := bulkClientRemovalBucketStart(server.NowUtc()) + + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{NetworkId: networkId}, + } + + // occupy this network's run_once key directly, simulating a run + // already in progress + scheduled, _ := task.ScheduleTaskIfAbsent( + RemoveNetworkClientsTask, + &RemoveNetworkClientsTaskArgs{ClientIds: []server.Id{server.NewId()}}, + sess, + runNetworkClientsTaskKey(networkId), + ) + assert.Equal(t, scheduled, true) + + // large enough to take the async path + clientIds := make([]server.Id, RemoveNetworkClientsBatchCount+1) + for i := range clientIds { + clientIds[i] = server.NewId() + } + + result, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: clientIds, + }, sess) + assert.Equal(t, err, nil) + assert.Equal(t, result.AlreadyInProgress, true) + + // the reservation made just before the run_once check must have + // been released -- the current bucket's full ceiling must still be + // available to a fresh request + _, bucketStart, err := ReserveBulkClientRemovalSlot(ctx, server.NewId(), MaxBulkClientRemovalsPerBucket) + assert.Equal(t, err, nil) + assert.Equal(t, bucketStart, currentBucket) + }) +} + +// RemoveNetworkClientsTaskPost, tested directly and in isolation (not +// through a full task.TaskWorker run, which the slow real-worker lifecycle +// test above already covers end-to-end): a non-empty RemainingClientIds must +// result in exactly one new pending task, scheduled under the same run_once +// key as the original request; an empty RemainingClientIds must schedule +// nothing at all. +func TestRemoveNetworkClientsTaskPostReschedulesRemainder(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + networkId := server.NewId() + + sess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{ + NetworkId: networkId, + }, + } + + remainingClientIds := []server.Id{server.NewId(), server.NewId()} + + server.Tx(ctx, func(tx server.PgTx) { + err := RemoveNetworkClientsTaskPost( + &RemoveNetworkClientsTaskArgs{ClientIds: remainingClientIds}, + &RemoveNetworkClientsTaskResult{RemainingClientIds: remainingClientIds}, + sess, + tx, + ) + assert.Equal(t, err, nil) + }) + + // the reschedule must be blocked by the run_once key (same key the + // original request would have used), proving the continuation was + // scheduled under it + scheduledAgain, _ := task.ScheduleTaskIfAbsent( + RemoveNetworkClientsTask, + &RemoveNetworkClientsTaskArgs{ClientIds: []server.Id{server.NewId()}}, + sess, + runNetworkClientsTaskKey(networkId), + ) + assert.Equal(t, scheduledAgain, false) + + // a DIFFERENT network's key must be unaffected + otherNetworkId := server.NewId() + otherSess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{NetworkId: otherNetworkId}, + } + scheduledOther, _ := task.ScheduleTaskIfAbsent( + RemoveNetworkClientsTask, + &RemoveNetworkClientsTaskArgs{ClientIds: []server.Id{server.NewId()}}, + otherSess, + runNetworkClientsTaskKey(otherNetworkId), + ) + assert.Equal(t, scheduledOther, true) + }) +} + +// A caller must not be able to deactivate another network's clients by +// passing their ids in the request body: `network_id = $2` in the query +// must be enforced server-side from the session, not trusted from input. +func TestRemoveNetworkClientsEnforcesNetworkScoping(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + + victimNetworkId := server.NewId() + victimDeviceId := server.NewId() + victimClientId := server.NewId() + Testing_CreateDevice(ctx, victimNetworkId, victimDeviceId, victimClientId, "victim", "test") + + attackerNetworkId := server.NewId() + attackerSess := &session.ClientSession{ + Ctx: ctx, + ByJwt: &jwt.ByJwt{ + NetworkId: attackerNetworkId, + }, + } + + _, err := RemoveNetworkClients(&RemoveNetworkClientsArgs{ + ClientIds: []server.Id{victimClientId}, + }, attackerSess) + assert.Equal(t, err, nil) + + // the victim's client must still be active + assert.NotEqual(t, GetNetworkClient(ctx, victimClientId), nil) + }) +} + // GetProvideModes / GetProvideSecretKey must fall back to postgres when the // redis cache is cold (data written before the redis layer existed, or evicted). // Regression: a redis miss used to leak a non-nil error even though the db diff --git a/model/network_create_rate_limit.go b/model/network_create_rate_limit.go new file mode 100644 index 00000000..cc418b0e --- /dev/null +++ b/model/network_create_rate_limit.go @@ -0,0 +1,91 @@ +package model + +import ( + "context" + "fmt" + "time" + + "github.com/urnetwork/server" + "github.com/urnetwork/server/session" +) + +const NetworkCreateDailyLimit = 5 +const NetworkCreateDailyWindow = 24 * time.Hour + +func maxNetworkCreateAttemptsError() error { + return fmt.Errorf("429 You have reached the maximum number of account creations for today. Please try again later.") +} + +// CheckNetworkCreateRateLimit checks if the IP has exceeded the daily account +// creation limit. It records the attempt and returns an error if over the limit. +// Must be called BEFORE the network is actually created (attempt recorded atomically). +func CheckNetworkCreateRateLimit( + ctx context.Context, + session *session.ClientSession, +) error { + clientAddressHash, _, err := session.ClientAddressHashPort() + if err != nil { + // can't determine client address — allow + return nil + } + + var count int + + server.Tx(ctx, func(tx server.PgTx) { + // Count how many network creates this IP has done in the last 24 hours + result, err := tx.Query( + ctx, + ` + SELECT COUNT(*) + FROM network_create_attempt + WHERE + client_address_hash = $1 AND + now() - INTERVAL '1 seconds' * $2 <= create_time + `, + clientAddressHash[:], + int(NetworkCreateDailyWindow/time.Second), + ) + server.WithPgResult(result, err, func() { + if result.Next() { + server.Raise(result.Scan(&count)) + } + }) + + if count >= NetworkCreateDailyLimit { + return + } + + // Record this attempt + server.RaisePgResult(tx.Exec( + ctx, + ` + INSERT INTO network_create_attempt + (network_create_attempt_id, client_address_hash, create_time) + VALUES ($1, $2, $3) + `, + server.NewId(), + clientAddressHash[:], + server.NowUtc(), + )) + }) + + if count >= NetworkCreateDailyLimit { + return maxNetworkCreateAttemptsError() + } + + return nil +} + +// RemoveExpiredNetworkCreateAttempts cleans up attempts older than the window. +func RemoveExpiredNetworkCreateAttempts(ctx context.Context, minTime time.Time) { + server.MaintenanceTx(ctx, func(tx server.PgTx) { + server.RaisePgResult(tx.Exec( + ctx, + ` + DELETE FROM network_create_attempt + WHERE create_time < $1 + `, + minTime.UTC(), + )) + }) +} diff --git a/model/network_model.go b/model/network_model.go index e514f0e5..17158f62 100644 --- a/model/network_model.go +++ b/model/network_model.go @@ -11,6 +11,7 @@ import ( // "github.com/urnetwork/glog" goaway "github.com/TwiN/go-away" + bip39 "github.com/tyler-smith/go-bip39" "github.com/urnetwork/glog" "github.com/urnetwork/server" "github.com/urnetwork/server/session" @@ -58,7 +59,7 @@ const ( func NetworkCheck(check *NetworkCheckArgs, session *session.ClientSession) (*NetworkCheckResult, error) { - _, err := validateNetworkName(check.NetworkName) + _, err := ValidateNetworkName(check.NetworkName) if err != nil { return &NetworkCheckResult{ Available: false, @@ -81,25 +82,16 @@ type NetworkCreateArgs struct { Password *string `json:"password,omitempty"` NetworkName string `json:"network_name"` Terms bool `json:"terms"` - GuestMode bool `json:"guest_mode"` VerifyUseNumeric bool `json:"verify_use_numeric"` ReferralCode *string `json:"referral_code,omitempty"` BalanceCode *string `json:"balance_code,omitempty"` WalletAuth *WalletAuthArgs `json:"wallet_auth,omitempty"` } -type UpgradeGuestArgs struct { - NetworkName string `json:"network_name"` - UserAuth *string `json:"user_auth,omitempty"` - AuthJwt *string `json:"auth_jwt,omitempty"` - AuthJwtType *string `json:"auth_jwt_type,omitempty"` - Password *string `json:"password,omitempty"` - WalletAuth *WalletAuthArgs `json:"wallet_auth,omitempty"` -} - type NetworkCreateResult struct { Network *NetworkCreateResultNetwork `json:"network,omitempty"` UserAuth *string `json:"user_auth,omitempty"` + Seedphrase *string `json:"seedphrase,omitempty"` VerificationRequired *NetworkCreateResultVerification `json:"verification_required,omitempty"` Error *NetworkCreateResultError `json:"error,omitempty"` IsPro bool `json:"is_pro,omitempty"` @@ -120,7 +112,7 @@ type NetworkCreateResultError struct { Message string `json:"message"` } -func validateNetworkName(networkName string) (string, error) { +func ValidateNetworkName(networkName string) (string, error) { trimmed := strings.TrimSpace(networkName) // to lowercase @@ -147,7 +139,6 @@ func validateNetworkName(networkName string) (string, error) { } return normalized, nil - } func NetworkCreate( @@ -156,62 +147,87 @@ func NetworkCreate( ) (*NetworkCreateResult, error) { userAuth, _ := NormalUserAuthV1(networkCreate.UserAuth) - userAuthAttemptId, allow := UserAuthAttempt(userAuth, session) - if !allow { - return nil, maxUserAuthAttemptsError() - } + // seedphrase path: no auth method provided + if networkCreate.UserAuth == nil && networkCreate.AuthJwt == nil && networkCreate.WalletAuth == nil { - if !networkCreate.Terms { - result := &NetworkCreateResult{ - Error: &NetworkCreateResultError{ - Message: AgreeToTerms, - }, + if !networkCreate.Terms { + result := &NetworkCreateResult{ + Error: &NetworkCreateResultError{ + Message: AgreeToTerms, + }, + } + return result, nil } - return result, nil - } - // create a guest network - if networkCreate.GuestMode { + // rate limit: max 5 seedphrase accounts per IP per day + if err := CheckNetworkCreateRateLimit(session.Ctx, session); err != nil { + return nil, err + } - resultNetworkCreate := networkCreateGuest(session.Ctx) + validatedNetworkName, err := generateRandomNetworkName() + if err != nil { + result := &NetworkCreateResult{ + Error: &NetworkCreateResultError{ + Message: "Failed to generate network name.", + }, + } + return result, nil + } + + resultNetworkCreate := networkCreateSeedphrase( + session.Ctx, + &networkCreate, + validatedNetworkName, + ) if resultNetworkCreate.Created { auditNetworkCreate(networkCreate, resultNetworkCreate.NetworkId, session) - // we should disable adding the guest network to name search? - // networkNameSearch.Add(session.Ctx, networkCreate.NetworkName, createdNetworkId, 0) - isPro := false - isGuest := true - byJwt := jwt.NewByJwt( resultNetworkCreate.NetworkId, resultNetworkCreate.UserId, - networkCreate.NetworkName, - isGuest, + validatedNetworkName, + false, isPro, ) byJwtSigned := byJwt.Sign() result := &NetworkCreateResult{ + Seedphrase: &resultNetworkCreate.Seedphrase, Network: &NetworkCreateResultNetwork{ ByJwt: &byJwtSigned, - NetworkName: resultNetworkCreate.NetworkName, + NetworkName: validatedNetworkName, NetworkId: resultNetworkCreate.NetworkId, + IsPro: resultNetworkCreate.IsPro, }, } - return result, nil } else { result := &NetworkCreateResult{ Error: &NetworkCreateResultError{ - Message: "An error occurred creating a guest network", + Message: "Account might already exist. Please start over.", }, } return result, nil } } - validatedNetworkName, error := validateNetworkName(networkCreate.NetworkName) + // email/SSO/wallet paths use the existing auth rate limit + userAuthAttemptId, allow := UserAuthAttempt(userAuth, session) + if !allow { + return nil, maxUserAuthAttemptsError() + } + + if !networkCreate.Terms { + result := &NetworkCreateResult{ + Error: &NetworkCreateResultError{ + Message: AgreeToTerms, + }, + } + return result, nil + } + + validatedNetworkName, error := ValidateNetworkName(networkCreate.NetworkName) if error != nil { result := &NetworkCreateResult{ @@ -483,6 +499,7 @@ type networkCreateResult struct { NetworkId server.Id NetworkName string UserId server.Id + Seedphrase string IsPro bool } @@ -841,66 +858,64 @@ func networkCreateUserAuth( } /** - * network create guest mode + * network create seedphrase */ - -func networkCreateGuest( +func networkCreateSeedphrase( ctx context.Context, + networkCreate *NetworkCreateArgs, + validatedNetworkName string, ) networkCreateResult { - created := false var createdNetworkId server.Id - var networkName string var createdUserId server.Id + var seedphrase string + isPro := false server.Tx(ctx, func(tx server.PgTx) { - var err error - createdUserId = server.NewId() createdNetworkId = server.NewId() - // remove dashes from the network name - networkName = fmt.Sprintf( - "g%s", - strings.ReplaceAll(server.NewId().String(), "-", ""), - ) + // generate seedphrase + entropy, err := bip39.NewEntropy(256) + server.Raise(err) + seedphrase, err = bip39.NewMnemonic(entropy) + server.Raise(err) _, err = tx.Exec( ctx, - ` - INSERT INTO network_user - (user_id, user_name, auth_type) - VALUES ($1, $2, $3) - `, - createdUserId, - "guest", - AuthTypeGuest, + `INSERT INTO network_user (user_id, user_name, auth_type) + VALUES ($1, $2, $3)`, + createdUserId, validatedNetworkName, AuthTypeSeedphrase, ) server.Raise(err) + err = CreateSeedphraseAuthInTx(tx, ctx, createdUserId, seedphrase) + server.Raise(err) + _, err = tx.Exec( ctx, - ` - INSERT INTO network - (network_id, network_name, admin_user_id) - VALUES ($1, $2, $3) - `, - createdNetworkId, - networkName, - createdUserId, + `INSERT INTO network (network_id, network_name, admin_user_id) + VALUES ($1, $2, $3)`, + createdNetworkId, validatedNetworkName, createdUserId, ) server.Raise(err) CreateNetworkReferralCodeInTx(ctx, tx, createdNetworkId) + isPro = networkCreateRedeemBalanceCodeInTx( + networkCreate, createdNetworkId, ctx, tx, + ) + created = true }) return networkCreateResult{ Created: created, NetworkId: createdNetworkId, - NetworkName: networkName, + NetworkName: validatedNetworkName, UserId: createdUserId, + Seedphrase: seedphrase, + IsPro: isPro, } } @@ -978,7 +993,7 @@ func checkNetworkNameAvailability( var existingNetworkId *server.Id - validatedNetworkName, validationErr := validateNetworkName(networkName) + validatedNetworkName, validationErr := ValidateNetworkName(networkName) if validationErr != nil { err = validationErr return @@ -1054,560 +1069,6 @@ func NetworkUpdate( return networkCreateResult, nil } -type UpgradeGuestResult struct { - Error *UpgradeGuestError `json:"error,omitempty"` - VerificationRequired *UpgradeGuestResultVerification `json:"verification_required,omitempty"` - Network *UpgradeGuestNetwork `json:"network,omitempty"` - UserAuth *string `json:"user_auth,omitempty"` -} - -type UpgradeGuestNetwork struct { - ByJwt *string `json:"by_jwt,omitempty"` -} - -type UpgradeGuestResultVerification struct { - UserAuth string `json:"user_auth"` -} - -type UpgradeGuestError struct { - Message string `json:"message"` -} - -func UpgradeGuest( - upgradeGuest UpgradeGuestArgs, - session *session.ClientSession, -) (*UpgradeGuestResult, error) { - - userAuth, _ := NormalUserAuthV1(upgradeGuest.UserAuth) - - userAuthAttemptId, allow := UserAuthAttempt(userAuth, session) - if !allow { - return nil, maxUserAuthAttemptsError() - } - - var result = &UpgradeGuestResult{} - networkName := strings.TrimSpace(upgradeGuest.NetworkName) - - err := checkNetworkNameAvailability(networkName, session) - if err != nil { - result = &UpgradeGuestResult{ - Error: &UpgradeGuestError{ - Message: err.Error(), - }, - } - return result, nil - } - - // Social/wallet upgrades issue a durable JWT. Resolve the fresh - // entitlement before the update transaction so the lookup's PostgreSQL - // and Redis I/O cannot nest under an open transaction. - isPro := false - if (upgradeGuest.AuthJwt != nil && upgradeGuest.AuthJwtType != nil) || - upgradeGuest.WalletAuth != nil { - isPro = IsProFresh(session.Ctx, &session.ByJwt.NetworkId) - } - server.Tx(session.Ctx, func(tx server.PgTx) { - - if upgradeGuest.UserAuth != nil { - - /** - * Upgrade from guest from email + password - */ - - if userAuth == nil { - result = &UpgradeGuestResult{ - Error: &UpgradeGuestError{ - Message: "Invalid email or phone number.", - }, - } - return - } - - /** - * Validate the user does not exist - */ - - var userId *server.Id - - userCheck, err := tx.Query( - session.Ctx, - ` - SELECT user_id FROM network_user WHERE user_auth = $1 - `, - userAuth, - ) - server.WithPgResult(userCheck, err, func() { - if userCheck.Next() { - server.Raise(userCheck.Scan(&userId)) - } - }) - - if userId != nil { - result = &UpgradeGuestResult{ - Error: &UpgradeGuestError{ - Message: "User already exists", - }, - } - return - } - - /** - * Handle password - */ - - passwordSalt := createPasswordSalt() - passwordHash := computePasswordHashV1([]byte(*upgradeGuest.Password), passwordSalt) - - /** - * Update network user from guest - */ - - server.RaisePgResult(tx.Exec( - session.Ctx, - ` - UPDATE network_user - SET - auth_type = $2, - user_auth = $3, - password_hash = $4, - password_salt = $5 - - WHERE - user_id = $1 - `, - session.ByJwt.UserId, - AuthTypePassword, - userAuth, - passwordHash, - passwordSalt, - )) - - // do we need to run auditNetworkCreate on the upgrade from guest -> normal account? - - networkNameSearch().Add(session.Ctx, networkName, session.ByJwt.NetworkId, 0) - - result = &UpgradeGuestResult{ - VerificationRequired: &UpgradeGuestResultVerification{ - UserAuth: *userAuth, - }, - } - - } else if upgradeGuest.AuthJwt != nil && upgradeGuest.AuthJwtType != nil { - - /** - * Upgrade from guest from social login - */ - - authJwt, _ := ParseAuthJwt(*upgradeGuest.AuthJwt, AuthType(*upgradeGuest.AuthJwtType)) - - if authJwt != nil { - - /** - * Validate the user does not exist - */ - normalJwtUserAuth, _ := NormalUserAuth(authJwt.UserAuth) - var userId *server.Id - - userCheck, err := tx.Query( - session.Ctx, - ` - SELECT user_id FROM network_user WHERE user_auth = $1 - `, - normalJwtUserAuth, - ) - server.WithPgResult(userCheck, err, func() { - if userCheck.Next() { - server.Raise(userCheck.Scan(&userId)) - } - }) - - if userId != nil { - result = &UpgradeGuestResult{ - Error: &UpgradeGuestError{ - Message: "User already exists", - }, - } - return - } - - /** - * Update network user from guest - */ - server.RaisePgResult(tx.Exec( - session.Ctx, - ` - UPDATE network_user - SET - auth_type = $2, - user_auth = $3, - auth_jwt = $4 - - WHERE - user_id = $1 - `, - session.ByJwt.UserId, - upgradeGuest.AuthJwtType, - normalJwtUserAuth, - upgradeGuest.AuthJwt, - )) - - setUserAuthAttemptSuccessInTx( - session.Ctx, - tx, - userAuthAttemptId, - true, - ) - - isGuest := false - - byJwt := jwt.NewByJwt( - session.ByJwt.NetworkId, - session.ByJwt.UserId, - networkName, - isGuest, - isPro, - ) - byJwtSigned := byJwt.Sign() - - result = &UpgradeGuestResult{ - Network: &UpgradeGuestNetwork{ - ByJwt: &byJwtSigned, - }, - UserAuth: &normalJwtUserAuth, - } - - } - - } else if upgradeGuest.WalletAuth != nil { - /** - * Upgrade from guest from wallet - */ - - /** - * verify the wallet signature (proof of key control) before binding the - * wallet address to this account, matching NetworkCreate and handleLoginWallet - */ - isValid, err := VerifySignature( - upgradeGuest.WalletAuth.Blockchain, - upgradeGuest.WalletAuth.PublicKey, - upgradeGuest.WalletAuth.Message, - upgradeGuest.WalletAuth.Signature, - ) - if err != nil || !isValid { - result = &UpgradeGuestResult{ - Error: &UpgradeGuestError{ - Message: "invalid wallet signature", - }, - } - return - } - - var userId *server.Id - - userCheck, err := tx.Query( - session.Ctx, - ` - SELECT user_id FROM network_user WHERE wallet_address = $1 - `, - upgradeGuest.WalletAuth.PublicKey, - ) - server.WithPgResult(userCheck, err, func() { - if userCheck.Next() { - server.Raise(userCheck.Scan(&userId)) - } - }) - - if userId != nil { - result = &UpgradeGuestResult{ - Error: &UpgradeGuestError{ - Message: "User already exists", - }, - } - return - } - - /** - * Update network user from guest - */ - server.RaisePgResult(tx.Exec( - session.Ctx, - ` - UPDATE network_user - SET - wallet_address = $2, - wallet_blockchain = $3, - auth_type = $4 - - WHERE - user_id = $1 - `, - session.ByJwt.UserId, - upgradeGuest.WalletAuth.PublicKey, - AuthTypeSolana, - AuthTypeSolana, - )) - - setUserAuthAttemptSuccessInTx( - session.Ctx, - tx, - userAuthAttemptId, - true, - ) - - isGuest := false - - byJwt := jwt.NewByJwt( - session.ByJwt.NetworkId, - session.ByJwt.UserId, - networkName, - isGuest, - isPro, - ) - byJwtSigned := byJwt.Sign() - - result = &UpgradeGuestResult{ - Network: &UpgradeGuestNetwork{ - ByJwt: &byJwtSigned, - }, - // UserAuth: &normalJwtUserAuth, - } - - } - - /** - * Update the network name - */ - server.RaisePgResult(tx.Exec( - session.Ctx, - ` - UPDATE network - SET - network_name = $2 - - WHERE - network_id = $1 - `, - session.ByJwt.NetworkId, - networkName, - )) - - }) - - return result, nil -} - -/** - * Upgrade guest with existing account - */ -type UpgradeGuestExistingArgs struct { - UserAuth *string `json:"user_auth,omitempty"` - Password *string `json:"password,omitempty"` - AuthJwt *string `json:"auth_jwt,omitempty"` - AuthJwtType *string `json:"auth_jwt_type,omitempty"` - WalletAuth *WalletAuthArgs `json:"wallet_auth,omitempty"` -} - -type UpgradeGuestExistingError struct { - Message string `json:"message"` -} - -type UpgradeGuestExistingResult struct { - Error *UpgradeGuestExistingError `json:"error,omitempty"` - VerificationRequired *UpgradeGuestExistingVerificationRequired `json:"verification_required,omitempty"` - Network *UpgradeGuestExistingResultNetwork `json:"network,omitempty"` - // Error *AuthLoginWithPasswordResultError `json:"error,omitempty"` -} - -type UpgradeGuestExistingVerificationRequired struct { - UserAuth string `json:"user_auth"` -} - -type UpgradeGuestExistingResultNetwork struct { - ByJwt *string `json:"by_jwt,omitempty"` - // NetworkName *string `json:"name,omitempty"` -} - -func UpgradeFromGuestExisting( - upgradeGuestExisting UpgradeGuestExistingArgs, - session *session.ClientSession, -) (*UpgradeGuestExistingResult, error) { - - if upgradeGuestExisting.UserAuth != nil && upgradeGuestExisting.Password != nil { - /** - * Upgrade from guest from email + password - */ - - args := AuthLoginWithPasswordArgs{ - UserAuth: *upgradeGuestExisting.UserAuth, - Password: *upgradeGuestExisting.Password, - } - - loginResult, err := AuthLoginWithPassword(args, session) - if err != nil { - return &UpgradeGuestExistingResult{ - Error: &UpgradeGuestExistingError{ - Message: "Invalid login", - }, - }, nil - } - - if loginResult.Error != nil { - return &UpgradeGuestExistingResult{ - Error: &UpgradeGuestExistingError{ - Message: loginResult.Error.Message, - }, - }, nil - } - - if loginResult.Network.ByJwt == nil { - - return &UpgradeGuestExistingResult{ - Error: &UpgradeGuestExistingError{ - Message: "Invalid network token", - }, - }, nil - } - - network, err := jwt.ParseByJwt(session.Ctx, *loginResult.Network.ByJwt) - if err != nil { - return &UpgradeGuestExistingResult{ - Error: &UpgradeGuestExistingError{ - Message: "Error parsing network token", - }, - }, nil - } - - err = markUpgradedNetworkId(network.NetworkId, session) - if err != nil { - return &UpgradeGuestExistingResult{ - Error: &UpgradeGuestExistingError{ - Message: "Error marking upgraded network id", - }, - }, nil - } - - result := &UpgradeGuestExistingResult{ - Network: &UpgradeGuestExistingResultNetwork{ - ByJwt: loginResult.Network.ByJwt, - }, - } - - if loginResult.VerificationRequired != nil { - result.VerificationRequired = &UpgradeGuestExistingVerificationRequired{ - UserAuth: loginResult.VerificationRequired.UserAuth, - } - } - - return result, nil - - } - - if upgradeGuestExisting.AuthJwt != nil && upgradeGuestExisting.AuthJwtType != nil { - /** - * Upgrade from guest from social login - */ - - args := AuthLoginArgs{ - AuthJwt: upgradeGuestExisting.AuthJwt, - AuthJwtType: upgradeGuestExisting.AuthJwtType, - } - - return handleAuthLoginUpgrade(args, session) - - } - - if upgradeGuestExisting.WalletAuth != nil { - - args := AuthLoginArgs{ - WalletAuth: upgradeGuestExisting.WalletAuth, - } - - return handleAuthLoginUpgrade(args, session) - - } - - return &UpgradeGuestExistingResult{ - Error: &UpgradeGuestExistingError{ - Message: "Invalid args", - }, - }, nil - -} - -func handleAuthLoginUpgrade( - args AuthLoginArgs, - session *session.ClientSession, -) (*UpgradeGuestExistingResult, error) { - loginResult, err := AuthLogin(args, session) - if err != nil { - return &UpgradeGuestExistingResult{ - Error: &UpgradeGuestExistingError{ - Message: "Invalid login", - }, - }, nil - } - - if loginResult.Error != nil { - return &UpgradeGuestExistingResult{ - Error: &UpgradeGuestExistingError{ - Message: loginResult.Error.Message, - }, - }, nil - } - - // in this case, we should navigate the user to the network creation view - if loginResult.Network == nil { - return &UpgradeGuestExistingResult{}, nil - } - - network, err := jwt.ParseByJwt(session.Ctx, loginResult.Network.ByJwt) - if err != nil { - return &UpgradeGuestExistingResult{ - Error: &UpgradeGuestExistingError{ - Message: "Error parsing network token", - }, - }, nil - } - - err = markUpgradedNetworkId(network.NetworkId, session) - if err != nil { - return &UpgradeGuestExistingResult{ - Error: &UpgradeGuestExistingError{ - Message: "Error marking upgraded network id", - }, - }, nil - } - - return &UpgradeGuestExistingResult{ - Network: &UpgradeGuestExistingResultNetwork{ - ByJwt: &loginResult.Network.ByJwt, - }, - }, nil -} - -func markUpgradedNetworkId( - upgradedNetworkId server.Id, - session *session.ClientSession, -) error { - - server.Tx(session.Ctx, func(tx server.PgTx) { - server.RaisePgResult(tx.Exec( - session.Ctx, - ` - UPDATE network - SET - guest_upgrade_network_id = $1 - WHERE - network_id = $2 - `, - upgradedNetworkId, - session.ByJwt.NetworkId, - )) - }) - - return nil -} - type Network struct { NetworkId *server.Id `json:"network_id"` NetworkName string `json:"network_name"` diff --git a/model/network_model_test.go b/model/network_model_test.go index 5c2d0c0e..1624866a 100644 --- a/model/network_model_test.go +++ b/model/network_model_test.go @@ -2,369 +2,20 @@ package model import ( "context" - "encoding/base64" - "fmt" - "regexp" "testing" - "github.com/gagliardetto/solana-go" - "github.com/urnetwork/connect" + "github.com/go-playground/assert/v2" "github.com/urnetwork/server" "github.com/urnetwork/server/jwt" "github.com/urnetwork/server/session" ) -func TestNetworkCreateGuestMode(t *testing.T) { - server.DefaultTestEnv().Run(t, func(t testing.TB) { - ctx := context.Background() - - networkCreate := NetworkCreateArgs{ - Terms: true, - GuestMode: true, - } - - byJwt := jwt.ByJwt{} - - clientSession := session.Testing_CreateClientSession(ctx, &byJwt) - - pattern := `^g[0-9a-f]+$` - - // Compile the regex - regex, err := regexp.Compile(pattern) - connect.AssertEqual(t, err, nil) - - result, err := NetworkCreate(networkCreate, clientSession) - connect.AssertEqual(t, err, nil) - - connect.AssertEqual(t, regex.MatchString(result.Network.NetworkName), true) - - }) -} - -func TestNetworkUpgradeGuestMode(t *testing.T) { - server.DefaultTestEnv().Run(t, func(t testing.TB) { - ctx := context.Background() - - networkId := server.NewId() - userId := server.NewId() - - byJwt := jwt.ByJwt{ - NetworkId: networkId, - UserId: userId, - } - - clientSession := session.Testing_CreateClientSession(ctx, &byJwt) - - Testing_CreateGuestNetwork( - ctx, - networkId, - "abcdef", - userId, - ) - - // fetch the network user and make sure it's a guest - - network := GetNetwork(clientSession) - connect.AssertNotEqual(t, network, nil) - - networkUser := GetNetworkUser(ctx, *network.AdminUserId) - - connect.AssertEqual(t, networkUser.AuthType, AuthTypeGuest) - - // upgrade to non-guest - - userAuth := "test@ur.io" - password := "abcdefg1234567" - upgradedNetworkName := "abcdef1234" - - upgradeGuestArgs := UpgradeGuestArgs{ - NetworkName: upgradedNetworkName, - UserAuth: &userAuth, - Password: &password, - } - - upgradeGuestResult, err := UpgradeGuest( - upgradeGuestArgs, - clientSession, - ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, upgradeGuestResult.Error, nil) - connect.AssertEqual(t, upgradeGuestResult.VerificationRequired.UserAuth, userAuth) - - // fetch the network user and make sure it's no longer a guest - networkUser = GetNetworkUser(ctx, userId) - connect.AssertEqual(t, networkUser.AuthType, AuthTypePassword) - - // ensure network name has been updated - network = GetNetwork(clientSession) - connect.AssertEqual(t, network.NetworkName, upgradedNetworkName) - - }) -} - -func TestUpgradeGuestExistingUser(t *testing.T) { - server.DefaultTestEnv().Run(t, func(t testing.TB) { - ctx := context.Background() - networkId := server.NewId() - userId := server.NewId() - // clientId := server.NewId() - networkName := "abcdef" - - Testing_CreateNetwork(ctx, networkId, networkName, userId) - - guestNetworkId := server.NewId() - guestUserId := server.NewId() - - byJwt := jwt.ByJwt{ - NetworkId: guestNetworkId, - UserId: guestUserId, - } - - clientSession := session.Testing_CreateClientSession(ctx, &byJwt) - - Testing_CreateGuestNetwork( - ctx, - guestNetworkId, - fmt.Sprintf("guest-%s", networkId.String()), - guestUserId, - ) - - // fetch the network - // it should have no guest upgrade network id - network := GetNetwork(clientSession) - connect.AssertNotEqual(t, network, nil) - connect.AssertEqual(t, network.GuestUpgradeNetworkId, nil) - - userAuth := fmt.Sprintf("%s@bringyour.com", networkId) // pulled from Testing_CreateNetwork - password := "password" // pulled from Testing_CreateNetwork - args := UpgradeGuestExistingArgs{ - UserAuth: &userAuth, - Password: &password, - } - - result, err := UpgradeFromGuestExisting(args, clientSession) - connect.AssertEqual(t, err, nil) - // fmt.Println("error UpgradeFromGuestExisting: ", result.Error.Message) - connect.AssertEqual(t, result.Error, nil) - - // fetch the network - // it should have the guest upgrade network id - network = GetNetwork(clientSession) - connect.AssertNotEqual(t, network, nil) - - connect.AssertEqual(t, network.GuestUpgradeNetworkId, networkId) - - }) -} - -func TestUpgradeGuestExistingWalletUser(t *testing.T) { - server.DefaultTestEnv().Run(t, func(t testing.TB) { - ctx := context.Background() - networkId := server.NewId() - userId := server.NewId() - networkName := "abcdef" - - privateKey, err := solana.NewRandomPrivateKey() - connect.AssertEqual(t, err, nil) - publicKey := privateKey.PublicKey().String() - - challenge := CreateWalletAuthChallenge(WalletAuthChallengeArgs{}, ctx) - connect.AssertEqual(t, challenge.Error, nil) - signature, err := privateKey.Sign([]byte(challenge.MessageTemplate)) - connect.AssertEqual(t, err, nil) - signatureB64 := base64.StdEncoding.EncodeToString(signature[:]) - - Testing_CreateNetworkByWallet( - ctx, - networkId, - networkName, - userId, - publicKey, - signatureB64, - challenge.MessageTemplate, - ) - - guestNetworkId := server.NewId() - guestUserId := server.NewId() - - byJwt := jwt.ByJwt{ - NetworkId: guestNetworkId, - UserId: guestUserId, - } - - clientSession := session.Testing_CreateClientSession(ctx, &byJwt) - - Testing_CreateGuestNetwork( - ctx, - guestNetworkId, - fmt.Sprintf("guest-%s", networkId.String()), - guestUserId, - ) - - // fetch the network - // it should have no guest upgrade network id - network := GetNetwork(clientSession) - connect.AssertNotEqual(t, network, nil) - connect.AssertEqual(t, network.GuestUpgradeNetworkId, nil) - - // login with a fresh challenge - challenge = CreateWalletAuthChallenge(WalletAuthChallengeArgs{}, ctx) - connect.AssertEqual(t, challenge.Error, nil) - signature, err = privateKey.Sign([]byte(challenge.MessageTemplate)) - connect.AssertEqual(t, err, nil) - signatureB64 = base64.StdEncoding.EncodeToString(signature[:]) - - args := UpgradeGuestExistingArgs{ - WalletAuth: &WalletAuthArgs{ - PublicKey: publicKey, - Signature: signatureB64, - Message: challenge.MessageTemplate, - Blockchain: AuthTypeSolana, - }, - } - - result, err := UpgradeFromGuestExisting(args, clientSession) - - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, result.Error, nil) - - // fetch the network - // it should have the guest upgrade network id - network = GetNetwork(clientSession) - connect.AssertNotEqual(t, network, nil) - - connect.AssertEqual(t, network.GuestUpgradeNetworkId, networkId) - - }) -} - -func TestUpgradeGuestByWallet(t *testing.T) { - server.DefaultTestEnv().Run(t, func(t testing.TB) { - - ctx := context.Background() - networkId := server.NewId() - userId := server.NewId() - - byJwt := jwt.ByJwt{ - NetworkId: networkId, - UserId: userId, - } - - clientSession := session.Testing_CreateClientSession(ctx, &byJwt) - - Testing_CreateGuestNetwork( - ctx, - networkId, - fmt.Sprintf("guest-%s", networkId.String()), - userId, - ) - - privateKey, err := solana.NewRandomPrivateKey() - connect.AssertEqual(t, err, nil) - publicKey := privateKey.PublicKey().String() - networkName := "abcdef" - - challenge := CreateWalletAuthChallenge(WalletAuthChallengeArgs{}, ctx) - connect.AssertEqual(t, challenge.Error, nil) - signature, err := privateKey.Sign([]byte(challenge.MessageTemplate)) - connect.AssertEqual(t, err, nil) - signatureB64 := base64.StdEncoding.EncodeToString(signature[:]) - - args := UpgradeGuestArgs{ - NetworkName: networkName, - WalletAuth: &WalletAuthArgs{ - PublicKey: publicKey, - Signature: signatureB64, - Message: challenge.MessageTemplate, - Blockchain: "solana", - }, - } - - networkUser := GetNetworkUser(ctx, userId) - connect.AssertEqual(t, networkUser.AuthType, AuthTypeGuest) - connect.AssertEqual(t, networkUser.WalletAddress, nil) - - result, err := UpgradeGuest(args, clientSession) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, result.Error, nil) - - network := GetNetwork(clientSession) - connect.AssertNotEqual(t, network, nil) - connect.AssertEqual(t, network.NetworkName, networkName) - - user := GetNetworkUser(ctx, *network.AdminUserId) - connect.AssertEqual(t, user.AuthType, "solana") - connect.AssertEqual(t, user.WalletAddress, publicKey) - - }) -} - -// TestUpgradeGuestByWalletInvalidSignatureRejected reproduces a Class-A auth bug -// adjacent to the adopt_secret / confirm-share fixes: UpgradeGuest's wallet branch -// (network_model.go ~1261) binds WalletAuth.PublicKey to the caller's own account -// WITHOUT ever calling VerifySignature — unlike NetworkCreate's wallet branch -// (network_model.go:365) and handleLoginWallet (auth_model.go:457), which both verify. -// So a guest can claim any wallet address (e.g. a victim's) with a bogus signature, -// squatting the wallet identity and blocking the real owner from ever registering it -// (networkCreateWalletAuth gates on network_user.wallet_address). -// -// This asserts the secure expectation (invalid signature => wallet NOT bound), so it -// FAILS on the current code and will pass once UpgradeGuest verifies the signature. -func TestUpgradeGuestByWalletInvalidSignatureRejected(t *testing.T) { - server.DefaultTestEnv().Run(t, func(t testing.TB) { - ctx := context.Background() - - networkId := server.NewId() - userId := server.NewId() - clientSession := session.Testing_CreateClientSession(ctx, &jwt.ByJwt{ - NetworkId: networkId, - UserId: userId, - }) - Testing_CreateGuestNetwork( - ctx, - networkId, - fmt.Sprintf("guest-%s", networkId.String()), - userId, - ) - - // A real-format Solana address the caller does NOT control, paired with a - // well-formed but cryptographically invalid signature (64 zero bytes). - victimWallet := "6UJtwDRMv2CCfVCKm6hgMDAGrFzv7z8WKEHut2u8dV8s" - invalidSignature := base64.StdEncoding.EncodeToString(make([]byte, 64)) - - before := GetNetworkUser(ctx, userId) - connect.AssertEqual(t, before.AuthType, AuthTypeGuest) - connect.AssertEqual(t, before.WalletAddress, nil) - - args := UpgradeGuestArgs{ - NetworkName: "abcdef", - WalletAuth: &WalletAuthArgs{ - PublicKey: victimWallet, - Signature: invalidSignature, - Message: "Welcome to URnetwork", - Blockchain: "solana", - }, - } - result, err := UpgradeGuest(args, clientSession) - - // Security invariant: without a valid signature proving control of the key, - // the guest account must NOT be bound to that wallet address. - user := GetNetworkUser(ctx, userId) - if user.WalletAddress != nil { - t.Fatalf("SECURITY: UpgradeGuest bound wallet_address=%q to the guest account with an "+ - "INVALID signature — no proof of key control. The wallet branch never calls "+ - "VerifySignature (NetworkCreate and handleLoginWallet do). auth_type=%q err=%v result=%+v", - *user.WalletAddress, user.AuthType, err, result) - } - }) -} - func TestNetworkCreateTermsFail(t *testing.T) { server.DefaultTestEnv().Run(t, func(t testing.TB) { ctx := context.Background() networkCreate := NetworkCreateArgs{ - GuestMode: true, + Terms: false, } byJwt := jwt.ByJwt{} @@ -372,8 +23,8 @@ func TestNetworkCreateTermsFail(t *testing.T) { clientSession := session.Testing_CreateClientSession(ctx, &byJwt) result, err := NetworkCreate(networkCreate, clientSession) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, result.Error.Message, AgreeToTerms) + assert.Equal(t, err, nil) + assert.Equal(t, result.Error.Message, AgreeToTerms) }) } @@ -400,8 +51,8 @@ func TestNetworkUpdate(t *testing.T) { NetworkName: networkName, } result, err := NetworkUpdate(networkUpdateArgs, sourceSession) - connect.AssertEqual(t, err, nil) - connect.AssertNotEqual(t, result.Error, nil) + assert.Equal(t, err, nil) + assert.NotEqual(t, result.Error, nil) // fail // network name should be at least 6 characters @@ -409,8 +60,8 @@ func TestNetworkUpdate(t *testing.T) { NetworkName: "a", } result, err = NetworkUpdate(networkUpdateArgs, sourceSession) - connect.AssertEqual(t, err, nil) - connect.AssertNotEqual(t, result.Error, nil) + assert.Equal(t, err, nil) + assert.NotEqual(t, result.Error, nil) // success newName := "uvwxyz" @@ -418,11 +69,11 @@ func TestNetworkUpdate(t *testing.T) { NetworkName: newName, } result, err = NetworkUpdate(networkUpdateArgs, sourceSession) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, result.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, result.Error, nil) network := GetNetwork(sourceSession) - connect.AssertEqual(t, network.NetworkName, newName) + assert.Equal(t, network.NetworkName, newName) }) } @@ -432,39 +83,39 @@ func TestNetworkNameValidation(t *testing.T) { // too short networkName := "" - _, err := validateNetworkName(networkName) - connect.AssertNotEqual(t, err, nil) + _, err := ValidateNetworkName(networkName) + assert.NotEqual(t, err, nil) // too long networkName = "a123456789012345678901234567890123456789012345678901" - _, err = validateNetworkName(networkName) - connect.AssertNotEqual(t, err, nil) + _, err = ValidateNetworkName(networkName) + assert.NotEqual(t, err, nil) /** * testing special characters */ networkName = "abcde$" - _, err = validateNetworkName(networkName) - connect.AssertNotEqual(t, err, nil) + _, err = ValidateNetworkName(networkName) + assert.NotEqual(t, err, nil) networkName = "abcdeé" - _, err = validateNetworkName(networkName) - connect.AssertNotEqual(t, err, nil) + _, err = ValidateNetworkName(networkName) + assert.NotEqual(t, err, nil) networkName = "東京タワー" - _, err = validateNetworkName(networkName) - connect.AssertNotEqual(t, err, nil) + _, err = ValidateNetworkName(networkName) + assert.NotEqual(t, err, nil) // test spaces networkName = "abc def" expected := "abc-def" - validated, err := validateNetworkName(networkName) - connect.AssertEqual(t, validated, expected) + validated, err := ValidateNetworkName(networkName) + assert.Equal(t, validated, expected) // valid name should pass networkName = "abcdef" - _, err = validateNetworkName(networkName) - connect.AssertEqual(t, err, nil) + _, err = ValidateNetworkName(networkName) + assert.Equal(t, err, nil) }) } diff --git a/model/network_user_model.go b/model/network_user_model.go index c23274c5..e18dba94 100644 --- a/model/network_user_model.go +++ b/model/network_user_model.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "time" "github.com/urnetwork/glog" "github.com/urnetwork/server" @@ -11,15 +12,17 @@ import ( ) type NetworkUser struct { - UserId server.Id `json:"user_id"` - UserAuth *string `json:"user_auth,omitempty"` - Verified bool `json:"verified"` - AuthType string `json:"auth_type"` - NetworkName string `json:"network_name"` - WalletAddress *string `json:"wallet_address,omitempty"` - UserAuths []NetworkUserUserAuth `json:"user_auths,omitempty"` - SsoAuths []NetworkUserSsoAuth `json:"sso_auths,omitempty"` - WalletAuths []NetworkUserWalletAuth `json:"wallet_auths,omitempty"` + UserId server.Id `json:"user_id"` + UserAuth *string `json:"user_auth,omitempty"` + Verified bool `json:"verified"` + AuthType string `json:"auth_type"` + NetworkName string `json:"network_name"` + WalletAddress *string `json:"wallet_address,omitempty"` + UserAuths []NetworkUserUserAuth `json:"user_auths,omitempty"` + SsoAuths []NetworkUserSsoAuth `json:"sso_auths,omitempty"` + WalletAuths []NetworkUserWalletAuth `json:"wallet_auths,omitempty"` + SeedphraseAuths []NetworkUserSeedphraseAuth `json:"seedphrase_auths,omitempty"` + AuthTypes []string `json:"auth_types"` } type NetworkUserUserAuth struct { @@ -42,6 +45,11 @@ type NetworkUserWalletAuth struct { Blockchain string `json:"blockchain"` } +type NetworkUserSeedphraseAuth struct { + Seedphrase string `json:"-"` + CreateTime time.Time `json:"create_time"` +} + func GetNetworkUser( ctx context.Context, userId server.Id, @@ -111,6 +119,29 @@ func GetNetworkUser( server.Raise(err) networkUser.WalletAuths = walletAuths + /** + * Get seedphrase auths for the user + */ + seedphraseAuths, err := getSeedphraseAuths(ctx, userId) + server.Raise(err) + networkUser.SeedphraseAuths = seedphraseAuths + + /** + * Build composite auth_types list + */ + for _, a := range userAuths { + networkUser.AuthTypes = append(networkUser.AuthTypes, string(a.AuthType)) + } + for _, a := range ssoAuths { + networkUser.AuthTypes = append(networkUser.AuthTypes, string(a.AuthType)) + } + if len(walletAuths) > 0 { + networkUser.AuthTypes = append(networkUser.AuthTypes, "solana") + } + if len(seedphraseAuths) > 0 { + networkUser.AuthTypes = append(networkUser.AuthTypes, "seedphrase") + } + }) return networkUser @@ -139,7 +170,26 @@ type AddAuthMethodError struct { func AddAuth( authArgs AddAuthMethod, session *session.ClientSession, -) (*AddAuthMethodResult, error) { +) (result *AddAuthMethodResult, resultErr error) { + if err := CheckAccountActionRateLimit( + session.Ctx, + session.ByJwt.UserId, + AccountActionAddAuth, + AccountActionAddAuthDailyLimit, + AccountActionDailyWindow, + ); err != nil { + return &AddAuthMethodResult{ + Error: &AddAuthMethodError{ + Message: err.Error(), + }, + }, nil + } + + defer func() { + if resultErr == nil && result != nil && result.Error == nil { + RecordAccountActionAttempt(session.Ctx, session.ByJwt.UserId, AccountActionAddAuth) + } + }() if authArgs.UserAuth != nil && authArgs.Password != nil { /** @@ -159,7 +209,7 @@ func AddAuth( passwordSalt := createPasswordSalt() passwordHash := computePasswordHashV1([]byte(*authArgs.Password), passwordSalt) - addUserAuth( + err := addUserAuth( &AddUserAuthArgs{ UserId: session.ByJwt.UserId, UserAuth: authArgs.UserAuth, @@ -168,6 +218,13 @@ func AddAuth( }, session.Ctx, ) + if err != nil { + return &AddAuthMethodResult{ + Error: &AddAuthMethodError{ + Message: err.Error(), + }, + }, nil + } return &AddAuthMethodResult{}, nil } else if authArgs.AuthJwt != nil && authArgs.AuthJwtType != nil { @@ -191,7 +248,7 @@ func AddAuth( }, nil } - addSsoAuth( + err = addSsoAuth( &AddSsoAuthArgs{ ParsedAuthJwt: *parsedAuthJwt, AuthJwtType: SsoAuthType(*authArgs.AuthJwtType), @@ -200,17 +257,31 @@ func AddAuth( }, session.Ctx, ) + if err != nil { + return &AddAuthMethodResult{ + Error: &AddAuthMethodError{ + Message: err.Error(), + }, + }, nil + } return &AddAuthMethodResult{}, nil } else if authArgs.WalletAuth != nil { // user is adding a wallet auth method - addWalletAuth( + err := addWalletAuth( &AddWalletAuthArgs{ WalletAuth: authArgs.WalletAuth, UserId: session.ByJwt.UserId, }, session.Ctx, ) + if err != nil { + return &AddAuthMethodResult{ + Error: &AddAuthMethodError{ + Message: err.Error(), + }, + }, nil + } return &AddAuthMethodResult{}, nil } @@ -663,6 +734,40 @@ func addWalletAuth( server.Tx(ctx, func(tx server.PgTx) { + // Check for an existing binding of this wallet to a different user + // before attempting the INSERT below, mirroring the + // validateUserAuthAvailability pre-check addUserAuthInTx does for + // email/phone auth. A raw unique-constraint violation from the + // INSERT would abort this transaction; server.Tx's default retry + // options blindly retry the subsequent failed COMMIT for up to 60s + // regardless of whether the underlying cause is actually + // retryable, so on the ordinary (non-racing) "wallet already taken" + // path this would otherwise stall the request for up to a minute + // with no error ever reaching the client instead of failing fast. + var conflictUserId *server.Id + result, queryErr := tx.Query( + ctx, + ` + SELECT user_id FROM network_user_auth_wallet + WHERE wallet_address = $1 AND blockchain = $2 + `, + walletAuth.PublicKey, + walletAuth.Blockchain, + ) + if queryErr != nil { + err = queryErr + return + } + server.WithPgResult(result, queryErr, func() { + if result.Next() { + server.Raise(result.Scan(&conflictUserId)) + } + }) + if conflictUserId != nil && *conflictUserId != addWalletAuth.UserId { + err = errors.New("This wallet is already linked to another account.") + return + } + _, dbErr := tx.Exec( ctx, ` @@ -690,6 +795,43 @@ func addWalletAuth( ) err = dbErr + return + } + + // Mirror onto network_user's top-level wallet columns, symmetric + // with RemoveAuth's solana branch which clears them. Without this, + // a remove-then-re-add cycle leaves wallet_address/wallet_blockchain + // nil even though network_user_auth_wallet has the wallet bound. + // + // network_user.wallet_address has a global (not per-user) unique + // index, so this can legitimately conflict -- e.g. a legacy account + // whose network_user.wallet_address was set before + // network_user_auth_wallet existed and was never cleared. Handle + // that the same way the INSERT above does (return a clean error) + // instead of panicking, which would otherwise surface as an + // uncaught 500 after server.Tx's transient-error retry loop burns + // up to a minute on a permanent constraint violation. + _, dbErr = tx.Exec( + ctx, + `UPDATE network_user + SET wallet_address = $2, wallet_blockchain = $3 + WHERE user_id = $1`, + addWalletAuth.UserId, + walletAuth.PublicKey, + walletAuth.Blockchain, + ) + if dbErr != nil { + + glog.Infof( + "Error mirroring wallet auth onto network_user: %s user_id=%s wallet_address=%s blockchain=%s", + dbErr.Error(), + addWalletAuth.UserId, + walletAuth.PublicKey, + walletAuth.Blockchain, + ) + + err = dbErr + return } }) @@ -736,6 +878,44 @@ func getWalletAuths( return walletAuths, nil } +func getSeedphraseAuths( + ctx context.Context, + userId server.Id, +) ([]NetworkUserSeedphraseAuth, error) { + + var seedphraseAuths []NetworkUserSeedphraseAuth + + server.Tx(ctx, func(tx server.PgTx) { + + result, err := tx.Query( + ctx, + ` + SELECT + create_time + FROM network_user_auth_seedphrase + WHERE user_id = $1 + `, + userId, + ) + if err != nil { + server.Raise(err) + } + + server.WithPgResult(result, err, func() { + for result.Next() { + sa := NetworkUserSeedphraseAuth{} + server.Raise(result.Scan( + &sa.CreateTime, + )) + seedphraseAuths = append(seedphraseAuths, sa) + } + }) + + }) + + return seedphraseAuths, nil +} + func getWalletAuthsByAddress( ctx context.Context, walletAddress string, @@ -1165,3 +1345,158 @@ func MigrateNetworkUserChildAuths( }) } + +// HasAnyAuthMethod reports whether userId has at least one bound +// password/phone, SSO, wallet, or seedphrase auth method. +// +// This exists because neither the JWT's GuestMode claim nor +// network_user.auth_type reliably reflect whether a legacy guest account +// (auth_type = 'guest', now silently treated as a normal account since guest +// signup was retired) has since added a real auth method via AddAuth: +// RefreshToken re-signs guest JWTs with GuestMode=false unconditionally, and +// AddAuth never updates network_user.auth_type when adding a guest's first +// real method. A live count across the auth tables is the only accurate +// signal, and it self-heals the moment a guest adds a method -- no refresh +// or auth_type backfill required. +func HasAnyAuthMethod(ctx context.Context, userId server.Id) bool { + var count int + server.Db(ctx, func(conn server.PgConn) { + result, err := conn.Query( + ctx, + `SELECT COUNT(*) FROM ( + SELECT 1 FROM network_user_auth_password WHERE user_id = $1 + UNION ALL + SELECT 1 FROM network_user_auth_sso WHERE user_id = $1 + UNION ALL + SELECT 1 FROM network_user_auth_wallet WHERE user_id = $1 + UNION ALL + SELECT 1 FROM network_user_auth_seedphrase WHERE user_id = $1 + ) AS auth_counts`, + userId, + ) + server.WithPgResult(result, err, func() { + if result.Next() { + server.Raise(result.Scan(&count)) + } + }) + }) + return 0 < count +} + +func RemoveAuth(ctx context.Context, userId server.Id, authType string) error { + if err := CheckAccountActionRateLimit( + ctx, + userId, + AccountActionRemoveAuth, + AccountActionRemoveAuthDailyLimit, + AccountActionDailyWindow, + ); err != nil { + return err + } + + var validationErr error + + server.Tx(ctx, func(tx server.PgTx) { + // Lock the user row for the duration of this tx so a concurrent + // RemoveAuth call blocks (and retries via server.Tx's serialization + // handling) instead of both calls reading the same stale auth-method + // count and both committing, which can leave the account with zero + // remaining auth methods. + server.RaisePgResult(tx.Exec( + ctx, + `SELECT 1 FROM network_user WHERE user_id = $1 FOR UPDATE`, + userId, + )) + + currentCount := 0 + countResult, err := tx.Query( + ctx, + `SELECT COUNT(*) FROM ( + SELECT 1 FROM network_user_auth_password WHERE user_id = $1 + UNION ALL + SELECT 1 FROM network_user_auth_sso WHERE user_id = $1 + UNION ALL + SELECT 1 FROM network_user_auth_wallet WHERE user_id = $1 + UNION ALL + SELECT 1 FROM network_user_auth_seedphrase WHERE user_id = $1 + ) AS auth_counts`, + userId, + ) + server.WithPgResult(countResult, err, func() { + if countResult.Next() { + server.Raise(countResult.Scan(¤tCount)) + } + }) + if currentCount <= 1 { + validationErr = fmt.Errorf("cannot remove your last auth method") + return + } + + switch authType { + case "email", "phone": + server.RaisePgResult(tx.Exec( + ctx, + `DELETE FROM network_user_auth_password + WHERE user_id = $1 AND auth_type = $2`, + userId, authType, + )) + case "apple", "google": + server.RaisePgResult(tx.Exec( + ctx, + `DELETE FROM network_user_auth_sso + WHERE user_id = $1 AND auth_type = $2`, + userId, authType, + )) + case "solana": + server.RaisePgResult(tx.Exec( + ctx, + `DELETE FROM network_user_auth_wallet + WHERE user_id = $1`, + userId, + )) + // Clear wallet fields on network_user so the wallet + // can be reused to create a fresh account + server.RaisePgResult(tx.Exec( + ctx, + `UPDATE network_user + SET wallet_address = NULL, wallet_blockchain = NULL + WHERE user_id = $1`, + userId, + )) + case "seedphrase": + server.RaisePgResult(tx.Exec( + ctx, + `DELETE FROM network_user_auth_seedphrase + WHERE user_id = $1`, + userId, + )) + default: + validationErr = fmt.Errorf("unknown auth type: %s", authType) + return + } + + // Update network_user.auth_type to match what's actually left. + // The app reads this field to show available sign-in methods, + // so it must reflect reality after removal. + server.RaisePgResult(tx.Exec( + ctx, + `UPDATE network_user + SET auth_type = COALESCE( + (SELECT auth_type FROM network_user_auth_password WHERE user_id = $1 LIMIT 1), + (SELECT 'apple' FROM network_user_auth_sso WHERE user_id = $1 AND auth_type = 'apple' LIMIT 1), + (SELECT 'google' FROM network_user_auth_sso WHERE user_id = $1 AND auth_type = 'google' LIMIT 1), + (SELECT 'solana' FROM network_user_auth_wallet WHERE user_id = $1 LIMIT 1), + (SELECT 'seedphrase' FROM network_user_auth_seedphrase WHERE user_id = $1 LIMIT 1), + auth_type + ) + WHERE user_id = $1`, + userId, + )) + }) + + if validationErr == nil { + RecordAccountActionAttempt(ctx, userId, AccountActionRemoveAuth) + } + + return validationErr +} diff --git a/model/network_user_model_test.go b/model/network_user_model_test.go index 48f1589e..93e470c1 100644 --- a/model/network_user_model_test.go +++ b/model/network_user_model_test.go @@ -6,7 +6,7 @@ import ( "strings" "testing" - "github.com/urnetwork/connect" + "github.com/go-playground/assert/v2" "github.com/urnetwork/server" "github.com/urnetwork/server/jwt" "github.com/urnetwork/server/session" @@ -26,33 +26,33 @@ func TestNetworkUser(t *testing.T) { networkUser := GetNetworkUser(ctx, userId) - connect.AssertNotEqual(t, networkUser, nil) - connect.AssertEqual(t, networkUser.UserId, userId) - connect.AssertEqual(t, networkUser.UserAuth, fmt.Sprintf("%s@bringyour.com", networkId)) - connect.AssertEqual(t, networkUser.Verified, true) - connect.AssertEqual(t, networkUser.AuthType, AuthTypePassword) - connect.AssertEqual(t, networkUser.NetworkName, networkName) + assert.NotEqual(t, networkUser, nil) + assert.Equal(t, networkUser.UserId, userId) + assert.Equal(t, networkUser.UserAuth, fmt.Sprintf("%s@bringyour.com", networkId)) + assert.Equal(t, networkUser.Verified, true) + assert.Equal(t, networkUser.AuthType, AuthTypePassword) + assert.Equal(t, networkUser.NetworkName, networkName) // test for invalid user id userId = server.NewId() networkUser = GetNetworkUser(ctx, userId) - connect.AssertEqual(t, networkUser, nil) + assert.Equal(t, networkUser, nil) - // create guest network - guestNetworkId := server.NewId() - guestUserId := server.NewId() - guestNetworkName := "guest_hello_world" + // a second, independent network+user via the same helper + secondNetworkId := server.NewId() + secondUserId := server.NewId() + secondNetworkName := "second_hello_world" - Testing_CreateGuestNetwork(ctx, guestNetworkId, guestNetworkName, guestUserId) + Testing_CreateNetwork(ctx, secondNetworkId, secondNetworkName, secondUserId) - networkUser = GetNetworkUser(ctx, guestUserId) + networkUser = GetNetworkUser(ctx, secondUserId) - connect.AssertNotEqual(t, networkUser, nil) - connect.AssertEqual(t, networkUser.UserId, guestUserId) - connect.AssertEqual(t, networkUser.UserAuth, nil) - connect.AssertEqual(t, networkUser.Verified, false) - connect.AssertEqual(t, networkUser.AuthType, AuthTypeGuest) - connect.AssertEqual(t, networkUser.NetworkName, guestNetworkName) + assert.NotEqual(t, networkUser, nil) + assert.Equal(t, networkUser.UserId, secondUserId) + assert.Equal(t, networkUser.UserAuth, fmt.Sprintf("%s@bringyour.com", secondNetworkId)) + assert.Equal(t, networkUser.Verified, true) + assert.Equal(t, networkUser.AuthType, AuthTypePassword) + assert.Equal(t, networkUser.NetworkName, secondNetworkName) }) } @@ -89,12 +89,12 @@ func TestAddUserAuthPassword(t *testing.T) { }, ctx, ) - connect.AssertNotEqual(t, err, nil) + assert.NotEqual(t, err, nil) networkUser := GetNetworkUser(ctx, userId) - connect.AssertNotEqual(t, networkUser, nil) - connect.AssertEqual(t, len(networkUser.UserAuths), 1) - connect.AssertEqual(t, networkUser.UserAuths[0].AuthType, UserAuthTypeEmail) + assert.NotEqual(t, networkUser, nil) + assert.Equal(t, len(networkUser.UserAuths), 1) + assert.Equal(t, networkUser.UserAuths[0].AuthType, UserAuthTypeEmail) /** * But adding a phone number should work @@ -110,13 +110,13 @@ func TestAddUserAuthPassword(t *testing.T) { }, ctx, ) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) networkUser = GetNetworkUser(ctx, userId) - connect.AssertNotEqual(t, networkUser, nil) - connect.AssertEqual(t, len(networkUser.UserAuths), 2) - connect.AssertEqual(t, strings.Contains(networkUser.UserAuths[1].UserAuth, userAuth), true) // adds "+1" to phone number, so doing a string check - connect.AssertEqual(t, networkUser.UserAuths[1].AuthType, UserAuthTypePhone) + assert.NotEqual(t, networkUser, nil) + assert.Equal(t, len(networkUser.UserAuths), 2) + assert.Equal(t, strings.Contains(networkUser.UserAuths[1].UserAuth, userAuth), true) // adds "+1" to phone number, so doing a string check + assert.Equal(t, networkUser.UserAuths[1].AuthType, UserAuthTypePhone) }) } @@ -154,15 +154,15 @@ func TestAddUserAuthWallet(t *testing.T) { }, session.Ctx, ) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) /** * Make sure it's being populated whe we fetch the network user */ networkUser := GetNetworkUser(ctx, userId) - connect.AssertNotEqual(t, networkUser, nil) - connect.AssertEqual(t, len(networkUser.WalletAuths), 1) - connect.AssertEqual(t, networkUser.WalletAuths[0].WalletAddress, pk) + assert.NotEqual(t, networkUser, nil) + assert.Equal(t, len(networkUser.WalletAuths), 1) + assert.Equal(t, networkUser.WalletAuths[0].WalletAddress, pk) /** * Overwrite the wallet auth with a different public key @@ -183,12 +183,12 @@ func TestAddUserAuthWallet(t *testing.T) { }, ctx, ) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) networkUser = GetNetworkUser(ctx, userId) - connect.AssertNotEqual(t, networkUser, nil) - connect.AssertEqual(t, len(networkUser.WalletAuths), 1) - connect.AssertEqual(t, networkUser.WalletAuths[0].WalletAddress, pk) + assert.NotEqual(t, networkUser, nil) + assert.Equal(t, len(networkUser.WalletAuths), 1) + assert.Equal(t, networkUser.WalletAuths[0].WalletAddress, pk) }) } @@ -204,8 +204,8 @@ func TestFindNetworkIdByEmail(t *testing.T) { userAuth := Testing_CreateNetwork(ctx, networkId, networkName, userId) retrievedNetworkId, err := FindNetworkIdByEmail(ctx, userAuth) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, *retrievedNetworkId, networkId) + assert.Equal(t, err, nil) + assert.Equal(t, *retrievedNetworkId, networkId) /** * Test SSO @@ -227,16 +227,16 @@ func TestFindNetworkIdByEmail(t *testing.T) { ) retrievedNetworkId, err = FindNetworkIdByEmail(ctx, email) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, *retrievedNetworkId, networkId) + assert.Equal(t, err, nil) + assert.Equal(t, *retrievedNetworkId, networkId) /** * Test not found */ retrievedNetworkId, err = FindNetworkIdByEmail(ctx, "unknown@email.com") - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, retrievedNetworkId, nil) + assert.Equal(t, err, nil) + assert.Equal(t, retrievedNetworkId, nil) }) } @@ -263,16 +263,16 @@ func TestFindNetworkIdByWalletAddress(t *testing.T) { ) retrievedNetworkId, err := FindNetworkIdByWalletAddress(ctx, publicKey) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, *retrievedNetworkId, networkId) + assert.Equal(t, err, nil) + assert.Equal(t, *retrievedNetworkId, networkId) /** * Test not found */ retrievedNetworkId, err = FindNetworkIdByWalletAddress(ctx, "unknown_wallet_address") - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, retrievedNetworkId, nil) + assert.Equal(t, err, nil) + assert.Equal(t, retrievedNetworkId, nil) }) } diff --git a/model/peer_model.go b/model/peer_model.go index 1a880186..2ebb89fd 100644 --- a/model/peer_model.go +++ b/model/peer_model.go @@ -1450,14 +1450,6 @@ type NetworkPeersError struct { // network-level non-guest sessions (all peers) and top-level client sessions // (peers excluding self). Reads only the redis peer registry. func GetNetworkPeersForSession(session *session.ClientSession) (*NetworkPeersResult, error) { - if session.ByJwt.GuestMode { - return &NetworkPeersResult{ - Error: &NetworkPeersError{ - Message: "Not allowed.", - }, - }, nil - } - var selfClientId *server.Id if session.ByJwt.ClientId != nil { // only top-level clients have peers diff --git a/model/peer_model_test.go b/model/peer_model_test.go index 2c559fd0..ad765cd9 100644 --- a/model/peer_model_test.go +++ b/model/peer_model_test.go @@ -4,12 +4,12 @@ import ( "context" "fmt" mathrand "math/rand" - "strings" + "slices" "sync" "testing" "time" - "github.com/urnetwork/connect" + "github.com/go-playground/assert/v2" "github.com/urnetwork/server" "github.com/urnetwork/server/jwt" @@ -111,7 +111,7 @@ func TestNetworkPeerLifecycle(t *testing.T) { select { case <-time.After(1 * time.Second): } - connect.AssertEqual(t, len(c.Connected()), 0) + assert.Equal(t, len(c.Connected()), 0) peer1 := &NetworkPeer{ ClientId: clientId1, @@ -130,69 +130,70 @@ func TestNetworkPeerLifecycle(t *testing.T) { AddNetworkPeer(ctx, networkId, peer2, residentId2, ttl) eventId, peers := GetNetworkPeers(ctx, networkId) - connect.AssertEqual(t, eventId, GetNetworkPeerEventId(ctx, networkId)) + assert.Equal(t, eventId, GetNetworkPeerEventId(ctx, networkId)) connected, markers := splitNetworkPeers(peers) - connect.AssertEqual(t, len(connected), 2) - connect.AssertEqual(t, len(markers), 0) - connect.AssertEqual(t, connected[clientId1].Principal, "svc-a") - connect.AssertEqual(t, connected[clientId1].Roles, []string{"role1", "role2"}) - connect.AssertEqual(t, connected[clientId1].ProvideModes, []ProvideMode{ProvideModeNetwork, ProvideModeStream}) - connect.AssertEqual(t, connected[clientId1].DeviceName, "device a") - connect.AssertEqual(t, connected[clientId1].DeviceSpec, "spec a") - connect.AssertEqual(t, connected[clientId2].Principal, "") - connect.AssertEqual(t, len(connected[clientId2].Roles), 0) + assert.Equal(t, len(connected), 2) + assert.Equal(t, len(markers), 0) + assert.Equal(t, connected[clientId1].Principal, "svc-a") + assert.Equal(t, connected[clientId1].Roles, []string{"role1", "role2"}) + assert.Equal(t, connected[clientId1].ProvideModes, []ProvideMode{ProvideModeNetwork, ProvideModeStream}) + assert.Equal(t, connected[clientId1].DeviceName, "device a") + assert.Equal(t, connected[clientId1].DeviceSpec, "spec a") + assert.Equal(t, connected[clientId2].Principal, "") + assert.Equal(t, len(connected[clientId2].Roles), 0) // the listener accumulates to the same state select { case <-time.After(1 * time.Second): } - connect.AssertEqual(t, c.Connected(), connected) - connect.AssertEqual(t, len(c.Markers()), 0) + assert.Equal(t, c.Connected(), connected) + assert.Equal(t, len(c.Markers()), 0) // refresh is resident-guarded - connect.AssertEqual(t, RefreshNetworkPeer(ctx, networkId, clientId1, residentId1, ttl), true) - connect.AssertEqual(t, RefreshNetworkPeer(ctx, networkId, clientId1, residentId2, ttl), false) - connect.AssertEqual(t, RefreshNetworkPeer(ctx, networkId, server.NewId(), residentId1, ttl), false) + assert.Equal(t, RefreshNetworkPeer(ctx, networkId, clientId1, residentId1, ttl), true) + assert.Equal(t, RefreshNetworkPeer(ctx, networkId, clientId1, residentId2, ttl), false) + assert.Equal(t, RefreshNetworkPeer(ctx, networkId, server.NewId(), residentId1, ttl), false) // remove is resident-guarded RemoveNetworkPeer(ctx, networkId, clientId1, residentId2) _, peers = GetNetworkPeers(ctx, networkId) connected, _ = splitNetworkPeers(peers) - connect.AssertEqual(t, len(connected), 2) + assert.Equal(t, len(connected), 2) RemoveNetworkPeer(ctx, networkId, clientId1, residentId1) _, peers = GetNetworkPeers(ctx, networkId) connected, markers = splitNetworkPeers(peers) - connect.AssertEqual(t, len(connected), 1) - connect.AssertEqual(t, len(markers), 1) - connect.AssertNotEqual(t, markers[clientId1].DisconnectTime, nil) + assert.Equal(t, len(connected), 1) + assert.Equal(t, len(markers), 1) + assert.NotEqual(t, markers[clientId1].DisconnectTime, nil) select { case <-time.After(1 * time.Second): } - connect.AssertEqual(t, len(c.Connected()), 1) - connect.AssertEqual(t, len(c.Markers()), 1) + assert.Equal(t, len(c.Connected()), 1) + assert.Equal(t, len(c.Markers()), 1) // a reconnect clears the marker AddNetworkPeer(ctx, networkId, peer1, residentId1, ttl) _, peers = GetNetworkPeers(ctx, networkId) connected, markers = splitNetworkPeers(peers) - connect.AssertEqual(t, len(connected), 2) - connect.AssertEqual(t, len(markers), 0) + assert.Equal(t, len(connected), 2) + assert.Equal(t, len(markers), 0) select { case <-time.After(1 * time.Second): } - connect.AssertEqual(t, c.Connected(), connected) - connect.AssertEqual(t, len(c.Markers()), 0) + assert.Equal(t, c.Connected(), connected) + assert.Equal(t, len(c.Markers()), 0) - // v2 (PEERS2.md): delivery is poll + full-read — every emitted event - // is a Reset snapshot (the accumulator diffs locally); incremental - // Updated/Removed events are no longer delivered - connect.AssertNotEqual(t, len(c.events), 0) + // event ids are monotonic with no gaps, so the listener never resets + // after the initial subscribe + eventTypes := []NetworkPeerEventType{} for _, event := range c.events { - connect.AssertEqual(t, event.NetworkPeerEventType, NetworkPeerEventTypeReset) + eventTypes = append(eventTypes, event.NetworkPeerEventType) } + assert.Equal(t, eventTypes[0], NetworkPeerEventTypeReset) + assert.Equal(t, slices.Contains(eventTypes[1:], NetworkPeerEventTypeReset), false) // a new listener syncs to the head state with a reset c2 := newTestNetworkPeerAccumulator() @@ -202,7 +203,7 @@ func TestNetworkPeerLifecycle(t *testing.T) { select { case <-time.After(1 * time.Second): } - connect.AssertEqual(t, c2.Connected(), connected) + assert.Equal(t, c2.Connected(), connected) }) } @@ -234,9 +235,9 @@ func TestNetworkPeerExpiry(t *testing.T) { // expired but not yet pruned entries read as disconnect markers _, peers := GetNetworkPeers(ctx, networkId) connected, markers := splitNetworkPeers(peers) - connect.AssertEqual(t, len(connected), 0) - connect.AssertEqual(t, len(markers), 1) - connect.AssertNotEqual(t, markers[clientId1].DisconnectTime, nil) + assert.Equal(t, len(connected), 0) + assert.Equal(t, len(markers), 1) + assert.NotEqual(t, markers[clientId1].DisconnectTime, nil) // another peer's activity prunes the expired entry and publishes // the disconnect marker @@ -249,11 +250,11 @@ func TestNetworkPeerExpiry(t *testing.T) { select { case <-time.After(1 * time.Second): } - connect.AssertEqual(t, len(c.Connected()), 1) - connect.AssertEqual(t, len(c.Markers()), 1) + assert.Equal(t, len(c.Connected()), 1) + assert.Equal(t, len(c.Markers()), 1) // the pruned registration is gone, so refresh reports not registered - connect.AssertEqual(t, RefreshNetworkPeer(ctx, networkId, clientId1, residentId1, 60*time.Second), false) + assert.Equal(t, RefreshNetworkPeer(ctx, networkId, clientId1, residentId1, 60*time.Second), false) }) } @@ -279,12 +280,12 @@ func TestNetworkPeerProvideModesUpdate(t *testing.T) { }, userSession, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, authClientResult.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, authClientResult.Error, nil) clientId = *authClientResult.ClientId _, topLevel, _, profile, _ := GetNetworkPeerProfile(ctx, clientId) - connect.AssertEqual(t, topLevel, true) + assert.Equal(t, topLevel, true) AddNetworkPeer(ctx, networkId, profile, residentId, 60*time.Second) c := newTestNetworkPeerAccumulator() @@ -301,8 +302,8 @@ func TestNetworkPeerProvideModesUpdate(t *testing.T) { case <-time.After(1 * time.Second): } connected := c.Connected() - connect.AssertEqual(t, len(connected), 1) - connect.AssertEqual(t, connected[clientId].ProvideModes, []ProvideMode{ProvideModeNetwork, ProvideModeStream}) + assert.Equal(t, len(connected), 1) + assert.Equal(t, connected[clientId].ProvideModes, []ProvideMode{ProvideModeNetwork, ProvideModeStream}) // no change publishes no event eventCount := c.EventCount() @@ -313,7 +314,7 @@ func TestNetworkPeerProvideModesUpdate(t *testing.T) { select { case <-time.After(1 * time.Second): } - connect.AssertEqual(t, c.EventCount(), eventCount) + assert.Equal(t, c.EventCount(), eventCount) // removing provide keys publishes the reduced modes SetProvide(ctx, clientId, map[ProvideMode][]byte{ @@ -323,7 +324,7 @@ func TestNetworkPeerProvideModesUpdate(t *testing.T) { case <-time.After(1 * time.Second): } connected = c.Connected() - connect.AssertEqual(t, connected[clientId].ProvideModes, []ProvideMode{ProvideModeStream}) + assert.Equal(t, connected[clientId].ProvideModes, []ProvideMode{ProvideModeStream}) }) } @@ -351,32 +352,32 @@ func TestNetworkPeerProfile(t *testing.T) { }, userSession, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, authClientResult.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, authClientResult.Error, nil) clientId := *authClientResult.ClientId profileNetworkId, topLevel, category, profile, peersEnabled := GetNetworkPeerProfile(ctx, clientId) - connect.AssertEqual(t, profileNetworkId, networkId) - connect.AssertEqual(t, topLevel, true) + assert.Equal(t, profileNetworkId, networkId) + assert.Equal(t, topLevel, true) // a network under the top-level limit is enabled for peers - connect.AssertEqual(t, peersEnabled, true) + assert.Equal(t, peersEnabled, true) // an ordinary client is the client category - connect.AssertEqual(t, category, NetworkPeerCategoryClient) - connect.AssertEqual(t, profile.ClientId, clientId) + assert.Equal(t, category, NetworkPeerCategoryClient) + assert.Equal(t, profile.ClientId, clientId) // roles are deduped and sorted - connect.AssertEqual(t, profile.Roles, []string{"role1", "role2"}) - connect.AssertEqual(t, profile.Principal, "svc-a") - connect.AssertEqual(t, profile.DeviceName, "test device") - connect.AssertEqual(t, profile.DeviceSpec, "test spec") + assert.Equal(t, profile.Roles, []string{"role1", "role2"}) + assert.Equal(t, profile.Principal, "svc-a") + assert.Equal(t, profile.DeviceName, "test device") + assert.Equal(t, profile.DeviceSpec, "test spec") // the identity read-through matches identity := GetClientIdentity(ctx, clientId) - connect.AssertEqual(t, identity.Roles, []string{"role1", "role2"}) - connect.AssertEqual(t, identity.Principal, "svc-a") + assert.Equal(t, identity.Roles, []string{"role1", "role2"}) + assert.Equal(t, identity.Principal, "svc-a") // and again from the cache identity = GetClientIdentity(ctx, clientId) - connect.AssertEqual(t, identity.Roles, []string{"role1", "role2"}) - connect.AssertEqual(t, identity.Principal, "svc-a") + assert.Equal(t, identity.Roles, []string{"role1", "role2"}) + assert.Equal(t, identity.Principal, "svc-a") // a derivative client is not top-level sourceClientResult, err := AuthNetworkClient( @@ -386,20 +387,20 @@ func TestNetworkPeerProfile(t *testing.T) { }, userSession, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, sourceClientResult.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, sourceClientResult.Error, nil) _, topLevel, _, profile, peersEnabled = GetNetworkPeerProfile(ctx, *sourceClientResult.ClientId) - connect.AssertEqual(t, topLevel, false) - connect.AssertNotEqual(t, profile, nil) + assert.Equal(t, topLevel, false) + assert.NotEqual(t, profile, nil) // a derivative client never resolves peers enabled - connect.AssertEqual(t, peersEnabled, false) + assert.Equal(t, peersEnabled, false) // a guest session cannot assign roles or principal guestSession := session.Testing_CreateClientSession(ctx, &jwt.ByJwt{ NetworkId: networkId, UserId: userId, - GuestMode: true, + GuestMode: false, }) authClientResult, err = AuthNetworkClient( &AuthNetworkClientArgs{ @@ -408,8 +409,8 @@ func TestNetworkPeerProfile(t *testing.T) { }, guestSession, ) - connect.AssertEqual(t, err, nil) - connect.AssertNotEqual(t, authClientResult.Error, nil) + assert.Equal(t, err, nil) + assert.NotEqual(t, authClientResult.Error, nil) // a client session cannot assign roles or principal deviceId := server.NewId() @@ -426,8 +427,8 @@ func TestNetworkPeerProfile(t *testing.T) { }, clientSession, ) - connect.AssertEqual(t, err, nil) - connect.AssertNotEqual(t, authClientResult.Error, nil) + assert.Equal(t, err, nil) + assert.NotEqual(t, authClientResult.Error, nil) // a session with roles and principal (e.g. from an auth code) passes // them to clients it creates @@ -443,12 +444,12 @@ func TestNetworkPeerProfile(t *testing.T) { }, serviceSession, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, authClientResult.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, authClientResult.Error, nil) _, _, _, profile, _ = GetNetworkPeerProfile(ctx, *authClientResult.ClientId) - connect.AssertEqual(t, profile.Roles, []string{"service-role"}) - connect.AssertEqual(t, profile.Principal, "svc-inherited") + assert.Equal(t, profile.Roles, []string{"service-role"}) + assert.Equal(t, profile.Principal, "svc-inherited") }) } @@ -469,12 +470,12 @@ func TestNetworkProxyPeer(t *testing.T) { // an ordinary client clientResult, err := AuthNetworkClient(&AuthNetworkClientArgs{Description: "client"}, userSession) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) clientId := *clientResult.ClientId // a proxy client: a top-level client with a proxy_device_config row proxyResult, err := AuthNetworkClient(&AuthNetworkClientArgs{Description: "proxy"}, userSession) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) proxyClientId := *proxyResult.ClientId proxyInstanceId := server.NewId() server.Tx(ctx, func(tx server.PgTx) { @@ -492,12 +493,12 @@ func TestNetworkProxyPeer(t *testing.T) { // the profile detects the proxy category _, topLevel, category, profile, peersEnabled := GetNetworkPeerProfile(ctx, proxyClientId) - connect.AssertEqual(t, topLevel, true) - connect.AssertEqual(t, category, NetworkPeerCategoryProxy) - connect.AssertNotEqual(t, profile, nil) - connect.AssertEqual(t, peersEnabled, true) + assert.Equal(t, topLevel, true) + assert.Equal(t, category, NetworkPeerCategoryProxy) + assert.NotEqual(t, profile, nil) + assert.Equal(t, peersEnabled, true) _, _, clientCategory, _, _ := GetNetworkPeerProfile(ctx, clientId) - connect.AssertEqual(t, clientCategory, NetworkPeerCategoryClient) + assert.Equal(t, clientCategory, NetworkPeerCategoryClient) residentId := server.NewId() ttl := 60 * time.Second @@ -517,28 +518,28 @@ func TestNetworkProxyPeer(t *testing.T) { // the peer list contains only the client _, peers := GetNetworkPeers(ctx, networkId) connected, _ := splitNetworkPeers(peers) - connect.AssertEqual(t, len(connected), 1) - connect.AssertNotEqual(t, connected[clientId], nil) - connect.AssertEqual(t, connected[proxyClientId], nil) + assert.Equal(t, len(connected), 1) + assert.NotEqual(t, connected[clientId], nil) + assert.Equal(t, connected[proxyClientId], nil) // the listener only saw the client peer listenerConnected := c.Connected() - connect.AssertEqual(t, len(listenerConnected), 1) - connect.AssertNotEqual(t, listenerConnected[clientId], nil) - connect.AssertEqual(t, listenerConnected[proxyClientId], nil) + assert.Equal(t, len(listenerConnected), 1) + assert.NotEqual(t, listenerConnected[clientId], nil) + assert.Equal(t, listenerConnected[proxyClientId], nil) // the combined count includes both - connect.AssertEqual(t, GetNetworkConnectedCount(ctx, networkId), 2) + assert.Equal(t, GetNetworkConnectedCount(ctx, networkId), 2) // removing the proxy peer drops the count but emits no marker/event eventCount := c.EventCount() RemoveNetworkProxyPeer(ctx, networkId, proxyClientId) - connect.AssertEqual(t, GetNetworkConnectedCount(ctx, networkId), 1) + assert.Equal(t, GetNetworkConnectedCount(ctx, networkId), 1) select { case <-time.After(1 * time.Second): } - connect.AssertEqual(t, c.EventCount(), eventCount) - connect.AssertEqual(t, len(c.Markers()), 0) + assert.Equal(t, c.EventCount(), eventCount) + assert.Equal(t, len(c.Markers()), 0) }) } @@ -562,13 +563,13 @@ func TestNetworkPeerEventGapReset(t *testing.T) { select { case <-time.After(1 * time.Second): } - connect.AssertEqual(t, len(c.Connected()), 1) + assert.Equal(t, len(c.Connected()), 1) // create a delivery gap: advance the event counter without publishing server.Redis(ctx, func(r server.RedisClient) { for range 2 { _, err := r.Incr(ctx, networkPeerEventIdKey(networkId)).Result() - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) } }) @@ -580,9 +581,9 @@ func TestNetworkPeerEventGapReset(t *testing.T) { case <-time.After(2 * time.Second): } connected := c.Connected() - connect.AssertEqual(t, len(connected), 2) - connect.AssertNotEqual(t, connected[clientId1], nil) - connect.AssertNotEqual(t, connected[clientId2], nil) + assert.Equal(t, len(connected), 2) + assert.NotEqual(t, connected[clientId1], nil) + assert.NotEqual(t, connected[clientId2], nil) // the reset event was used to recover resetCount := 0 @@ -596,7 +597,7 @@ func TestNetworkPeerEventGapReset(t *testing.T) { } }() // initial subscribe reset plus the gap recovery reset - connect.AssertEqual(t, 2 <= resetCount, true) + assert.Equal(t, 2 <= resetCount, true) }) } @@ -624,7 +625,7 @@ func TestNetworkPeerRegistryFlushRecovery(t *testing.T) { select { case <-time.After(1 * time.Second): } - connect.AssertEqual(t, len(c.Connected()), 1) + assert.Equal(t, len(c.Connected()), 1) // flush the registry (e.g. a redis loss). The event counter restarts, // so subsequent event ids move backward. @@ -636,7 +637,7 @@ func TestNetworkPeerRegistryFlushRecovery(t *testing.T) { networkPeerDisconnectedKey(networkId), networkPeerEventIdKey(networkId), ).Err() - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) }) // the registration is lost: the heartbeat refresh reports not @@ -645,11 +646,11 @@ func TestNetworkPeerRegistryFlushRecovery(t *testing.T) { // intermediate empty state (missing counter reads as 0, below the // synced value -> resync to empty), then the re-add diverges the // version again. - connect.AssertEqual(t, RefreshNetworkPeer(ctx, networkId, clientId, residentId, ttl), false) + assert.Equal(t, RefreshNetworkPeer(ctx, networkId, clientId, residentId, ttl), false) select { case <-time.After(1 * time.Second): } - connect.AssertEqual(t, len(c.Connected()), 0) + assert.Equal(t, len(c.Connected()), 0) // the resident recovery branch re-adds (a different peer here to show // the resync carries fresh data, not stale accumulator state) @@ -659,16 +660,16 @@ func TestNetworkPeerRegistryFlushRecovery(t *testing.T) { _, peers := GetNetworkPeers(ctx, networkId) connected, _ := splitNetworkPeers(peers) - connect.AssertEqual(t, len(connected), 1) - connect.AssertEqual(t, connected[clientId2].Principal, "svc-b") + assert.Equal(t, len(connected), 1) + assert.Equal(t, connected[clientId2].Principal, "svc-b") // the already-subscribed listener resyncs to the rebuilt state select { case <-time.After(3 * time.Second): } connectedAccumulated := c.Connected() - connect.AssertEqual(t, len(connectedAccumulated), 1) - connect.AssertNotEqual(t, connectedAccumulated[clientId2], nil) + assert.Equal(t, len(connectedAccumulated), 1) + assert.NotEqual(t, connectedAccumulated[clientId2], nil) // a fresh listener converges too c2 := newTestNetworkPeerAccumulator() @@ -678,132 +679,7 @@ func TestNetworkPeerRegistryFlushRecovery(t *testing.T) { select { case <-time.After(1 * time.Second): } - connect.AssertEqual(t, len(c2.Connected()), 1) - }) -} - -// TestNetworkPeerListenerNoConnectionGrowth guards the PEERS2 property whose -// absence caused the 2026-07-15 outage: the poll listeners must NOT open a -// standing connection per listener (v1 held one pubsub subscription each, -// O(clients) connections that melted the cluster). Many listeners share the -// pool; connected_clients stays ~pool-sized, and idle polls (no registry -// change) deliver no events. -func TestNetworkPeerListenerNoConnectionGrowth(t *testing.T) { - server.DefaultTestEnv().Run(t, func(t testing.TB) { - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() - - clientsInfo := func(field string) int { - count := -1 - server.Redis(ctx, func(r server.RedisClient) { - info, err := r.Info(ctx, "clients").Result() - connect.AssertEqual(t, err, nil) - for _, line := range strings.Split(info, "\n") { - if strings.HasPrefix(line, field+":") { - fmt.Sscanf(strings.TrimSpace(strings.TrimPrefix(line, field+":")), "%d", &count) - } - } - }) - return count - } - connectedClients := func() int { return clientsInfo("connected_clients") } - - baseline := connectedClients() - - // stand up many listeners on distinct networks, each with one peer - const listenerCount = 40 - accumulators := make([]*testNetworkPeerAccumulator, listenerCount) - for i := range listenerCount { - networkId := server.NewId() - AddNetworkPeer(ctx, networkId, &NetworkPeer{ClientId: server.NewId()}, server.NewId(), 60*time.Second) - c := newTestNetworkPeerAccumulator() - accumulators[i] = c - listener := NewNetworkPeerListener(ctx, networkId, c.Event, 200*time.Millisecond, 5) - defer listener.Close() - } - - // let every listener poll many times (2s / 200ms = ~10 ticks each) - select { - case <-time.After(2 * time.Second): - } - - // connections did NOT grow ~1 per listener. The pool cap (test config - // max_connections=16) plus a small margin is the ceiling regardless of - // listener count — v1 would have added ~40 subscription connections. - grown := connectedClients() - baseline - if grown >= listenerCount/2 { - t.Fatalf("connected_clients grew by %d for %d listeners — listeners are not sharing the pool (v1 regression?)", grown, listenerCount) - } - - // the defining v2 invariant, literally: ZERO pubsub subscriptions - // exist while listeners run (v1 held one per listener) - if n := clientsInfo("pubsub_clients"); n != 0 { - t.Fatalf("pubsub_clients = %d, want 0 — a subscription-based listener path is back (v1 regression)", n) - } - - // every listener synced its one peer - for i, c := range accumulators { - if len(c.Connected()) != 1 { - t.Fatalf("listener %d saw %d peers, want 1", i, len(c.Connected())) - } - // idle polls (no registry change after the initial sync) do not - // deliver per tick: only the initial reset + the 1/5 insurance - // full-reads (~2-3 over 2s), never one event per poll tick (~10) - if n := c.EventCount(); n > 6 { - t.Fatalf("listener %d delivered %d events for a static registry — polls are over-delivering", i, n) - } - } - }) -} - -// TestNetworkPeerListenerSurvivesRedisError guards the dead-listener fix: a -// redis error inside a poll must be contained to that tick (logged, backed -// off), never kill the listener. 2026-07-15: a panic in the listener run -// goroutine killed it permanently and the client silently stopped receiving -// peer updates. -func TestNetworkPeerListenerSurvivesRedisError(t *testing.T) { - server.DefaultTestEnv().Run(t, func(t testing.TB) { - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() - - networkId := server.NewId() - clientId1 := server.NewId() - - c := newTestNetworkPeerAccumulator() - listener := NewNetworkPeerListener(ctx, networkId, c.Event, 200*time.Millisecond, 5) - defer listener.Close() - - AddNetworkPeer(ctx, networkId, &NetworkPeer{ClientId: clientId1}, server.NewId(), 60*time.Second) - select { - case <-time.After(1 * time.Second): - } - connect.AssertEqual(t, len(c.Connected()), 1) - - // corrupt the version counter to a non-integer: every poll's - // GetNetworkPeerEventId now panics on the Int64 parse - server.Redis(ctx, func(r server.RedisClient) { - connect.AssertEqual(t, r.Set(ctx, networkPeerEventIdKey(networkId), "not-a-number", 0).Err(), nil) - }) - - // the listener rides several failing polls without dying - select { - case <-time.After(2 * time.Second): - } - - // recover: drop the corrupt counter (a legitimate reset). The next - // poll reads a missing (0) counter, mismatches its synced value, and - // resyncs from the intact registry — proving the listener survived the - // error window and still delivers. - clientId2 := server.NewId() - server.Redis(ctx, func(r server.RedisClient) { - connect.AssertEqual(t, r.Del(ctx, networkPeerEventIdKey(networkId)).Err(), nil) - }) - AddNetworkPeer(ctx, networkId, &NetworkPeer{ClientId: clientId2}, server.NewId(), 60*time.Second) - - select { - case <-time.After(3 * time.Second): - } - connect.AssertEqual(t, len(c.Connected()), 2) + assert.Equal(t, len(c2.Connected()), 1) }) } @@ -891,20 +767,20 @@ func TestNetworkPeerChurn(t *testing.T) { registryConnected, registryMarkers := splitNetworkPeers(peers) // the registry truth matches the intended end state - connect.AssertEqual(t, len(registryConnected), peerCount/2) - connect.AssertEqual(t, len(registryMarkers), peerCount-peerCount/2) + assert.Equal(t, len(registryConnected), peerCount/2) + assert.Equal(t, len(registryMarkers), peerCount-peerCount/2) for clientId, endConnected := range expectedConnected { if endConnected { - connect.AssertNotEqual(t, registryConnected[clientId], nil) + assert.NotEqual(t, registryConnected[clientId], nil) } else { - connect.AssertNotEqual(t, registryMarkers[clientId], nil) + assert.NotEqual(t, registryMarkers[clientId], nil) } } // the accumulated listener state converges to the registry truth - connect.AssertEqual(t, c.Connected(), registryConnected) + assert.Equal(t, c.Connected(), registryConnected) for clientId := range registryMarkers { - connect.AssertNotEqual(t, c.Markers()[clientId], nil) + assert.NotEqual(t, c.Markers()[clientId], nil) } // a fresh listener converges to the same state @@ -915,7 +791,7 @@ func TestNetworkPeerChurn(t *testing.T) { select { case <-time.After(1 * time.Second): } - connect.AssertEqual(t, c2.Connected(), registryConnected) + assert.Equal(t, c2.Connected(), registryConnected) }) } @@ -941,8 +817,8 @@ func TestNetworkClientReauthIdentity(t *testing.T) { }, userSession, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, authClientResult.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, authClientResult.Error, nil) clientId := *authClientResult.ClientId // re-auth mints the client's stored identity into the client jwt @@ -953,12 +829,12 @@ func TestNetworkClientReauthIdentity(t *testing.T) { }, userSession, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, reauthResult.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, reauthResult.Error, nil) reauthByJwt, err := jwt.ParseByJwt(ctx, *reauthResult.ByClientJwt) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, reauthByJwt.Roles, []string{"role1", "role2"}) - connect.AssertEqual(t, reauthByJwt.Principal, "svc-a") + assert.Equal(t, err, nil) + assert.Equal(t, reauthByJwt.Roles, []string{"role1", "role2"}) + assert.Equal(t, reauthByJwt.Principal, "svc-a") // a session with its own identity claims does not override the // client's stored identity on re-auth @@ -975,12 +851,12 @@ func TestNetworkClientReauthIdentity(t *testing.T) { }, serviceSession, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, reauthResult.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, reauthResult.Error, nil) reauthByJwt, err = jwt.ParseByJwt(ctx, *reauthResult.ByClientJwt) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, reauthByJwt.Roles, []string{"role1", "role2"}) - connect.AssertEqual(t, reauthByJwt.Principal, "svc-a") + assert.Equal(t, err, nil) + assert.Equal(t, reauthByJwt.Roles, []string{"role1", "role2"}) + assert.Equal(t, reauthByJwt.Principal, "svc-a") // roles and principal are immutable post-create reauthResult, err = AuthNetworkClient( @@ -991,8 +867,8 @@ func TestNetworkClientReauthIdentity(t *testing.T) { }, userSession, ) - connect.AssertEqual(t, err, nil) - connect.AssertNotEqual(t, reauthResult.Error, nil) + assert.Equal(t, err, nil) + assert.NotEqual(t, reauthResult.Error, nil) }) } @@ -1010,13 +886,6 @@ func TestNetworkPeerTopLevelClientLimit(t *testing.T) { UserId: userId, }) - // the top-level client cap is gated by enforce_concurrent_clients (dark by - // default; see pro.yml). Enable it so this test exercises the cap. No - // client is connected here, so the plan concurrent-connected limit stays - // at zero and never interferes. - defer Testing_SetEnforceConcurrentClients(true)() - Testing_ClearNetworkPeersEnabledCache() - var firstClientId server.Id for i := range LimitTopLevelClientIdsPerNetwork { authClientResult, err := AuthNetworkClient( @@ -1025,8 +894,8 @@ func TestNetworkPeerTopLevelClientLimit(t *testing.T) { }, userSession, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, authClientResult.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, authClientResult.Error, nil) if i == 0 { firstClientId = *authClientResult.ClientId } @@ -1039,9 +908,9 @@ func TestNetworkPeerTopLevelClientLimit(t *testing.T) { }, userSession, ) - connect.AssertEqual(t, err, nil) - connect.AssertNotEqual(t, authClientResult.Error, nil) - connect.AssertEqual(t, authClientResult.Error.ClientLimitExceeded, true) + assert.Equal(t, err, nil) + assert.NotEqual(t, authClientResult.Error, nil) + assert.Equal(t, authClientResult.Error.ClientLimitExceeded, true) // derivative clients are not limited by the top-level limit authClientResult, err = AuthNetworkClient( @@ -1051,25 +920,25 @@ func TestNetworkPeerTopLevelClientLimit(t *testing.T) { }, userSession, ) - connect.AssertEqual(t, err, nil) - connect.AssertEqual(t, authClientResult.Error, nil) + assert.Equal(t, err, nil) + assert.Equal(t, authClientResult.Error, nil) // a network at the limit still gets peer subscriptions - connect.AssertEqual(t, NetworkPeersEnabled(ctx, networkId), true) + assert.Equal(t, NetworkPeersEnabled(ctx, networkId), true) // a network over the limit (created before the limit) does not. // The decision is cached per network, so the cached value holds until // the ttl (cleared here) Testing_CreateDevice(ctx, networkId, server.NewId(), server.NewId(), "grandfathered", "grandfathered") - connect.AssertEqual(t, NetworkPeersEnabled(ctx, networkId), true) + assert.Equal(t, NetworkPeersEnabled(ctx, networkId), true) Testing_ClearNetworkPeersEnabledCache() - connect.AssertEqual(t, NetworkPeersEnabled(ctx, networkId), false) + assert.Equal(t, NetworkPeersEnabled(ctx, networkId), false) // the profile resolves the same decision _, topLevel, _, profile, peersEnabled := GetNetworkPeerProfile(ctx, firstClientId) - connect.AssertEqual(t, topLevel, true) - connect.AssertNotEqual(t, profile, nil) - connect.AssertEqual(t, peersEnabled, false) + assert.Equal(t, topLevel, true) + assert.NotEqual(t, profile, nil) + assert.Equal(t, peersEnabled, false) }) } @@ -1112,17 +981,17 @@ func TestNetworkTopLevelClientLimitDisabled(t *testing.T) { &AuthNetworkClientArgs{Description: fmt.Sprintf("provider %d", i)}, userSession, ) - connect.AssertEqual(t, err, nil) + assert.Equal(t, err, nil) if authClientResult.Error != nil { t.Fatalf("provider %d refused while limit disabled: %s (ClientLimitExceeded=%v)", i, authClientResult.Error.Message, authClientResult.Error.ClientLimitExceeded) } clientIds = append(clientIds, *authClientResult.ClientId) } - connect.AssertEqual(t, len(clientIds), overLimit) + assert.Equal(t, len(clientIds), overLimit) // the plan gates report no limit while disabled, at a count over the cap - connect.AssertEqual(t, NetworkConcurrentClientsExceeded(ctx, networkId), false) + assert.Equal(t, NetworkConcurrentClientsExceeded(ctx, networkId), false) // every client may connect — the connection activation gate is dark for i, clientId := range clientIds { if !CanConnectNetworkPeer(ctx, clientId) { @@ -1133,13 +1002,13 @@ func TestNetworkTopLevelClientLimitDisabled(t *testing.T) { // enforce_concurrent_clients: an over-limit network gets NO peer // registrations/subscriptions even while disabled, so its O(size^2) // full-read fan-out never lands on redis (2026-07-17 fix). - connect.AssertEqual(t, NetworkPeersEnabled(ctx, networkId), false) + assert.Equal(t, NetworkPeersEnabled(ctx, networkId), false) // the per-client profile resolves the same peers-off decision while // still reporting the client as a valid top-level client _, topLevel, _, profile, peersEnabled := GetNetworkPeerProfile(ctx, clientIds[0]) - connect.AssertEqual(t, topLevel, true) - connect.AssertNotEqual(t, profile, nil) - connect.AssertEqual(t, peersEnabled, false) + assert.Equal(t, topLevel, true) + assert.NotEqual(t, profile, nil) + assert.Equal(t, peersEnabled, false) }) } @@ -1167,12 +1036,12 @@ func TestNetworkProviderConnectionExemptFromLimit(t *testing.T) { ordinaryId := server.NewId() Testing_CreateDevice(ctx, networkId, server.NewId(), ordinaryId, "ordinary", "ordinary") AddNetworkPeer(ctx, networkId, &NetworkPeer{ClientId: ordinaryId}, server.NewId(), 60*time.Second) - connect.AssertEqual(t, GetNetworkEnforceableConnectedCount(ctx, networkId), 1) + assert.Equal(t, GetNetworkEnforceableConnectedCount(ctx, networkId), 1) // a second ordinary client would exceed the limit -> refused ordinary2Id := server.NewId() Testing_CreateDevice(ctx, networkId, server.NewId(), ordinary2Id, "ordinary2", "ordinary2") - connect.AssertEqual(t, CanConnectNetworkPeer(ctx, ordinary2Id), false) + assert.Equal(t, CanConnectNetworkPeer(ctx, ordinary2Id), false) // a PUBLIC PROVIDER connects regardless: exempt from the count and the // activation gate. SetProvide gives the DB provide modes the gate reads; @@ -1197,6 +1066,6 @@ func TestNetworkProviderConnectionExemptFromLimit(t *testing.T) { // providers did not consume the enforceable count (still 1: the lone // ordinary client) - connect.AssertEqual(t, GetNetworkEnforceableConnectedCount(ctx, networkId), 1) + assert.Equal(t, GetNetworkEnforceableConnectedCount(ctx, networkId), 1) }) } diff --git a/model/seedphrase_auth_model.go b/model/seedphrase_auth_model.go new file mode 100644 index 00000000..661b9b68 --- /dev/null +++ b/model/seedphrase_auth_model.go @@ -0,0 +1,179 @@ +package model + +import ( + "bytes" + "context" + "crypto/sha256" + "errors" + "strings" + + bip39 "github.com/tyler-smith/go-bip39" + + "github.com/urnetwork/server" + "github.com/urnetwork/server/jwt" +) + +type SeedphraseLoginResult struct { + ByJwt string +} + +func normalizeSeedphrase(seedphrase string) string { + words := strings.Fields(strings.ToLower(strings.TrimSpace(seedphrase))) + return strings.Join(words, " ") +} + +func computeSeedphraseLookup(normalized string) []byte { + h := sha256.Sum256([]byte(normalized)) + return h[:] +} + +func CreateSeedphraseAuthInTx( + tx server.PgTx, + ctx context.Context, + userId server.Id, + seedphrase string, +) error { + normalized := normalizeSeedphrase(seedphrase) + lookup := computeSeedphraseLookup(normalized) + salt := createSeedphraseSalt() + hash := computeSeedphraseHash([]byte(normalized), salt) + + _, err := tx.Exec( + ctx, + `INSERT INTO network_user_auth_seedphrase + (user_id, seedphrase_lookup, seedphrase_hash, seedphrase_salt) + VALUES ($1, $2, $3, $4)`, + userId, lookup, hash, salt, + ) + return err +} + +func LoginWithSeedphrase( + ctx context.Context, + seedphrase string, +) (*SeedphraseLoginResult, error) { + normalized := normalizeSeedphrase(seedphrase) + lookup := computeSeedphraseLookup(normalized) + + var userId server.Id + var hash []byte + var salt []byte + + server.Db(ctx, func(conn server.PgConn) { + result, err := conn.Query( + ctx, + `SELECT user_id, seedphrase_hash, seedphrase_salt + FROM network_user_auth_seedphrase + WHERE seedphrase_lookup = $1`, + lookup, + ) + server.WithPgResult(result, err, func() { + if result.Next() { + server.Raise(result.Scan(&userId, &hash, &salt)) + } + }) + }) + + if userId == (server.Id{}) { + return nil, errors.New("unknown seedphrase") + } + + expectedHash := computeSeedphraseHash([]byte(normalized), salt) + if !bytes.Equal(hash, expectedHash) { + return nil, errors.New("unknown seedphrase") + } + + var networkId server.Id + var networkName string + + server.Db(ctx, func(conn server.PgConn) { + result, err := conn.Query( + ctx, + `SELECT n.network_id, n.network_name + FROM network_user nu + INNER JOIN network n ON n.admin_user_id = nu.user_id + WHERE nu.user_id = $1`, + userId, + ) + server.WithPgResult(result, err, func() { + if result.Next() { + server.Raise(result.Scan(&networkId, &networkName)) + } + }) + }) + + if networkId == (server.Id{}) { + // The seedphrase auth row outlived its network_user (e.g. the + // network was deleted but the auth row wasn't cleaned up by some + // other path). jwt.NewByJwt panics on a zero network id, so fail + // cleanly here instead of turning an orphaned-row edge case into + // a 500. + return nil, errors.New("unknown seedphrase") + } + + isPro := IsPro(ctx, &networkId) + + byJwt := jwt.NewByJwt(networkId, userId, networkName, false, isPro) + return &SeedphraseLoginResult{ByJwt: byJwt.Sign()}, nil +} + +func RegenerateSeedphrase(ctx context.Context, userId server.Id) (string, error) { + entropy, err := bip39.NewEntropy(256) + if err != nil { + return "", err + } + newSeedphrase, err := bip39.NewMnemonic(entropy) + if err != nil { + return "", err + } + + normalized := normalizeSeedphrase(newSeedphrase) + lookup := computeSeedphraseLookup(normalized) + salt := createSeedphraseSalt() + hash := computeSeedphraseHash([]byte(normalized), salt) + + server.Tx(ctx, func(tx server.PgTx) { + server.RaisePgResult(tx.Exec( + ctx, + `UPDATE network_user_auth_seedphrase + SET seedphrase_lookup = $2, seedphrase_hash = $3, seedphrase_salt = $4 + WHERE user_id = $1`, + userId, lookup, hash, salt, + )) + }) + + return newSeedphrase, nil +} + +func GenerateSeedphrase(ctx context.Context, userId server.Id) (string, error) { + entropy, err := bip39.NewEntropy(256) + if err != nil { + return "", err + } + seedphrase, err := bip39.NewMnemonic(entropy) + if err != nil { + return "", err + } + + server.Tx(ctx, func(tx server.PgTx) { + err = CreateSeedphraseAuthInTx(tx, ctx, userId, seedphrase) + server.Raise(err) + }) + + return seedphrase, nil +} + +func HasSeedphraseAuth(ctx context.Context, userId server.Id) (bool, error) { + exists := false + server.Db(ctx, func(conn server.PgConn) { + result, err := conn.Query( + ctx, + `SELECT 1 FROM network_user_auth_seedphrase WHERE user_id = $1`, + userId, + ) + server.WithPgResult(result, err, func() { + exists = result.Next() + }) + }) + return exists, nil +} diff --git a/model/seedphrase_auth_model_test.go b/model/seedphrase_auth_model_test.go new file mode 100644 index 00000000..38f3e153 --- /dev/null +++ b/model/seedphrase_auth_model_test.go @@ -0,0 +1,158 @@ +package model + +import ( + "context" + "testing" + + "github.com/go-playground/assert/v2" + "github.com/urnetwork/server" + "github.com/urnetwork/server/jwt" +) + +func TestSeedphraseCreateAndLogin(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + + userId := server.NewId() + networkId := server.NewId() + networkName := "seedphrase-test" + + server.Tx(ctx, func(tx server.PgTx) { + server.RaisePgResult(tx.Exec(ctx, + `INSERT INTO network_user (user_id, user_name, auth_type) + VALUES ($1, $2, $3)`, + userId, networkName, AuthTypeSeedphrase, + )) + server.RaisePgResult(tx.Exec(ctx, + `INSERT INTO network (network_id, network_name, admin_user_id) + VALUES ($1, $2, $3)`, + networkId, networkName, userId, + )) + }) + + // Test 1: Generate seedphrase + seedphrase, err := GenerateSeedphrase(ctx, userId) + assert.Equal(t, err, nil) + assert.NotEqual(t, len(seedphrase), 0) + + // Test 2: Has seedphrase + hasSeedphrase, err := HasSeedphraseAuth(ctx, userId) + assert.Equal(t, err, nil) + assert.Equal(t, hasSeedphrase, true) + + // Test 3: Login with seedphrase + loginResult, err := LoginWithSeedphrase(ctx, seedphrase) + assert.Equal(t, err, nil) + assert.NotEqual(t, loginResult.ByJwt, "") + + // Test 4: Parse the JWT + parsed, err := jwt.ParseByJwt(ctx, loginResult.ByJwt) + assert.Equal(t, err, nil) + assert.Equal(t, parsed.NetworkId, networkId) + assert.Equal(t, parsed.UserId, userId) + + // Test 5: Regenerate + newSeedphrase, err := RegenerateSeedphrase(ctx, userId) + assert.Equal(t, err, nil) + assert.NotEqual(t, newSeedphrase, seedphrase) + + // Test 6: Login with new seedphrase + loginResult2, err := LoginWithSeedphrase(ctx, newSeedphrase) + assert.Equal(t, err, nil) + assert.NotEqual(t, loginResult2.ByJwt, "") + + // Test 7: Login with old seedphrase should fail + _, err = LoginWithSeedphrase(ctx, seedphrase) + assert.NotEqual(t, err, nil) + }) +} + +func TestSeedphraseBadLogin(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + + _, err := LoginWithSeedphrase(ctx, "this is not a valid bip39 mnemonic at all") + assert.NotEqual(t, err, nil) + }) +} + +func TestSeedphraseAlreadyExists(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + + userId := server.NewId() + networkId := server.NewId() + networkName := "seedphrase-exists" + + server.Tx(ctx, func(tx server.PgTx) { + server.RaisePgResult(tx.Exec(ctx, + `INSERT INTO network_user (user_id, user_name, auth_type) + VALUES ($1, $2, $3)`, + userId, networkName, AuthTypeSeedphrase, + )) + server.RaisePgResult(tx.Exec(ctx, + `INSERT INTO network (network_id, network_name, admin_user_id) + VALUES ($1, $2, $3)`, + networkId, networkName, userId, + )) + }) + + _, err := GenerateSeedphrase(ctx, userId) + assert.Equal(t, err, nil) + + // Second generate should fail + _, err = GenerateSeedphrase(ctx, userId) + assert.NotEqual(t, err, nil) + }) +} + +func TestRemoveAuth(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + + userId := server.NewId() + networkId := server.NewId() + networkName := "remove-auth-test" + + server.Tx(ctx, func(tx server.PgTx) { + server.RaisePgResult(tx.Exec(ctx, + `INSERT INTO network_user (user_id, user_name, auth_type) + VALUES ($1, $2, $3)`, + userId, networkName, AuthTypeSeedphrase, + )) + server.RaisePgResult(tx.Exec(ctx, + `INSERT INTO network (network_id, network_name, admin_user_id) + VALUES ($1, $2, $3)`, + networkId, networkName, userId, + )) + }) + + // Add seedphrase auth + _, err := GenerateSeedphrase(ctx, userId) + assert.Equal(t, err, nil) + + // Can't remove last auth method + err = RemoveAuth(ctx, userId, "seedphrase") + assert.NotEqual(t, err, nil) + + // Add a second auth method (email) + email := "test@removeauth.com" + passwordSalt := createPasswordSalt() + passwordHash := computePasswordHashV1([]byte("password"), passwordSalt) + addUserAuth(&AddUserAuthArgs{ + UserId: userId, + UserAuth: &email, + PasswordHash: passwordHash, + PasswordSalt: passwordSalt, + Verified: true, + }, ctx) + + // Now can remove seedphrase + err = RemoveAuth(ctx, userId, "seedphrase") + assert.Equal(t, err, nil) + + // Can't remove last auth method (only email left) + err = RemoveAuth(ctx, userId, "email") + assert.NotEqual(t, err, nil) + }) +} diff --git a/model/util.go b/model/util.go index 5979b193..77dc08d6 100644 --- a/model/util.go +++ b/model/util.go @@ -2,6 +2,8 @@ package model import ( "crypto/rand" + "fmt" + "math/big" "strings" // TODO replace this with fewer works from a larger dictionary @@ -26,3 +28,34 @@ func newCode() (string, error) { func newCodeBase32() (string, error) { return rand.Text(), nil } + +func randomBip39Word() (string, error) { + wordList := bip39.GetWordList() + n, err := rand.Int(rand.Reader, big.NewInt(int64(len(wordList)))) + if err != nil { + return "", err + } + return wordList[n.Int64()], nil +} + +func generateRandomNetworkName() (string, error) { + for { + w1, err := randomBip39Word() + if err != nil { + return "", err + } + w2, err := randomBip39Word() + if err != nil { + return "", err + } + r, err := rand.Int(rand.Reader, big.NewInt(100)) + if err != nil { + return "", err + } + num := r.Int64() + name := fmt.Sprintf("%s-%s-%d", w1, w2, num) + if len(name) >= 5 && len(name) <= 50 { + return name, nil + } + } +} diff --git a/router/handler_utils.go b/router/handler_utils.go index b8e1e589..b17af795 100644 --- a/router/handler_utils.go +++ b/router/handler_utils.go @@ -107,37 +107,6 @@ func implName[R any](impl ImplFunction[R]) string { return name } -// guarantees NetworkId+UserId -// denies guest mode requests -func WrapRequireAuthNoGuest[R any]( - impl ImplFunction[R], - w http.ResponseWriter, - req *http.Request, - formatters ...FormatFunction[R], -) { - wrap( - func(session *session.ClientSession) (R, error) { - if err := session.Auth(req); err != nil { - var empty R - return empty, fmt.Errorf("%d Not authorized.", http.StatusUnauthorized) - } - if session.ByJwt.GuestMode { - var empty R - return empty, fmt.Errorf("%d Not authorized.", http.StatusUnauthorized) - } - r, err := impl(session) - if err != nil { - // wrap the error to tag the impl - err = fmt.Errorf("[%s]%s", implName(impl), err) - } - return r, err - }, - w, - req, - formatters..., - ) -} - // allow guest mode, or authenticated requests func WrapRequireAuth[R any]( impl ImplFunction[R], @@ -279,51 +248,6 @@ func wrapWithInput[T any, R any]( writeJsonResponse(w, result) } -// guarantees NetworkId+UserId -// denies guest mode requests -func WrapWithInputRequireAuthNoGuest[T any, R any]( - impl ImplWithInputFunction[T, R], - w http.ResponseWriter, - req *http.Request, - formatters ...FormatFunction[R], -) { - WrapWithInputBodyFormatterRequireAuthNoGuest( - RequestBodyFormatter, - impl, - w, - req, - formatters..., - ) -} - -func WrapWithInputBodyFormatterRequireAuthNoGuest[T any, R any]( - bodyFormatter BodyFormatFunction, - impl ImplWithInputFunction[T, R], - w http.ResponseWriter, - req *http.Request, - formatters ...FormatFunction[R], -) { - wrapWithInput( - bodyFormatter, - func(arg T, session *session.ClientSession) (R, error) { - if err := session.Auth(req); err != nil { - var empty R - return empty, fmt.Errorf("%d Not authorized.", http.StatusUnauthorized) - } - - if session.ByJwt.GuestMode { - var empty R - return empty, fmt.Errorf("%d Not authorized.", http.StatusUnauthorized) - } - - return impl(arg, session) - }, - w, - req, - formatters..., - ) -} - // guarantees NetworkId+UserId func WrapWithInputRequireAuth[T any, R any]( impl ImplWithInputFunction[T, R], diff --git a/router/router_test.go b/router/router_test.go index 59486b20..31b79792 100644 --- a/router/router_test.go +++ b/router/router_test.go @@ -46,17 +46,6 @@ func TestRouterBasic(t *testing.T) { WrapRequireAuth(impl, w, r) } - AuthNoGuest := func(w http.ResponseWriter, r *http.Request) { - impl := func(clientSession *session.ClientSession) (map[string]any, error) { - if clientSession.ByJwt == nil { - return nil, errors.New("Missing auth.") - } - - return map[string]any{}, nil - } - WrapRequireAuthNoGuest(impl, w, r) - } - Client := func(w http.ResponseWriter, r *http.Request) { impl := func(clientSession *session.ClientSession) (map[string]any, error) { if clientSession.ByJwt == nil { @@ -90,16 +79,6 @@ func TestRouterBasic(t *testing.T) { WrapWithInputRequireAuth(impl, w, r) } - InputAuthNoGuest := func(w http.ResponseWriter, r *http.Request) { - impl := func(input map[string]any, clientSession *session.ClientSession) (map[string]any, error) { - if clientSession.ByJwt == nil { - return nil, errors.New("Missing auth.") - } - return map[string]any{}, nil - } - WrapWithInputRequireAuthNoGuest(impl, w, r) - } - InputClient := func(w http.ResponseWriter, r *http.Request) { impl := func(input map[string]any, clientSession *session.ClientSession) (map[string]any, error) { if clientSession.ByJwt == nil { @@ -116,11 +95,9 @@ func TestRouterBasic(t *testing.T) { routes := []*Route{ NewRoute("GET", "/noauth", NoAuth), NewRoute("GET", "/auth", Auth), - NewRoute("GET", "/noguest", AuthNoGuest), NewRoute("GET", "/client", Client), NewRoute("POST", "/inputnoauth", InputNoAuth), NewRoute("POST", "/inputauth", InputAuth), - NewRoute("POST", "/inputauth-no-guest", InputAuthNoGuest), NewRoute("POST", "/inputclient", InputClient), } @@ -146,17 +123,6 @@ func TestRouterBasic(t *testing.T) { header.Add("Authorization", fmt.Sprintf("Bearer %s", byJwt.Sign())) } - byJwtGuestMode := jwt.NewByJwt( - networkId, - userId, - "test", - true, // guest mode true - false, // pro is false - ) - authGuestMode := func(header http.Header) { - header.Add("Authorization", fmt.Sprintf("Bearer %s", byJwtGuestMode.Sign())) - } - deviceId := server.NewId() clientId := server.NewId() byClientJwt := byJwt.Client(deviceId, clientId) @@ -182,31 +148,6 @@ func TestRouterBasic(t *testing.T) { ) connect.AssertEqual(t, err, nil) - _, err = server.HttpGet( - ctx, - fmt.Sprintf("http://127.0.0.1:%d/noauth", port), - authGuestMode, - server.HttpResponseRequireStatusOk(server.ResponseJsonObject[map[string]any]), - ) - connect.AssertEqual(t, err, nil) - - _, err = server.HttpGet( - ctx, - fmt.Sprintf("http://127.0.0.1:%d/noauth", port), - authGuestMode, - server.HttpResponseRequireStatusOk(server.ResponseJsonObject[map[string]any]), - ) - connect.AssertEqual(t, err, nil) - - // users in guest mode should be restricted to authenticated level routes - _, err = server.HttpGet( - ctx, - fmt.Sprintf("http://127.0.0.1:%d/auth", port), - authGuestMode, - server.HttpResponseRequireStatusOk(server.ResponseJsonObject[map[string]any]), - ) - connect.AssertNotEqual(t, err, nil) - _, err = server.HttpGet( ctx, fmt.Sprintf("http://127.0.0.1:%d/auth", port), @@ -215,31 +156,6 @@ func TestRouterBasic(t *testing.T) { ) connect.AssertNotEqual(t, err, nil) - _, err = server.HttpGet( - ctx, - fmt.Sprintf("http://127.0.0.1:%d/noguest", port), - authGuestMode, - server.HttpResponseRequireStatusOk(server.ResponseJsonObject[map[string]any]), - ) - connect.AssertNotEqual(t, err, nil) - - // authenticated users should be able to access guest level routes - _, err = server.HttpGet( - ctx, - fmt.Sprintf("http://127.0.0.1:%d/noguest", port), - auth, - server.HttpResponseRequireStatusOk(server.ResponseJsonObject[map[string]any]), - ) - connect.AssertEqual(t, err, nil) - - _, err = server.HttpGet( - ctx, - fmt.Sprintf("http://127.0.0.1:%d/noguest", port), - server.NoCustomHeaders, - server.HttpResponseRequireStatusOk(server.ResponseJsonObject[map[string]any]), - ) - connect.AssertNotEqual(t, err, nil) - _, err = server.HttpGet( ctx, fmt.Sprintf("http://127.0.0.1:%d/client", port), @@ -283,16 +199,6 @@ func TestRouterBasic(t *testing.T) { ) connect.AssertEqual(t, err, nil) - // should allow guest requests - _, err = server.HttpPost( - ctx, - fmt.Sprintf("http://127.0.0.1:%d/inputauth", port), - map[string]any{}, - authGuestMode, - server.HttpResponseRequireStatusOk(server.ResponseJsonObject[map[string]any]), - ) - connect.AssertEqual(t, err, nil) - _, err = server.HttpPost( ctx, fmt.Sprintf("http://127.0.0.1:%d/inputauth", port), @@ -302,34 +208,6 @@ func TestRouterBasic(t *testing.T) { ) connect.AssertNotEqual(t, err, nil) - _, err = server.HttpPost( - ctx, - fmt.Sprintf("http://127.0.0.1:%d/inputauth-no-guest", port), - map[string]any{}, - authGuestMode, - server.HttpResponseRequireStatusOk(server.ResponseJsonObject[map[string]any]), - ) - connect.AssertNotEqual(t, err, nil) - - // should deny guest requests - _, err = server.HttpPost( - ctx, - fmt.Sprintf("http://127.0.0.1:%d/inputauth-no-guest", port), - map[string]any{}, - auth, - server.HttpResponseRequireStatusOk(server.ResponseJsonObject[map[string]any]), - ) - connect.AssertEqual(t, err, nil) - - _, err = server.HttpPost( - ctx, - fmt.Sprintf("http://127.0.0.1:%d/inputauth-no-guest", port), - map[string]any{}, - server.NoCustomHeaders, - server.HttpResponseRequireStatusOk(server.ResponseJsonObject[map[string]any]), - ) - connect.AssertNotEqual(t, err, nil) - _, err = server.HttpPost( ctx, fmt.Sprintf("http://127.0.0.1:%d/inputclient", port), @@ -375,15 +253,6 @@ func TestRouterBasic(t *testing.T) { ) connect.AssertEqual(t, err, nil) - // API keys are non-guest, so /noguest should also succeed - _, err = server.HttpGet( - ctx, - fmt.Sprintf("http://127.0.0.1:%d/noguest", port), - authApiKey, - server.HttpResponseRequireStatusOk(server.ResponseJsonObject[map[string]any]), - ) - connect.AssertEqual(t, err, nil) - // API keys have no ClientId, so /client should fail _, err = server.HttpGet( ctx, diff --git a/session/client_session.go b/session/client_session.go index 5489d33b..355cd1cf 100644 --- a/session/client_session.go +++ b/session/client_session.go @@ -99,7 +99,7 @@ func (self *ClientSession) Auth(req *http.Request) error { network.NetworkId, network.UserId, network.NetworkName, - false, // guest mode + false, false, // pro mode - for api keys we don't need to thread this for now ) glog.V(2).Infof("[session]authed via api key as (%s %s)\n", network.NetworkName, network.NetworkId) diff --git a/task/task.go b/task/task.go index e9ad7037..d0e64d00 100644 --- a/task/task.go +++ b/task/task.go @@ -165,13 +165,23 @@ func ScheduleTask[T any, R any]( return } -func ScheduleTaskInTx[T any, R any]( - tx server.PgTx, +type preparedTask struct { + taskId server.Id + functionName string + argsJson []byte + byJwtJson *string + runAt time.Time + runOnceKey *string + priority TaskPriority + maxTimeSeconds int +} + +func prepareTask[T any, R any]( taskFunction TaskFunction[T, R], args T, clientSession *session.ClientSession, opts ...any, -) (taskId server.Id) { +) preparedTask { taskTarget := NewTaskTarget(taskFunction) argsJson, err := json.Marshal(args) @@ -226,11 +236,29 @@ func ScheduleTaskInTx[T any, R any]( runOnceKey_ := runOnce.String() runOnceKey = &runOnceKey_ } - maxTimeSeconds := int(runMaxTime.MaxTime / time.Second) - claimTime := time.Time{} + return preparedTask{ + taskId: server.NewId(), + functionName: taskTarget.TargetFunctionName(), + argsJson: argsJson, + byJwtJson: byJwtJson, + runAt: runAt.At.UTC(), + runOnceKey: runOnceKey, + priority: runPriority.Priority, + maxTimeSeconds: int(runMaxTime.MaxTime / time.Second), + } +} - taskId = server.NewId() +func ScheduleTaskInTx[T any, R any]( + tx server.PgTx, + taskFunction TaskFunction[T, R], + args T, + clientSession *session.ClientSession, + opts ...any, +) (taskId server.Id) { + p := prepareTask(taskFunction, args, clientSession, opts...) + + claimTime := time.Time{} server.RaisePgResult(tx.Exec( clientSession.Ctx, @@ -253,17 +281,94 @@ func ScheduleTaskInTx[T any, R any]( run_priority = LEAST(pending_task.run_priority, $8), run_max_time_seconds = GREATEST(pending_task.run_max_time_seconds, $9) `, - taskId, - taskTarget.TargetFunctionName(), - argsJson, + p.taskId, + p.functionName, + p.argsJson, + clientSession.ClientAddress, + p.byJwtJson, + p.runAt, + p.runOnceKey, + p.priority, + p.maxTimeSeconds, + claimTime, + )) + return p.taskId +} + +// ScheduleTaskInTxIfAbsent is like ScheduleTaskInTx but for callers that need +// an atomic "only schedule if not already pending under this key" guarantee, +// instead of RunOnce's merge-on-conflict semantics. RunOnce's +// `ON CONFLICT (run_once_key) DO UPDATE` only merges run_at/run_priority/ +// run_max_time_seconds into an existing pending row -- crucially not +// args_json -- so if two different calls share a run_once key while the +// first is still pending, scheduling both would silently drop the second +// call's args while still reporting success. This does a single +// `INSERT ... ON CONFLICT (run_once_key) DO NOTHING` and reports via +// `scheduled` whether the row was actually inserted, so the caller can +// reject a duplicate outright -- atomically, in one round trip -- instead of +// a separate check-then-act that can itself race. runOnce is required (not +// optional via opts) since the whole point is a key-scoped guarantee. +func ScheduleTaskInTxIfAbsent[T any, R any]( + tx server.PgTx, + taskFunction TaskFunction[T, R], + args T, + clientSession *session.ClientSession, + runOnce *RunOnceOption, + opts ...any, +) (scheduled bool, taskId server.Id) { + if runOnce == nil { + panic("ScheduleTaskInTxIfAbsent requires a non-nil runOnce key") + } + p := prepareTask(taskFunction, args, clientSession, append(opts, runOnce)...) + + claimTime := time.Time{} + + tag := server.RaisePgResult(tx.Exec( + clientSession.Ctx, + ` + INSERT INTO pending_task ( + task_id, + function_name, + args_json, + client_address, + client_by_jwt_json, + run_at, + run_once_key, + run_priority, + run_max_time_seconds, + claim_time, + release_time + ) VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $10) + ON CONFLICT (run_once_key) DO NOTHING + `, + p.taskId, + p.functionName, + p.argsJson, clientSession.ClientAddress, - byJwtJson, - runAt.At.UTC(), - runOnceKey, - runPriority.Priority, - maxTimeSeconds, + p.byJwtJson, + p.runAt, + p.runOnceKey, + p.priority, + p.maxTimeSeconds, claimTime, )) + scheduled = 0 < tag.RowsAffected() + if !scheduled { + return scheduled, server.Id{} + } + return scheduled, p.taskId +} + +func ScheduleTaskIfAbsent[T any, R any]( + taskFunction TaskFunction[T, R], + args T, + clientSession *session.ClientSession, + runOnce *RunOnceOption, + opts ...any, +) (scheduled bool, taskId server.Id) { + server.Tx(clientSession.Ctx, func(tx server.PgTx) { + scheduled, taskId = ScheduleTaskInTxIfAbsent[T, R](tx, taskFunction, args, clientSession, runOnce, opts...) + }) return } @@ -490,6 +595,41 @@ func ListRescheduledTasks(ctx context.Context) []server.Id { return taskIds } +// CountAvailableByFunctionName reports how many pending_task rows targeting +// the given task function are currently available to run (run_at has +// passed), across all run_once keys -- this excludes rows scheduled for a +// future run_at, which are queued/waiting, not consuming any worker +// capacity yet. Callers use this to enforce a global concurrency cap on a +// specific background task type (e.g. capping how many networks can have a +// bulk operation actually in flight at once), independent of any single +// run_once key. Counting future-scheduled rows here would make a large +// backlog of merely-queued work block admission of brand new requests, even +// though nothing is actually running yet. +func CountAvailableByFunctionName[T any, R any](ctx context.Context, taskFunction TaskFunction[T, R]) int { + functionName := NewTaskTarget(taskFunction).TargetFunctionName() + // computed in Go, not `now()` in SQL: run_at is a naive `timestamp` + // column holding UTC values, and comparing it against `now()` + // (timestamptz) would force a timezone-dependent cast on the session's + // TimeZone setting -- the same class of bug fixed in + // model/account_action_rate_limit.go and avoided in + // model.ReserveBulkClientRemovalSlot. + asOf := server.NowUtc() + count := 0 + server.Db(ctx, func(conn server.PgConn) { + result, err := conn.Query( + ctx, + `SELECT COUNT(*) FROM pending_task WHERE function_name = $1 AND run_at <= $2`, + functionName, asOf, + ) + server.WithPgResult(result, err, func() { + if result.Next() { + server.Raise(result.Scan(&count)) + } + }) + }) + return count +} + func ListClaimedTasks(ctx context.Context) []server.Id { taskIds := []server.Id{} diff --git a/task/task_test.go b/task/task_test.go index e53d16c8..9491b1ca 100644 --- a/task/task_test.go +++ b/task/task_test.go @@ -2,6 +2,7 @@ package task import ( "context" + "encoding/json" "errors" "fmt" "math" @@ -10,6 +11,7 @@ import ( "testing" "time" + "github.com/go-playground/assert/v2" "github.com/urnetwork/connect" // "github.com/urnetwork/server/jwt" @@ -461,3 +463,76 @@ func TestTaskLeaseNotShortenedByKeepalive(t *testing.T) { } }) } + +type Work2Args struct { + Tag string +} + +type Work2Result struct{} + +func Work2( + work2 *Work2Args, + clientSession *session.ClientSession, +) (*Work2Result, error) { + return &Work2Result{}, nil +} + +// ScheduleTaskIfAbsent must atomically insert-or-detect-conflict on the +// run_once key: a first call inserts and reports scheduled == true; a second +// call with the SAME key while the first is still pending (unclaimed) must +// report scheduled == false and must NOT touch the first call's persisted +// args -- unlike plain ScheduleTask+RunOnce, whose ON CONFLICT DO UPDATE +// silently merges only timing/priority into the existing row while leaving +// (and thus never surfacing) the second call's args. +func TestScheduleTaskIfAbsent(t *testing.T) { + server.DefaultTestEnv().Run(t, func(t testing.TB) { + ctx := context.Background() + clientSession := session.Testing_CreateClientSession(ctx, nil) + defer clientSession.Cancel() + + key := RunOnce("test_schedule_task_if_absent", server.NewId()) + + scheduled, firstTaskId := ScheduleTaskIfAbsent( + Work2, + &Work2Args{Tag: "first"}, + clientSession, + key, + ) + assert.Equal(t, scheduled, true) + + // a second call with the same key, while the first is still pending, + // must be rejected -- not merged + scheduledAgain, _ := ScheduleTaskIfAbsent( + Work2, + &Work2Args{Tag: "second"}, + clientSession, + key, + ) + assert.Equal(t, scheduledAgain, false) + + // the persisted task must still be the FIRST call's args; the + // second call's args must never have been written anywhere + tasks := GetTasks(ctx, firstTaskId) + task, ok := tasks[firstTaskId] + if !ok { + t.Fatal("first task not found") + } + var args Work2Args + if err := json.Unmarshal([]byte(task.ArgsJson), &args); err != nil { + t.Fatal(err) + } + assert.Equal(t, args.Tag, "first") + + // once the pending row is gone (simulating the run finishing), the + // same key must be schedulable again + RemovePendingTask(ctx, firstTaskId) + + scheduledAfterClear, _ := ScheduleTaskIfAbsent( + Work2, + &Work2Args{Tag: "third"}, + clientSession, + key, + ) + assert.Equal(t, scheduledAfterClear, true) + }) +} diff --git a/taskworker/taskworker.go b/taskworker/taskworker.go index 3364dce0..4e9f54fa 100644 --- a/taskworker/taskworker.go +++ b/taskworker/taskworker.go @@ -6,6 +6,7 @@ import ( "github.com/urnetwork/server" "github.com/urnetwork/server/controller" + "github.com/urnetwork/server/model" "github.com/urnetwork/server/session" "github.com/urnetwork/server/stats" "github.com/urnetwork/server/task" @@ -52,6 +53,7 @@ func InitTasks(ctx context.Context) { work.ScheduleRemoveExpiredAuthAttempts(clientSession, tx) work.ScheduleRemoveExpiredWalletAuthChallenges(clientSession, tx) work.ScheduleRemoveExpiredWalletNonces(clientSession, tx) + work.ScheduleRemoveExpiredBulkClientRemovalQuota(clientSession, tx) work.ScheduleRemoveOldAuditNetworkEvents(clientSession, tx) work.ScheduleRemoveOldAuditEvents(clientSession, tx) work.ScheduleRemoveOldClientReliabilityStats(clientSession, tx) @@ -169,6 +171,10 @@ func InitTaskWorkerWithSettings(ctx context.Context, settings *task.TaskWorkerSe work.SweepOrphanNetworkClientData, work.SweepOrphanNetworkClientDataPost, ), + task.NewTaskTargetWithPost( + model.RemoveNetworkClientsTask, + model.RemoveNetworkClientsTaskPost, + ), task.NewTaskTargetWithPost( work.SweepOrphanContractData, work.SweepOrphanContractDataPost, @@ -236,6 +242,10 @@ func InitTaskWorkerWithSettings(ctx context.Context, settings *task.TaskWorkerSe work.RemoveExpiredWalletNoncesPost, "github.com/urnetwork/server/taskworker/work.RemoveExpiredWalletNonces", ), + task.NewTaskTargetWithPost( + work.RemoveExpiredBulkClientRemovalQuota, + work.RemoveExpiredBulkClientRemovalQuotaPost, + ), task.NewTaskTargetWithPost( work.RemoveOldAuditNetworkEvents, work.RemoveOldAuditNetworkEventsPost, diff --git a/taskworker/work/bulk_client_removal_quota_work.go b/taskworker/work/bulk_client_removal_quota_work.go new file mode 100644 index 00000000..00e2164f --- /dev/null +++ b/taskworker/work/bulk_client_removal_quota_work.go @@ -0,0 +1,50 @@ +package work + +import ( + "time" + + "github.com/urnetwork/server" + "github.com/urnetwork/server/model" + "github.com/urnetwork/server/session" + "github.com/urnetwork/server/task" +) + +type RemoveExpiredBulkClientRemovalQuotaArgs struct{} + +type RemoveExpiredBulkClientRemovalQuotaResult struct{} + +func ScheduleRemoveExpiredBulkClientRemovalQuota(clientSession *session.ClientSession, tx server.PgTx) { + task.ScheduleTaskInTx( + tx, + RemoveExpiredBulkClientRemovalQuota, + &RemoveExpiredBulkClientRemovalQuotaArgs{}, + clientSession, + task.RunOnce("remove_expired_bulk_client_removal_quota"), + task.RunAt(server.NowUtc().Add(1*time.Hour)), + ) +} + +// RemoveExpiredBulkClientRemovalQuota reaps bulk_client_removal_quota rows +// whose bucket is safely in the past. ReserveBulkClientRemovalSlot only ever +// looks forward from the current hour, so a bucket older than the current +// hour is already irrelevant to any future reservation -- the 2-day margin +// here is just to keep a short audit trail, not because older rows are ever +// consulted again. +func RemoveExpiredBulkClientRemovalQuota( + _ *RemoveExpiredBulkClientRemovalQuotaArgs, + clientSession *session.ClientSession, +) (*RemoveExpiredBulkClientRemovalQuotaResult, error) { + minBucketStart := server.NowUtc().Add(-48 * time.Hour) + model.RemoveExpiredBulkClientRemovalQuota(clientSession.Ctx, minBucketStart) + return &RemoveExpiredBulkClientRemovalQuotaResult{}, nil +} + +func RemoveExpiredBulkClientRemovalQuotaPost( + _ *RemoveExpiredBulkClientRemovalQuotaArgs, + _ *RemoveExpiredBulkClientRemovalQuotaResult, + clientSession *session.ClientSession, + tx server.PgTx, +) error { + ScheduleRemoveExpiredBulkClientRemovalQuota(clientSession, tx) + return nil +}