mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-01 12:31:59 +00:00
fix(security): verify passwords, bind row-security args, fail closed on panic
Verify bcrypt passwords in Direct mode and the shipped procedures, hash on register/reset, ignore client-supplied roles and level at registration and drop the password from the jwt_login payload. Legacy cleartext upgrade is opt-in. Row security templates now bind the user as a parameter, validate identifiers, attach via common.SelectQuery and fail the request if the filter cannot be attached. ApplyColumnSecurity and GetRowSecurityTemplate convert panics to errors and the hooks fail closed. Update audit status.
This commit is contained in:
@@ -2,8 +2,10 @@ package security
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"reflect"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
@@ -145,3 +147,50 @@ func TestSplitTagDropsEmpty(t *testing.T) {
|
||||
t.Fatalf("got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestColumnSecurityPanicFailsClosed(t *testing.T) {
|
||||
type Rec struct {
|
||||
JSONCol string `json:"json_col" bun:"json_col"`
|
||||
}
|
||||
sl, _ := NewSecurityList(&slowProvider{})
|
||||
sl.ColumnSecurity["public.t@1"] = []ColumnSecurity{{
|
||||
Schema: "public", Tablename: "t", Path: []string{"JSONCol"}, Accesstype: "mask", UserID: 1,
|
||||
}}
|
||||
// A struct boxed in an interface is not addressable, so SetString panics.
|
||||
recs := []any{Rec{JSONCol: "secret"}}
|
||||
|
||||
out, err := sl.ApplyColumnSecurity(reflect.ValueOf(recs), reflect.TypeOf(Rec{}), 1, "public", "t")
|
||||
if err == nil {
|
||||
t.Fatalf("panic must be returned as an error, got out=%v", out)
|
||||
}
|
||||
if errors.Is(err, ErrNoColumnSecurity) {
|
||||
t.Fatal("a panic must not look like 'no rules'")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNoRulesIsNotAnError(t *testing.T) {
|
||||
sl, _ := NewSecurityList(&slowProvider{})
|
||||
if _, err := sl.GetRowSecurityTemplate(1, "s", "t"); !errors.Is(err, ErrNoRowSecurity) {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
if _, err := sl.ApplyColumnSecurity(reflect.ValueOf([]int{}), reflect.TypeOf(0), 1, "s", "t"); !errors.Is(err, ErrNoColumnSecurity) {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyColumnSecurityHookFailsClosedOnPanic(t *testing.T) {
|
||||
type Rec struct {
|
||||
JSONCol string `bun:"json_col"`
|
||||
}
|
||||
sl, _ := NewSecurityList(&slowProvider{})
|
||||
sl.ColumnSecurity["public.t@1"] = []ColumnSecurity{{
|
||||
Schema: "public", Tablename: "t", Path: []string{"JSONCol"}, Accesstype: "mask", UserID: 1,
|
||||
}}
|
||||
secCtx := &mockSecurityContext{
|
||||
ctx: context.Background(), userID: 1, hasUser: true, schema: "public", entity: "t",
|
||||
model: &Rec{}, result: []any{Rec{JSONCol: "secret"}},
|
||||
}
|
||||
if err := ApplyColumnSecurity(secCtx, sl); err == nil {
|
||||
t.Fatal("a panic during masking must fail the request, not return unmasked data")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,12 +1,16 @@
|
||||
-- Database Schema for DatabaseAuthenticator
|
||||
-- ============================================
|
||||
|
||||
-- pgcrypto provides gen_random_bytes(), crypt() and gen_salt(); it is required
|
||||
-- for session token generation and for password hashing/verification below.
|
||||
CREATE EXTENSION IF NOT EXISTS pgcrypto;
|
||||
|
||||
-- Users table
|
||||
CREATE TABLE IF NOT EXISTS users (
|
||||
id SERIAL PRIMARY KEY,
|
||||
username VARCHAR(255) NOT NULL UNIQUE,
|
||||
email VARCHAR(255) NOT NULL UNIQUE,
|
||||
password VARCHAR(255), -- bcrypt hashed password (nullable for OAuth2 users)
|
||||
password VARCHAR(255), -- bcrypt hash (nullable for OAuth2 users); legacy cleartext is accepted at login (upgrade to bcrypt is opt-in)
|
||||
user_level INTEGER DEFAULT 0,
|
||||
roles VARCHAR(500), -- Comma-separated roles: "admin,manager,user"
|
||||
is_active BOOLEAN DEFAULT true,
|
||||
@@ -98,6 +102,8 @@ DECLARE
|
||||
v_user_level INTEGER;
|
||||
v_roles TEXT;
|
||||
v_password_hash TEXT;
|
||||
v_supplied_password TEXT;
|
||||
v_password_ok BOOLEAN := false;
|
||||
v_session_token TEXT;
|
||||
v_expires_at TIMESTAMP;
|
||||
v_ip_address TEXT;
|
||||
@@ -107,6 +113,7 @@ DECLARE
|
||||
BEGIN
|
||||
-- Extract login request fields
|
||||
v_username := p_request->>'username';
|
||||
v_supplied_password := p_request->>'password';
|
||||
v_ip_address := p_request->'claims'->>'ip_address';
|
||||
v_user_agent := p_request->'claims'->>'user_agent';
|
||||
|
||||
@@ -121,12 +128,30 @@ BEGIN
|
||||
RETURN;
|
||||
END IF;
|
||||
|
||||
-- TODO: Verify password hash using pgcrypto extension
|
||||
-- Enable pgcrypto: CREATE EXTENSION IF NOT EXISTS pgcrypto;
|
||||
-- IF NOT (crypt(p_request->>'password', v_password_hash) = v_password_hash) THEN
|
||||
-- RETURN QUERY SELECT false, 'Invalid credentials'::text, NULL::jsonb;
|
||||
-- RETURN;
|
||||
-- END IF;
|
||||
-- Verify the password. bcrypt hashes are checked with crypt(); a legacy
|
||||
-- cleartext value is still accepted (and only rewritten as bcrypt if the
|
||||
-- upgrade is explicitly enabled).
|
||||
-- bcrypt only uses the first 72 bytes, so longer input is rejected.
|
||||
IF v_password_hash IS NOT NULL AND v_password_hash <> ''
|
||||
AND v_supplied_password IS NOT NULL AND v_supplied_password <> ''
|
||||
AND octet_length(v_supplied_password) <= 72 THEN
|
||||
IF v_password_hash ~ '^\$2[aby]\$' THEN
|
||||
v_password_ok := (crypt(v_supplied_password, v_password_hash) = v_password_hash);
|
||||
ELSE
|
||||
v_password_ok := (v_password_hash = v_supplied_password);
|
||||
-- Upgrading the stored value is opt-in:
|
||||
-- ALTER DATABASE <db> SET resolvespec.upgrade_password_hash = 'on';
|
||||
IF v_password_ok AND COALESCE(current_setting('resolvespec.upgrade_password_hash', true), 'off') = 'on' THEN
|
||||
UPDATE users SET password = crypt(v_supplied_password, gen_salt('bf')), updated_at = now()
|
||||
WHERE id = v_user_id;
|
||||
END IF;
|
||||
END IF;
|
||||
END IF;
|
||||
|
||||
IF NOT v_password_ok THEN
|
||||
RETURN QUERY SELECT false, 'Invalid credentials'::text, NULL::jsonb;
|
||||
RETURN;
|
||||
END IF;
|
||||
|
||||
-- Generate session token
|
||||
v_session_token := 'sess_' || encode(gen_random_bytes(32), 'hex') || '_' || extract(epoch from now())::bigint::text;
|
||||
@@ -336,6 +361,7 @@ DECLARE
|
||||
v_username TEXT;
|
||||
v_email TEXT;
|
||||
v_password TEXT;
|
||||
v_password_ok BOOLEAN := false;
|
||||
v_user_level INTEGER;
|
||||
v_roles TEXT;
|
||||
BEGIN
|
||||
@@ -350,11 +376,26 @@ BEGIN
|
||||
RETURN;
|
||||
END IF;
|
||||
|
||||
-- TODO: Verify password hash
|
||||
-- IF NOT (crypt(p_password, v_password) = v_password) THEN
|
||||
-- RETURN QUERY SELECT false, 'Invalid credentials'::text, NULL::jsonb;
|
||||
-- RETURN;
|
||||
-- END IF;
|
||||
-- Verify the password (bcrypt, or legacy cleartext).
|
||||
IF v_password IS NOT NULL AND v_password <> ''
|
||||
AND p_password IS NOT NULL AND p_password <> ''
|
||||
AND octet_length(p_password) <= 72 THEN
|
||||
IF v_password ~ '^\$2[aby]\$' THEN
|
||||
v_password_ok := (crypt(p_password, v_password) = v_password);
|
||||
ELSE
|
||||
v_password_ok := (v_password = p_password);
|
||||
-- Upgrading the stored value is opt-in (see resolvespec_login).
|
||||
IF v_password_ok AND COALESCE(current_setting('resolvespec.upgrade_password_hash', true), 'off') = 'on' THEN
|
||||
UPDATE users SET password = crypt(p_password, gen_salt('bf')), updated_at = now()
|
||||
WHERE id = v_user_id;
|
||||
END IF;
|
||||
END IF;
|
||||
END IF;
|
||||
|
||||
IF NOT v_password_ok THEN
|
||||
RETURN QUERY SELECT false, 'Invalid credentials'::text, NULL::jsonb;
|
||||
RETURN;
|
||||
END IF;
|
||||
|
||||
-- Return user data for JWT token generation
|
||||
RETURN QUERY SELECT
|
||||
@@ -364,7 +405,6 @@ BEGIN
|
||||
'id', v_user_id,
|
||||
'username', v_username,
|
||||
'email', v_email,
|
||||
'password', v_password,
|
||||
'user_level', v_user_level,
|
||||
'roles', v_roles
|
||||
);
|
||||
@@ -442,7 +482,8 @@ END;
|
||||
$$ LANGUAGE plpgsql;
|
||||
|
||||
-- 10. resolvespec_register - Registers a new user and creates session
|
||||
-- Input: RegisterRequest as jsonb {username: string, password: string, email: string, user_level: int, roles: array, claims: object, meta: object}
|
||||
-- Input: RegisterRequest as jsonb {username: string, password: string, email: string, claims: object, meta: object}
|
||||
-- (user_level / roles in the request are ignored; new users are unprivileged)
|
||||
-- Output: p_success (bool), p_error (text), p_data (LoginResponse as jsonb)
|
||||
CREATE OR REPLACE FUNCTION resolvespec_register(p_request jsonb)
|
||||
RETURNS TABLE(p_success boolean, p_error text, p_data jsonb) AS $$
|
||||
@@ -465,15 +506,14 @@ BEGIN
|
||||
v_username := p_request->>'username';
|
||||
v_email := p_request->>'email';
|
||||
v_password := p_request->>'password';
|
||||
v_user_level := COALESCE((p_request->>'user_level')::integer, 0);
|
||||
-- Privileges are never taken from the request: self-registration always
|
||||
-- creates an unprivileged user (level 0, no roles, no program user link).
|
||||
v_user_level := 0;
|
||||
v_roles := '';
|
||||
v_ip_address := p_request->'claims'->>'ip_address';
|
||||
v_user_agent := p_request->'claims'->>'user_agent';
|
||||
v_program_user_id := COALESCE((p_request->>'program_user_id')::integer, 0);
|
||||
v_program_user_table := COALESCE(p_request->>'program_user_table', '');
|
||||
|
||||
-- Convert roles array from JSON to comma-separated string
|
||||
SELECT array_to_string(ARRAY(SELECT jsonb_array_elements_text(p_request->'roles')), ',')
|
||||
INTO v_roles;
|
||||
v_program_user_id := 0;
|
||||
v_program_user_table := '';
|
||||
|
||||
-- Validate required fields
|
||||
IF v_username IS NULL OR v_username = '' THEN
|
||||
@@ -491,6 +531,11 @@ BEGIN
|
||||
RETURN;
|
||||
END IF;
|
||||
|
||||
IF octet_length(v_password) > 72 THEN
|
||||
RETURN QUERY SELECT false, 'Password must be at most 72 bytes'::text, NULL::jsonb;
|
||||
RETURN;
|
||||
END IF;
|
||||
|
||||
-- Check if username already exists
|
||||
IF EXISTS (SELECT 1 FROM users WHERE username = v_username) THEN
|
||||
RETURN QUERY SELECT false, 'Username already exists'::text, NULL::jsonb;
|
||||
@@ -503,9 +548,7 @@ BEGIN
|
||||
RETURN;
|
||||
END IF;
|
||||
|
||||
-- TODO: Hash password using pgcrypto extension
|
||||
-- Enable pgcrypto: CREATE EXTENSION IF NOT EXISTS pgcrypto;
|
||||
-- v_password := crypt(v_password, gen_salt('bf'));
|
||||
v_password := crypt(v_password, gen_salt('bf'));
|
||||
|
||||
-- Create new user
|
||||
INSERT INTO users (username, email, password, user_level, roles, is_active, created_at, updated_at, program_user_id, program_user_table)
|
||||
@@ -1520,8 +1563,7 @@ $$ LANGUAGE plpgsql;
|
||||
-- 2. resolvespec_password_reset - Validates the token and updates the user's password
|
||||
-- Input: p_request jsonb {token: string, new_password: string}
|
||||
-- Output: p_success (bool), p_error (text)
|
||||
-- NOTE: Hash the new_password with bcrypt before storing (pgcrypto crypt/gen_salt).
|
||||
-- The TODO below mirrors the convention used in resolvespec_register.
|
||||
-- NOTE: The new password is hashed with bcrypt (pgcrypto crypt/gen_salt) before storing.
|
||||
CREATE OR REPLACE FUNCTION resolvespec_password_reset(p_request jsonb)
|
||||
RETURNS TABLE(p_success boolean, p_error text) AS $$
|
||||
DECLARE
|
||||
@@ -1563,9 +1605,11 @@ BEGIN
|
||||
RETURN;
|
||||
END IF;
|
||||
|
||||
-- TODO: Hash new password with pgcrypto before storing
|
||||
-- Enable pgcrypto: CREATE EXTENSION IF NOT EXISTS pgcrypto;
|
||||
-- v_new_pw := crypt(v_new_pw, gen_salt('bf'));
|
||||
IF octet_length(v_new_pw) > 72 THEN
|
||||
RETURN QUERY SELECT false, 'new_password must be at most 72 bytes'::text;
|
||||
RETURN;
|
||||
END IF;
|
||||
v_new_pw := crypt(v_new_pw, gen_salt('bf'));
|
||||
|
||||
-- Update password and invalidate all sessions
|
||||
UPDATE users SET password = v_new_pw, updated_at = now() WHERE id = v_user_id;
|
||||
|
||||
@@ -7,7 +7,7 @@ CREATE TABLE IF NOT EXISTS users (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
username VARCHAR(255) NOT NULL UNIQUE,
|
||||
email VARCHAR(255) NOT NULL UNIQUE,
|
||||
password VARCHAR(255),
|
||||
password VARCHAR(255), -- bcrypt hash (nullable for OAuth2 users); legacy cleartext is accepted at login (upgrade to bcrypt is opt-in)
|
||||
user_level INTEGER DEFAULT 0,
|
||||
roles VARCHAR(500),
|
||||
is_active BOOLEAN DEFAULT 1,
|
||||
|
||||
@@ -53,11 +53,12 @@ func TestDirectMode_RegisterThenLogin(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
regResp, err := auth.Register(ctx, RegisterRequest{
|
||||
Username: "alice",
|
||||
Password: "hunter2",
|
||||
Email: "alice@example.com",
|
||||
Roles: []string{"user", "admin"},
|
||||
Claims: map[string]any{"ip_address": "127.0.0.1", "user_agent": "test-agent"},
|
||||
Username: "alice",
|
||||
Password: "hunter2",
|
||||
Email: "alice@example.com",
|
||||
Roles: []string{"user", "admin"},
|
||||
UserLevel: 99,
|
||||
Claims: map[string]any{"ip_address": "127.0.0.1", "user_agent": "test-agent"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Register() error = %v", err)
|
||||
@@ -65,8 +66,27 @@ func TestDirectMode_RegisterThenLogin(t *testing.T) {
|
||||
if regResp.Token == "" || regResp.User == nil {
|
||||
t.Fatalf("Register() returned incomplete response: %+v", regResp)
|
||||
}
|
||||
if len(regResp.User.Roles) != 2 {
|
||||
t.Errorf("expected 2 roles, got %v", regResp.User.Roles)
|
||||
if len(regResp.User.Roles) != 0 || regResp.User.UserLevel != 0 {
|
||||
t.Errorf("client-supplied privileges must be ignored, got level=%d roles=%v", regResp.User.UserLevel, regResp.User.Roles)
|
||||
}
|
||||
|
||||
// Password must be stored as a bcrypt hash, not cleartext.
|
||||
var stored string
|
||||
if err := db.QueryRow(`SELECT password FROM users WHERE username = 'alice'`).Scan(&stored); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !isBcryptHash(stored) || stored == "hunter2" {
|
||||
t.Errorf("password not hashed: %q", stored)
|
||||
}
|
||||
|
||||
if _, err := auth.Login(ctx, LoginRequest{Username: "alice", Password: "wrong"}); err == nil {
|
||||
t.Error("login with wrong password must fail")
|
||||
}
|
||||
if _, err := auth.Login(ctx, LoginRequest{Username: "alice"}); err == nil {
|
||||
t.Error("login with empty password must fail")
|
||||
}
|
||||
if _, err := auth.Login(ctx, LoginRequest{Username: "nobody", Password: "hunter2"}); err == nil {
|
||||
t.Error("login for unknown user must fail")
|
||||
}
|
||||
|
||||
loginResp, err := auth.Login(ctx, LoginRequest{Username: "alice", Password: "hunter2"})
|
||||
@@ -465,3 +485,51 @@ func TestDirectMode_OAuthServerClientAndCode(t *testing.T) {
|
||||
t.Error("expected token to be inactive after revoke")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDirectMode_LegacyPlaintextUpgradeIsOptIn(t *testing.T) {
|
||||
for _, enabled := range []bool{false, true} {
|
||||
db := newDirectTestDB(t)
|
||||
auth := NewDatabaseAuthenticatorWithOptions(db, DatabaseAuthenticatorOptions{QueryMode: ModeDirect, UpgradePasswordHash: enabled})
|
||||
ctx := context.Background()
|
||||
|
||||
if _, err := db.Exec(`DELETE FROM users`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := db.Exec(`INSERT INTO users (username, email, password, user_level, roles, is_active, created_at, updated_at) VALUES ('legacy', 'l@example.com', 'oldpass', 0, '', 1, datetime('now'), datetime('now'))`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := auth.Login(ctx, LoginRequest{Username: "legacy", Password: "nope"}); err == nil {
|
||||
t.Fatal("wrong password must fail for legacy row")
|
||||
}
|
||||
if _, err := auth.Login(ctx, LoginRequest{Username: "legacy", Password: "oldpass"}); err != nil {
|
||||
t.Fatalf("legacy login failed (upgrade=%v): %v", enabled, err)
|
||||
}
|
||||
var stored string
|
||||
_ = db.QueryRow(`SELECT password FROM users WHERE username = 'legacy'`).Scan(&stored)
|
||||
if enabled && !isBcryptHash(stored) {
|
||||
t.Fatalf("upgrade enabled but password not upgraded: %q", stored)
|
||||
}
|
||||
if !enabled && stored != "oldpass" {
|
||||
t.Fatalf("upgrade must not happen unless enabled, stored=%q", stored)
|
||||
}
|
||||
if _, err := auth.Login(ctx, LoginRequest{Username: "legacy", Password: "oldpass"}); err != nil {
|
||||
t.Fatalf("second login failed (upgrade=%v): %v", enabled, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestVerifyPasswordEdgeCases(t *testing.T) {
|
||||
h, _ := hashPassword("pw")
|
||||
if ok, _ := verifyPassword(h, "pw"); !ok {
|
||||
t.Error("bcrypt match failed")
|
||||
}
|
||||
if ok, _ := verifyPassword("", "pw"); ok {
|
||||
t.Error("empty stored must not match")
|
||||
}
|
||||
if ok, _ := verifyPassword("pw", ""); ok {
|
||||
t.Error("empty supplied must not match")
|
||||
}
|
||||
if _, err := hashPassword(string(make([]byte, 73))); err == nil {
|
||||
t.Error("73-byte password must be rejected")
|
||||
}
|
||||
}
|
||||
|
||||
+30
-20
@@ -7,6 +7,7 @@ import (
|
||||
"reflect"
|
||||
"strings"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||
)
|
||||
@@ -85,9 +86,13 @@ func applyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error
|
||||
// Get row security template
|
||||
rowSec, err := securityList.GetRowSecurityTemplate(userRef, schema, tablename)
|
||||
if err != nil {
|
||||
// No row security defined, allow query to proceed
|
||||
logger.Debug("No row security for %s.%s@%v: %v", schema, tablename, userRef, err)
|
||||
return nil
|
||||
if errors.Is(err, ErrNoRowSecurity) {
|
||||
// No row security defined for this user/table: nothing to apply.
|
||||
logger.Debug("No row security for %s.%s", schema, tablename)
|
||||
return nil
|
||||
}
|
||||
// Anything else (including a recovered panic) fails closed.
|
||||
return fmt.Errorf("row security failed for %s.%s: %w", schema, tablename, err)
|
||||
}
|
||||
|
||||
// Check if user has a blocking rule
|
||||
@@ -125,21 +130,21 @@ func applyRowSecurity(secCtx SecurityContext, securityList *SecurityList) error
|
||||
}
|
||||
}
|
||||
|
||||
// Generate the WHERE clause from template
|
||||
whereClause := rowSec.GetTemplate(pkName, modelType)
|
||||
|
||||
logger.Info("Applying row security filter for user %v on %s.%s: %s",
|
||||
userRef, schema, tablename, whereClause)
|
||||
|
||||
// Apply the WHERE clause to the query
|
||||
query := secCtx.GetQuery()
|
||||
if selectQuery, ok := query.(interface {
|
||||
Where(string, ...interface{}) interface{}
|
||||
}); ok {
|
||||
secCtx.SetQuery(selectQuery.Where(whereClause))
|
||||
} else {
|
||||
logger.Debug("Query doesn't support Where method, skipping row security")
|
||||
// Generate the WHERE clause and bind arguments from the template
|
||||
whereClause, whereArgs, err := rowSec.GetTemplate(pkName, modelType)
|
||||
if err != nil {
|
||||
return fmt.Errorf("row security failed for %s.%s: %w", schema, tablename, err)
|
||||
}
|
||||
|
||||
logger.Debug("Applying row security filter on %s.%s: %s", schema, tablename, whereClause)
|
||||
|
||||
// A filter that cannot be attached must fail the request; silently
|
||||
// skipping it would expose every row.
|
||||
selectQuery, ok := secCtx.GetQuery().(common.SelectQuery)
|
||||
if !ok {
|
||||
return fmt.Errorf("row security: query type %T on %s.%s does not support Where", secCtx.GetQuery(), schema, tablename)
|
||||
}
|
||||
secCtx.SetQuery(selectQuery.Where(whereClause, whereArgs...))
|
||||
}
|
||||
|
||||
return nil
|
||||
@@ -183,9 +188,14 @@ func applyColumnSecurity(secCtx SecurityContext, securityList *SecurityList) err
|
||||
|
||||
maskedResult, err := securityList.ApplyColumnSecurity(resultValue, modelType, userID, schema, tablename)
|
||||
if err != nil {
|
||||
logger.Warn("Column security error: %v", err)
|
||||
// Don't fail the request, just log the issue
|
||||
return nil
|
||||
if errors.Is(err, ErrNoColumnSecurity) {
|
||||
// No rules for this user/table: nothing to mask.
|
||||
logger.Debug("No column security for %s.%s", schema, tablename)
|
||||
return nil
|
||||
}
|
||||
// Anything else (including a recovered panic) fails closed rather
|
||||
// than returning unmasked data.
|
||||
return fmt.Errorf("column security failed for %s.%s: %w", schema, tablename, err)
|
||||
}
|
||||
|
||||
// Update the result with masked data
|
||||
|
||||
+99
-18
@@ -3,19 +3,23 @@ package security
|
||||
import (
|
||||
"context"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
)
|
||||
|
||||
// Mock SecurityContext for testing hooks
|
||||
type mockSecurityContext struct {
|
||||
ctx context.Context
|
||||
userID int
|
||||
hasUser bool
|
||||
schema string
|
||||
entity string
|
||||
model interface{}
|
||||
query interface{}
|
||||
result interface{}
|
||||
ctx context.Context
|
||||
userID int
|
||||
hasUser bool
|
||||
schema string
|
||||
entity string
|
||||
model interface{}
|
||||
query interface{}
|
||||
result interface{}
|
||||
userRef any
|
||||
}
|
||||
|
||||
func (m *mockSecurityContext) GetContext() context.Context {
|
||||
@@ -27,6 +31,9 @@ func (m *mockSecurityContext) GetUserID() (int, bool) {
|
||||
}
|
||||
|
||||
func (m *mockSecurityContext) GetUserRef() (any, bool) {
|
||||
if m.userRef != nil {
|
||||
return m.userRef, m.hasUser
|
||||
}
|
||||
return m.userID, m.hasUser
|
||||
}
|
||||
|
||||
@@ -194,6 +201,19 @@ func TestLoadSecurityRules(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
// recordingQuery is a common.SelectQuery that records Where calls.
|
||||
type recordingQuery struct {
|
||||
common.SelectQuery
|
||||
clauses []string
|
||||
args [][]any
|
||||
}
|
||||
|
||||
func (q *recordingQuery) Where(query string, args ...interface{}) common.SelectQuery {
|
||||
q.clauses = append(q.clauses, query)
|
||||
q.args = append(q.args, args)
|
||||
return q
|
||||
}
|
||||
|
||||
// Test applyRowSecurity
|
||||
func TestApplyRowSecurity(t *testing.T) {
|
||||
type TestModel struct {
|
||||
@@ -207,6 +227,7 @@ func TestApplyRowSecurity(t *testing.T) {
|
||||
Tablename: "orders",
|
||||
Template: "user_id = {UserID}",
|
||||
HasBlock: false,
|
||||
UserID: 1,
|
||||
},
|
||||
}
|
||||
secList, _ := NewSecurityList(provider)
|
||||
@@ -215,11 +236,7 @@ func TestApplyRowSecurity(t *testing.T) {
|
||||
// Load row security
|
||||
_, _ = secList.LoadRowSecurity(ctx, 1, "public", "orders", false)
|
||||
|
||||
// Mock query that supports Where
|
||||
type MockQuery struct {
|
||||
whereClause string
|
||||
}
|
||||
mockQuery := &MockQuery{}
|
||||
mockQuery := &recordingQuery{}
|
||||
|
||||
secCtx := &mockSecurityContext{
|
||||
ctx: ctx,
|
||||
@@ -236,8 +253,61 @@ func TestApplyRowSecurity(t *testing.T) {
|
||||
t.Fatalf("expected no error, got %v", err)
|
||||
}
|
||||
|
||||
// Note: The actual WHERE clause application requires a query type that implements Where()
|
||||
// In a real scenario, this would be a bun.SelectQuery or similar
|
||||
if len(mockQuery.clauses) != 1 || mockQuery.clauses[0] != "user_id = ?" {
|
||||
t.Fatalf("expected filter to be attached as %q, got %v", "user_id = ?", mockQuery.clauses)
|
||||
}
|
||||
if len(mockQuery.args[0]) != 1 || mockQuery.args[0][0] != 1 {
|
||||
t.Fatalf("expected bound arg [1], got %v", mockQuery.args[0])
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("fails closed when filter cannot be attached", func(t *testing.T) {
|
||||
provider := &mockSecurityProvider{rowSecurity: RowSecurity{
|
||||
Schema: "public", Tablename: "orders", Template: "user_id = {UserID}", UserID: 1,
|
||||
}}
|
||||
secList, _ := NewSecurityList(provider)
|
||||
ctx := context.Background()
|
||||
_, _ = secList.LoadRowSecurity(ctx, 1, "public", "orders", false)
|
||||
|
||||
secCtx := &mockSecurityContext{
|
||||
ctx: ctx, userID: 1, hasUser: true, schema: "public", entity: "orders",
|
||||
model: &TestModel{}, query: struct{}{},
|
||||
}
|
||||
if err := ApplyRowSecurity(secCtx, secList); err == nil {
|
||||
t.Fatal("expected an error when the query does not support Where")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("user context is bound as its id, never rendered into SQL", func(t *testing.T) {
|
||||
uc := &UserContext{UserID: 7, SessionID: "sess_secret", UserName: "x' OR '1'='1"}
|
||||
provider := &mockSecurityProvider{rowSecurity: RowSecurity{
|
||||
Schema: "public", Tablename: "orders", Template: "user_id = {UserID}", UserID: uc,
|
||||
}}
|
||||
secList, _ := NewSecurityList(provider)
|
||||
ctx := context.Background()
|
||||
_, _ = secList.LoadRowSecurity(ctx, uc, "public", "orders", false)
|
||||
|
||||
q := &recordingQuery{}
|
||||
secCtx := &mockSecurityContext{
|
||||
ctx: ctx, userID: 7, hasUser: true, schema: "public", entity: "orders",
|
||||
model: &TestModel{}, query: q, userRef: uc,
|
||||
}
|
||||
if err := ApplyRowSecurity(secCtx, secList); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(q.clauses) != 1 || strings.Contains(q.clauses[0], "sess_secret") || strings.Contains(q.clauses[0], "OR") {
|
||||
t.Fatalf("user data leaked into SQL: %v", q.clauses)
|
||||
}
|
||||
if q.args[0][0] != 7 {
|
||||
t.Fatalf("expected bound user id 7, got %v", q.args[0])
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("invalid identifier is rejected", func(t *testing.T) {
|
||||
rs := RowSecurity{Schema: "public", Tablename: "orders; DROP TABLE x", Template: "{TableName}.uid = 1"}
|
||||
if _, _, err := rs.GetTemplate("id", nil); err == nil {
|
||||
t.Fatal("expected invalid identifier error")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("block access", func(t *testing.T) {
|
||||
@@ -472,6 +542,7 @@ func TestSecurityIntegration(t *testing.T) {
|
||||
Tablename: "orders",
|
||||
Template: "user_id = {UserID}",
|
||||
HasBlock: false,
|
||||
UserID: 1,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -486,6 +557,7 @@ func TestSecurityIntegration(t *testing.T) {
|
||||
schema: "public",
|
||||
entity: "orders",
|
||||
model: &Order{},
|
||||
query: &recordingQuery{},
|
||||
}
|
||||
|
||||
// Step 1: Load security rules
|
||||
@@ -549,6 +621,7 @@ func TestRowSecurityGetTemplateIntegration(t *testing.T) {
|
||||
rowSec RowSecurity
|
||||
pkName string
|
||||
expectedPart string // Part of the expected output
|
||||
expectedArgs []any
|
||||
}{
|
||||
{
|
||||
name: "with all placeholders",
|
||||
@@ -559,7 +632,8 @@ func TestRowSecurityGetTemplateIntegration(t *testing.T) {
|
||||
Template: "{PrimaryKeyName} IN (SELECT {PrimaryKeyName} FROM {SchemaName}.{TableName}_access WHERE user_id = {UserID})",
|
||||
},
|
||||
pkName: "order_id",
|
||||
expectedPart: "order_id IN (SELECT order_id FROM sales.orders_access WHERE user_id = 42)",
|
||||
expectedPart: "order_id IN (SELECT order_id FROM sales.orders_access WHERE user_id = ?)",
|
||||
expectedArgs: []any{42},
|
||||
},
|
||||
{
|
||||
name: "simple user filter",
|
||||
@@ -570,18 +644,25 @@ func TestRowSecurityGetTemplateIntegration(t *testing.T) {
|
||||
Template: "user_id = {UserID}",
|
||||
},
|
||||
pkName: "id",
|
||||
expectedPart: "user_id = 1",
|
||||
expectedPart: "user_id = ?",
|
||||
expectedArgs: []any{1},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
modelType := reflect.TypeOf(Model{})
|
||||
result := tt.rowSec.GetTemplate(tt.pkName, modelType)
|
||||
result, args, err := tt.rowSec.GetTemplate(tt.pkName, modelType)
|
||||
if err != nil {
|
||||
t.Fatalf("GetTemplate() error = %v", err)
|
||||
}
|
||||
|
||||
if result != tt.expectedPart {
|
||||
t.Errorf("GetTemplate() = %q, want %q", result, tt.expectedPart)
|
||||
}
|
||||
if !reflect.DeepEqual(args, tt.expectedArgs) {
|
||||
t.Errorf("GetTemplate() args = %v, want %v", args, tt.expectedArgs)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
package security
|
||||
|
||||
import (
|
||||
"crypto/subtle"
|
||||
"errors"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
// bcrypt only considers the first 72 bytes of input; longer passwords are
|
||||
// rejected rather than silently truncated.
|
||||
const maxPasswordBytes = 72
|
||||
|
||||
var errPasswordTooLong = errors.New("password must be at most 72 bytes")
|
||||
|
||||
func hashPassword(password string) (string, error) {
|
||||
if len(password) > maxPasswordBytes {
|
||||
return "", errPasswordTooLong
|
||||
}
|
||||
h, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(h), nil
|
||||
}
|
||||
|
||||
func isBcryptHash(s string) bool {
|
||||
return strings.HasPrefix(s, "$2a$") || strings.HasPrefix(s, "$2b$") || strings.HasPrefix(s, "$2y$")
|
||||
}
|
||||
|
||||
// verifyPassword checks supplied against the stored value. A stored bcrypt hash
|
||||
// is compared with bcrypt. A legacy cleartext value (written before hashing was
|
||||
// implemented) is compared in constant time and, on a match, needsRehash is true
|
||||
// so the caller can upgrade the row to a bcrypt hash. An empty stored value
|
||||
// (e.g. an OAuth2-only user) never matches.
|
||||
func verifyPassword(stored, supplied string) (ok, needsRehash bool) {
|
||||
if stored == "" || supplied == "" || len(supplied) > maxPasswordBytes {
|
||||
return false, false
|
||||
}
|
||||
if isBcryptHash(stored) {
|
||||
return bcrypt.CompareHashAndPassword([]byte(stored), []byte(supplied)) == nil, false
|
||||
}
|
||||
if subtle.ConstantTimeCompare([]byte(stored), []byte(supplied)) == 1 {
|
||||
return true, true
|
||||
}
|
||||
return false, false
|
||||
}
|
||||
|
||||
var (
|
||||
dummyHashOnce sync.Once
|
||||
dummyHash string
|
||||
)
|
||||
|
||||
// burnPasswordCheck spends roughly one bcrypt comparison so an unknown username
|
||||
// costs about the same as a wrong password.
|
||||
func burnPasswordCheck(supplied string) {
|
||||
dummyHashOnce.Do(func() {
|
||||
h, _ := bcrypt.GenerateFromPassword([]byte("resolvespec-dummy"), bcrypt.DefaultCost)
|
||||
dummyHash = string(h)
|
||||
})
|
||||
if len(supplied) > maxPasswordBytes {
|
||||
supplied = supplied[:maxPasswordBytes]
|
||||
}
|
||||
_ = bcrypt.CompareHashAndPassword([]byte(dummyHash), []byte(supplied))
|
||||
}
|
||||
+92
-13
@@ -2,8 +2,10 @@ package security
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"regexp"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -41,13 +43,76 @@ type RowSecurity struct {
|
||||
UserID any `json:"user_id"`
|
||||
}
|
||||
|
||||
func (m *RowSecurity) GetTemplate(pPrimaryKeyName string, pModelType reflect.Type) string {
|
||||
// safeIdentRe matches an unquoted SQL identifier. Identifiers substituted into a
|
||||
// row-security template must match it; anything else is rejected.
|
||||
var safeIdentRe = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`)
|
||||
|
||||
// ErrNoRowSecurity is returned by GetRowSecurityTemplate when no row security
|
||||
// entry is loaded for the user and table. It means "no rules", as opposed to a
|
||||
// failure, which callers must treat as fatal.
|
||||
var ErrNoRowSecurity = errors.New("no row security data")
|
||||
|
||||
// ErrNoColumnSecurity is the column-security equivalent of ErrNoRowSecurity.
|
||||
var ErrNoColumnSecurity = errors.New("no column security data")
|
||||
|
||||
// userIDScalar reduces the opaque user reference to a scalar that is safe to
|
||||
// bind as a query argument. A *UserContext is reduced to its UserID; other
|
||||
// structured values are rejected rather than stringified into SQL.
|
||||
func userIDScalar(ref any) (any, error) {
|
||||
switch v := ref.(type) {
|
||||
case nil:
|
||||
return nil, fmt.Errorf("row security: no user reference")
|
||||
case *UserContext:
|
||||
if v == nil {
|
||||
return nil, fmt.Errorf("row security: nil user context")
|
||||
}
|
||||
return v.UserID, nil
|
||||
case UserContext:
|
||||
return v.UserID, nil
|
||||
case int, int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64:
|
||||
return v, nil
|
||||
case string:
|
||||
return v, nil
|
||||
default:
|
||||
return nil, fmt.Errorf("row security: unsupported user reference type %T", ref)
|
||||
}
|
||||
}
|
||||
|
||||
// GetTemplate expands the row-security template into a WHERE clause and its
|
||||
// bind arguments. {PrimaryKeyName}, {TableName} and {SchemaName} are validated
|
||||
// identifiers substituted in place; every {UserID} becomes a `?` placeholder
|
||||
// with the user reference bound as an argument, so user data never reaches the
|
||||
// SQL text.
|
||||
func (m *RowSecurity) GetTemplate(pPrimaryKeyName string, pModelType reflect.Type) (string, []any, error) {
|
||||
str := m.Template
|
||||
str = strings.ReplaceAll(str, "{PrimaryKeyName}", pPrimaryKeyName)
|
||||
str = strings.ReplaceAll(str, "{TableName}", m.Tablename)
|
||||
str = strings.ReplaceAll(str, "{SchemaName}", m.Schema)
|
||||
str = strings.ReplaceAll(str, "{UserID}", fmt.Sprintf("%v", m.UserID))
|
||||
return str
|
||||
|
||||
for placeholder, ident := range map[string]string{
|
||||
"{PrimaryKeyName}": pPrimaryKeyName,
|
||||
"{TableName}": m.Tablename,
|
||||
"{SchemaName}": m.Schema,
|
||||
} {
|
||||
if !strings.Contains(str, placeholder) {
|
||||
continue
|
||||
}
|
||||
if !safeIdentRe.MatchString(ident) {
|
||||
return "", nil, fmt.Errorf("row security: invalid identifier %q for %s", ident, placeholder)
|
||||
}
|
||||
str = strings.ReplaceAll(str, placeholder, ident)
|
||||
}
|
||||
|
||||
n := strings.Count(str, "{UserID}")
|
||||
if n == 0 {
|
||||
return str, nil, nil
|
||||
}
|
||||
uid, err := userIDScalar(m.UserID)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
args := make([]any, n)
|
||||
for i := range args {
|
||||
args[i] = uid
|
||||
}
|
||||
return strings.ReplaceAll(str, "{UserID}", "?"), args, nil
|
||||
}
|
||||
|
||||
// SecurityList manages security state and caching
|
||||
@@ -158,7 +223,7 @@ func (m *SecurityList) ColumSecurityApplyOnRecord(prevRecord reflect.Value, newR
|
||||
|
||||
colsecList, ok := m.ColumnSecurity[fmt.Sprintf("%s.%s@%d", pSchema, pTablename, pUserID)]
|
||||
if !ok || colsecList == nil {
|
||||
return cols, fmt.Errorf("no column security data")
|
||||
return cols, ErrNoColumnSecurity
|
||||
}
|
||||
|
||||
for i := range colsecList {
|
||||
@@ -318,8 +383,15 @@ func setColSecValue(fieldsrc reflect.Value, colsec ColumnSecurity, fieldTypeName
|
||||
return 0, fieldsrc
|
||||
}
|
||||
|
||||
func (m *SecurityList) ApplyColumnSecurity(records reflect.Value, modelType reflect.Type, pUserID int, pSchema, pTablename string) (reflect.Value, error) {
|
||||
defer logger.CatchPanic("ApplyColumnSecurity")()
|
||||
func (m *SecurityList) ApplyColumnSecurity(records reflect.Value, modelType reflect.Type, pUserID int, pSchema, pTablename string) (out reflect.Value, err error) {
|
||||
// A panic must surface as an error: recovering into zero results would
|
||||
// read as "success, nothing to mask" and let the response go out unmasked.
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
out = reflect.Value{}
|
||||
err = logger.HandlePanic("ApplyColumnSecurity", r)
|
||||
}
|
||||
}()
|
||||
|
||||
m.ColumnSecurityMutex.RLock()
|
||||
defer m.ColumnSecurityMutex.RUnlock()
|
||||
@@ -330,7 +402,7 @@ func (m *SecurityList) ApplyColumnSecurity(records reflect.Value, modelType refl
|
||||
|
||||
colsecList, ok := m.ColumnSecurity[fmt.Sprintf("%s.%s@%d", pSchema, pTablename, pUserID)]
|
||||
if !ok || colsecList == nil {
|
||||
return records, fmt.Errorf("nocolumn security data")
|
||||
return records, ErrNoColumnSecurity
|
||||
}
|
||||
|
||||
for i := range colsecList {
|
||||
@@ -508,8 +580,15 @@ func (m *SecurityList) LoadRowSecurity(ctx context.Context, pUserRef any, pSchem
|
||||
return record, nil
|
||||
}
|
||||
|
||||
func (m *SecurityList) GetRowSecurityTemplate(pUserRef any, pSchema, pTablename string) (RowSecurity, error) {
|
||||
defer logger.CatchPanic("GetRowSecurityTemplate")()
|
||||
func (m *SecurityList) GetRowSecurityTemplate(pUserRef any, pSchema, pTablename string) (out RowSecurity, err error) {
|
||||
// A panic must surface as an error: recovering into zero results would
|
||||
// read as "no row security" and unblock the user.
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
out = RowSecurity{}
|
||||
err = logger.HandlePanic("GetRowSecurityTemplate", r)
|
||||
}
|
||||
}()
|
||||
|
||||
m.RowSecurityMutex.RLock()
|
||||
defer m.RowSecurityMutex.RUnlock()
|
||||
@@ -520,7 +599,7 @@ func (m *SecurityList) GetRowSecurityTemplate(pUserRef any, pSchema, pTablename
|
||||
|
||||
rowSec, ok := m.RowSecurity[fmt.Sprintf("%s.%s@%v", pSchema, pTablename, pUserRef)]
|
||||
if !ok {
|
||||
return RowSecurity{}, fmt.Errorf("no row security data")
|
||||
return RowSecurity{}, ErrNoRowSecurity
|
||||
}
|
||||
|
||||
return rowSec, nil
|
||||
|
||||
@@ -38,7 +38,8 @@ func (m *mockSecurityProvider) Authenticate(r *http.Request) (*UserContext, erro
|
||||
return m.authUser, m.authError
|
||||
}
|
||||
|
||||
func (m *mockSecurityProvider) SetAuthenticateCallback(_ func(r *http.Request) (*UserContext, error)) {}
|
||||
func (m *mockSecurityProvider) SetAuthenticateCallback(_ func(r *http.Request) (*UserContext, error)) {
|
||||
}
|
||||
|
||||
func (m *mockSecurityProvider) GetColumnSecurity(ctx context.Context, userID int, schema, table string) ([]ColumnSecurity, error) {
|
||||
return m.columnSecurity, nil
|
||||
@@ -78,13 +79,13 @@ func TestNewSecurityList(t *testing.T) {
|
||||
// Test maskString function
|
||||
func TestMaskString(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input string
|
||||
maskStart int
|
||||
maskEnd int
|
||||
maskChar string
|
||||
invert bool
|
||||
expected string
|
||||
name string
|
||||
input string
|
||||
maskStart int
|
||||
maskEnd int
|
||||
maskChar string
|
||||
invert bool
|
||||
expected string
|
||||
}{
|
||||
{
|
||||
name: "mask first 3 characters",
|
||||
@@ -299,12 +300,18 @@ func TestRowSecurityGetTemplate(t *testing.T) {
|
||||
UserID: 42,
|
||||
}
|
||||
|
||||
result := rowSec.GetTemplate("order_id", nil)
|
||||
result, args, err := rowSec.GetTemplate("order_id", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("GetTemplate() error = %v", err)
|
||||
}
|
||||
|
||||
expected := "order_id IN (SELECT order_id FROM public.orders_access WHERE user_id = 42)"
|
||||
expected := "order_id IN (SELECT order_id FROM public.orders_access WHERE user_id = ?)"
|
||||
if result != expected {
|
||||
t.Errorf("GetTemplate() = %q, want %q", result, expected)
|
||||
}
|
||||
if len(args) != 1 || args[0] != 42 {
|
||||
t.Errorf("GetTemplate() args = %v, want [42]", args)
|
||||
}
|
||||
}
|
||||
|
||||
// Test ClearSecurity
|
||||
|
||||
@@ -88,6 +88,10 @@ type DatabaseAuthenticator struct {
|
||||
queryMode QueryMode
|
||||
capability *dbCapability
|
||||
|
||||
// upgradePasswordHash enables rewriting legacy cleartext passwords as bcrypt
|
||||
// on successful login (opt-in, see DatabaseAuthenticatorOptions).
|
||||
upgradePasswordHash bool
|
||||
|
||||
// activityWG tracks in-flight asynchronous session activity updates
|
||||
activityWG sync.WaitGroup
|
||||
|
||||
@@ -130,6 +134,11 @@ type DatabaseAuthenticatorOptions struct {
|
||||
// CookieOptions.Name (default "session_token") in addition to the Authorization header,
|
||||
// and LoginWithCookie / LogoutWithCookie automatically set / clear the cookie.
|
||||
EnableCookieSession bool
|
||||
// UpgradePasswordHash, when true, rewrites a legacy cleartext password as a
|
||||
// bcrypt hash after a successful login. It is off by default and is never
|
||||
// enabled automatically: legacy cleartext values are still accepted at login,
|
||||
// but stored rows are left untouched unless this is set.
|
||||
UpgradePasswordHash bool
|
||||
// CookieOptions configures the session cookie written by LoginWithCookie.
|
||||
// Only used when EnableCookieSession is true.
|
||||
CookieOptions SessionCookieOptions
|
||||
@@ -169,6 +178,7 @@ func NewDatabaseAuthenticatorWithOptions(db *sql.DB, opts DatabaseAuthenticatorO
|
||||
capability: newDBCapability(),
|
||||
passkeyProvider: opts.PasskeyProvider,
|
||||
enableCookieSession: opts.EnableCookieSession,
|
||||
upgradePasswordHash: opts.UpgradePasswordHash,
|
||||
cookieOptions: opts.CookieOptions,
|
||||
authenticateCallback: opts.AuthenticateCallback,
|
||||
}
|
||||
@@ -610,6 +620,17 @@ type JWTAuthenticator struct {
|
||||
tableNames *TableNames
|
||||
queryMode QueryMode
|
||||
capability *dbCapability
|
||||
|
||||
// upgradePasswordHash enables rewriting legacy cleartext passwords as bcrypt
|
||||
// on successful login. Off by default; enable with WithPasswordHashUpgrade.
|
||||
upgradePasswordHash bool
|
||||
}
|
||||
|
||||
// WithPasswordHashUpgrade explicitly enables (or disables) upgrading legacy
|
||||
// cleartext passwords to bcrypt after a successful login. Off by default.
|
||||
func (a *JWTAuthenticator) WithPasswordHashUpgrade(enabled bool) *JWTAuthenticator {
|
||||
a.upgradePasswordHash = enabled
|
||||
return a
|
||||
}
|
||||
|
||||
func NewJWTAuthenticator(secretKey string, db *sql.DB, names ...*SQLNames) *JWTAuthenticator {
|
||||
@@ -698,7 +719,6 @@ func (a *JWTAuthenticator) Login(ctx context.Context, req LoginRequest) (*LoginR
|
||||
ID int `json:"id"`
|
||||
Username string `json:"username"`
|
||||
Email string `json:"email"`
|
||||
Password string `json:"password"`
|
||||
UserLevel int `json:"user_level"`
|
||||
Roles string `json:"roles"`
|
||||
}
|
||||
@@ -707,10 +727,8 @@ func (a *JWTAuthenticator) Login(ctx context.Context, req LoginRequest) (*LoginR
|
||||
return nil, fmt.Errorf("failed to parse user data: %w", err)
|
||||
}
|
||||
|
||||
// TODO: Verify password
|
||||
// if err := bcrypt.CompareHashAndPassword([]byte(user.Password), []byte(req.Password)); err != nil {
|
||||
// return nil, fmt.Errorf("invalid credentials")
|
||||
// }
|
||||
// The password is verified inside resolvespec_jwt_login; the hash is never
|
||||
// returned to Go.
|
||||
|
||||
// Generate token (placeholder - implement JWT signing when library is available)
|
||||
expiresAt := time.Now().Add(24 * time.Hour)
|
||||
|
||||
@@ -10,6 +10,8 @@ import (
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
)
|
||||
|
||||
// Direct-mode implementations for DatabaseAuthenticator and JWTAuthenticator.
|
||||
@@ -17,10 +19,10 @@ import (
|
||||
// parameterized SQL against the configured TableNames, so they work on
|
||||
// SQLite, MySQL, or Postgres without the resolvespec_* functions installed.
|
||||
//
|
||||
// Password verification is intentionally not implemented here: the stored
|
||||
// procedures never verify the password hash either (see the TODOs in
|
||||
// database_schema.sql), so Direct mode matches that behavior exactly rather
|
||||
// than introducing a mismatch between modes.
|
||||
// Passwords are verified with bcrypt (see password.go). Legacy cleartext rows
|
||||
// are still accepted at login; they are only rewritten as bcrypt when the
|
||||
// upgrade is explicitly enabled (UpgradePasswordHash). Registration
|
||||
// never honours client-supplied user_level/roles.
|
||||
|
||||
var (
|
||||
errUsernameExists = errors.New("username already exists")
|
||||
@@ -29,22 +31,34 @@ var (
|
||||
|
||||
func (a *DatabaseAuthenticator) loginDirect(ctx context.Context, req LoginRequest) (*LoginResponse, error) {
|
||||
var userID int
|
||||
var email, roles, programUserTable sql.NullString
|
||||
var email, roles, programUserTable, storedPassword sql.NullString
|
||||
var userLevel, programUserID sql.NullInt64
|
||||
|
||||
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
|
||||
query := rewritePlaceholders(db, fmt.Sprintf(
|
||||
`SELECT id, email, user_level, roles, program_user_id, program_user_table FROM %s WHERE username = ? AND is_active = ?`,
|
||||
`SELECT id, email, user_level, roles, program_user_id, program_user_table, password FROM %s WHERE username = ? AND is_active = ?`,
|
||||
a.tableNames.Users))
|
||||
return db.QueryRowContext(ctx, query, req.Username, true).Scan(&userID, &email, &userLevel, &roles, &programUserID, &programUserTable)
|
||||
return db.QueryRowContext(ctx, query, req.Username, true).Scan(&userID, &email, &userLevel, &roles, &programUserID, &programUserTable, &storedPassword)
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
burnPasswordCheck(req.Password)
|
||||
return nil, fmt.Errorf("invalid credentials")
|
||||
}
|
||||
return nil, fmt.Errorf("login query failed: %w", err)
|
||||
}
|
||||
|
||||
ok, needsRehash := verifyPassword(storedPassword.String, req.Password)
|
||||
if !ok {
|
||||
if storedPassword.String == "" {
|
||||
burnPasswordCheck(req.Password)
|
||||
}
|
||||
return nil, fmt.Errorf("invalid credentials")
|
||||
}
|
||||
if needsRehash && a.upgradePasswordHash {
|
||||
a.upgradePasswordHashFor(ctx, userID, req.Password)
|
||||
}
|
||||
|
||||
sessionToken, err := generateSessionToken()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to generate session token: %w", err)
|
||||
@@ -86,6 +100,24 @@ func (a *DatabaseAuthenticator) loginDirect(ctx context.Context, req LoginReques
|
||||
}, nil
|
||||
}
|
||||
|
||||
// upgradePasswordHashFor replaces a legacy cleartext password with a bcrypt hash.
|
||||
// Only called when the upgrade has been explicitly enabled. Failure is logged
|
||||
// and ignored: the login itself already succeeded.
|
||||
func (a *DatabaseAuthenticator) upgradePasswordHashFor(ctx context.Context, userID int, password string) {
|
||||
h, err := hashPassword(password)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
|
||||
q := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET password = ?, updated_at = ? WHERE id = ?`, a.tableNames.Users))
|
||||
_, err := db.ExecContext(ctx, q, h, time.Now(), userID)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
logger.Warn("failed to upgrade legacy password hash for user %d: %v", userID, err)
|
||||
}
|
||||
}
|
||||
|
||||
func (a *DatabaseAuthenticator) registerDirect(ctx context.Context, req RegisterRequest) (*LoginResponse, error) {
|
||||
if req.Username == "" {
|
||||
return nil, fmt.Errorf("username is required")
|
||||
@@ -97,12 +129,20 @@ func (a *DatabaseAuthenticator) registerDirect(ctx context.Context, req Register
|
||||
return nil, fmt.Errorf("password is required")
|
||||
}
|
||||
|
||||
rolesStr := strings.Join(req.Roles, ",")
|
||||
passwordHash, err := hashPassword(req.Password)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Privileges are never taken from the request: self-registration always
|
||||
// creates an unprivileged user.
|
||||
const userLevel = 0
|
||||
const rolesStr = ""
|
||||
now := time.Now()
|
||||
ipAddress, userAgent := claimStrings(req.Claims)
|
||||
|
||||
var userID int64
|
||||
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
|
||||
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
|
||||
var count int
|
||||
checkQuery := rewritePlaceholders(db, fmt.Sprintf(`SELECT COUNT(*) FROM %s WHERE username = ?`, a.tableNames.Users))
|
||||
if err := db.QueryRowContext(ctx, checkQuery, req.Username).Scan(&count); err != nil {
|
||||
@@ -122,7 +162,7 @@ func (a *DatabaseAuthenticator) registerDirect(ctx context.Context, req Register
|
||||
insertQuery := rewritePlaceholders(db, fmt.Sprintf(
|
||||
`INSERT INTO %s (username, email, password, user_level, roles, is_active, created_at, updated_at, program_user_id, program_user_table) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
|
||||
a.tableNames.Users))
|
||||
res, err := db.ExecContext(ctx, insertQuery, req.Username, req.Email, req.Password, req.UserLevel, rolesStr, true, now, now, 0, "")
|
||||
res, err := db.ExecContext(ctx, insertQuery, req.Username, req.Email, passwordHash, userLevel, rolesStr, true, now, now, 0, "")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -164,7 +204,7 @@ func (a *DatabaseAuthenticator) registerDirect(ctx context.Context, req Register
|
||||
UserID: int(userID),
|
||||
UserName: req.Username,
|
||||
Email: req.Email,
|
||||
UserLevel: req.UserLevel,
|
||||
UserLevel: userLevel,
|
||||
Roles: parseRoles(rolesStr),
|
||||
SessionID: sessionToken,
|
||||
ProgramUserID: 0,
|
||||
@@ -367,12 +407,17 @@ func (a *DatabaseAuthenticator) completePasswordResetDirect(ctx context.Context,
|
||||
return fmt.Errorf("new_password is required")
|
||||
}
|
||||
|
||||
newHash, err := hashPassword(req.NewPassword)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
hash := sha256.Sum256([]byte(req.Token))
|
||||
tokenHash := hex.EncodeToString(hash[:])
|
||||
|
||||
var resetID, userID int
|
||||
var expiresAt time.Time
|
||||
err := a.runDBOpWithReconnect(func(db *sql.DB) error {
|
||||
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
|
||||
query := rewritePlaceholders(db, fmt.Sprintf(`SELECT id, user_id, expires_at FROM %s WHERE token_hash = ? AND used = ?`, a.tableNames.UserPasswordResets))
|
||||
return db.QueryRowContext(ctx, query, tokenHash, false).Scan(&resetID, &userID, &expiresAt)
|
||||
})
|
||||
@@ -389,7 +434,7 @@ func (a *DatabaseAuthenticator) completePasswordResetDirect(ctx context.Context,
|
||||
now := time.Now()
|
||||
err = a.runDBOpWithReconnect(func(db *sql.DB) error {
|
||||
updUser := rewritePlaceholders(db, fmt.Sprintf(`UPDATE %s SET password = ?, updated_at = ? WHERE id = ?`, a.tableNames.Users))
|
||||
if _, err := db.ExecContext(ctx, updUser, req.NewPassword, now, userID); err != nil {
|
||||
if _, err := db.ExecContext(ctx, updUser, newHash, now, userID); err != nil {
|
||||
return err
|
||||
}
|
||||
delSessions := rewritePlaceholders(db, fmt.Sprintf(`DELETE FROM %s WHERE user_id = ?`, a.tableNames.UserSessions))
|
||||
@@ -424,12 +469,12 @@ func claimStrings(claims map[string]any) (ipAddress, userAgent string) {
|
||||
// jwtLoginDirect mirrors resolvespec_jwt_login.
|
||||
func (a *JWTAuthenticator) jwtLoginDirect(ctx context.Context, req LoginRequest) (*LoginResponse, error) {
|
||||
var userID int
|
||||
var email, roles sql.NullString
|
||||
var email, roles, storedPassword sql.NullString
|
||||
var userLevel sql.NullInt64
|
||||
|
||||
runQuery := func() error {
|
||||
query := rewritePlaceholders(a.getDB(), fmt.Sprintf(`SELECT id, email, user_level, roles FROM %s WHERE username = ? AND is_active = ?`, a.tableNames.Users))
|
||||
return a.getDB().QueryRowContext(ctx, query, req.Username, true).Scan(&userID, &email, &userLevel, &roles)
|
||||
query := rewritePlaceholders(a.getDB(), fmt.Sprintf(`SELECT id, email, user_level, roles, password FROM %s WHERE username = ? AND is_active = ?`, a.tableNames.Users))
|
||||
return a.getDB().QueryRowContext(ctx, query, req.Username, true).Scan(&userID, &email, &userLevel, &roles, &storedPassword)
|
||||
}
|
||||
err := runQuery()
|
||||
if isDBClosed(err) {
|
||||
@@ -439,11 +484,28 @@ func (a *JWTAuthenticator) jwtLoginDirect(ctx context.Context, req LoginRequest)
|
||||
}
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
burnPasswordCheck(req.Password)
|
||||
return nil, fmt.Errorf("invalid credentials")
|
||||
}
|
||||
return nil, fmt.Errorf("login query failed: %w", err)
|
||||
}
|
||||
|
||||
ok, needsRehash := verifyPassword(storedPassword.String, req.Password)
|
||||
if !ok {
|
||||
if storedPassword.String == "" {
|
||||
burnPasswordCheck(req.Password)
|
||||
}
|
||||
return nil, fmt.Errorf("invalid credentials")
|
||||
}
|
||||
if needsRehash && a.upgradePasswordHash {
|
||||
if h, herr := hashPassword(req.Password); herr == nil {
|
||||
q := rewritePlaceholders(a.getDB(), fmt.Sprintf(`UPDATE %s SET password = ?, updated_at = ? WHERE id = ?`, a.tableNames.Users))
|
||||
if _, uerr := a.getDB().ExecContext(ctx, q, h, time.Now(), userID); uerr != nil {
|
||||
logger.Warn("failed to upgrade legacy password hash for user %d: %v", userID, uerr)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
expiresAt := time.Now().Add(24 * time.Hour)
|
||||
tokenString := fmt.Sprintf("token_%d_%d", userID, expiresAt.Unix())
|
||||
|
||||
|
||||
Reference in New Issue
Block a user