1012 lines
29 KiB
Go
1012 lines
29 KiB
Go
package auth
|
|
|
|
import (
|
|
"context"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/stretchr/testify/assert"
|
|
"github.com/stretchr/testify/require"
|
|
|
|
"github.com/gochat/gochat/internal/config"
|
|
"github.com/gochat/gochat/internal/model"
|
|
)
|
|
|
|
// --- JWT Service Tests ---
|
|
|
|
func TestNewJWTService(t *testing.T) {
|
|
cfg := &config.JWTConfig{Secret: "testsecret", ExpiryHours: 1, RefreshExpiryHours: 24}
|
|
svc := NewJWTService(cfg)
|
|
assert.NotNil(t, svc)
|
|
}
|
|
|
|
func TestJWTGenerateTokenPair(t *testing.T) {
|
|
cfg := &config.JWTConfig{Secret: "testsecret", ExpiryHours: 1, RefreshExpiryHours: 24}
|
|
svc := NewJWTService(cfg)
|
|
|
|
customRoleID := uint(5)
|
|
user := &model.User{
|
|
Base: model.Base{ID: 1},
|
|
Provider: "email",
|
|
CustomRoleID: &customRoleID,
|
|
}
|
|
|
|
pair, err := svc.GenerateTokenPair(user, 10, "agent")
|
|
require.NoError(t, err)
|
|
assert.NotEmpty(t, pair.AccessToken)
|
|
assert.NotEmpty(t, pair.RefreshToken)
|
|
assert.True(t, pair.ExpiresAt.After(time.Now()))
|
|
}
|
|
|
|
func TestJWTGenerateTokenPairForClientSuperAdmin(t *testing.T) {
|
|
cfg := &config.JWTConfig{Secret: "testsecret", ExpiryHours: 1, RefreshExpiryHours: 24}
|
|
svc := NewJWTService(cfg)
|
|
|
|
user := &model.User{
|
|
Base: model.Base{ID: 1},
|
|
Provider: "email",
|
|
Role: "super_admin",
|
|
}
|
|
|
|
pair, err := svc.GenerateTokenPairForClient(user, 10, "super_admin", "client123")
|
|
require.NoError(t, err)
|
|
assert.NotEmpty(t, pair.AccessToken)
|
|
|
|
// Validate the token and check user_type
|
|
claims, err := svc.ValidateAccessToken(pair.AccessToken)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "super_admin", claims.UserType)
|
|
assert.Equal(t, "client123", claims.ClientID)
|
|
}
|
|
|
|
func TestJWTGenerateTokenPairForClientSuperAdminType(t *testing.T) {
|
|
cfg := &config.JWTConfig{Secret: "testsecret", ExpiryHours: 1, RefreshExpiryHours: 24}
|
|
svc := NewJWTService(cfg)
|
|
|
|
user := &model.User{
|
|
Base: model.Base{ID: 1},
|
|
Provider: "email",
|
|
Type: "SuperAdmin",
|
|
}
|
|
|
|
pair, err := svc.GenerateTokenPairForClient(user, 10, "agent", "")
|
|
require.NoError(t, err)
|
|
|
|
claims, err := svc.ValidateAccessToken(pair.AccessToken)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "super_admin", claims.UserType)
|
|
}
|
|
|
|
func TestJWTValidateAccessToken(t *testing.T) {
|
|
cfg := &config.JWTConfig{Secret: "testsecret", ExpiryHours: 1, RefreshExpiryHours: 24}
|
|
svc := NewJWTService(cfg)
|
|
|
|
user := &model.User{Base: model.Base{ID: 1}, Provider: "email"}
|
|
pair, err := svc.GenerateTokenPair(user, 10, "agent")
|
|
require.NoError(t, err)
|
|
|
|
claims, err := svc.ValidateAccessToken(pair.AccessToken)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, uint(1), claims.UserID)
|
|
assert.Equal(t, uint(10), claims.AccountID)
|
|
assert.Equal(t, "agent", claims.Role)
|
|
}
|
|
|
|
func TestJWTValidateAccessTokenInvalid(t *testing.T) {
|
|
cfg := &config.JWTConfig{Secret: "testsecret", ExpiryHours: 1, RefreshExpiryHours: 24}
|
|
svc := NewJWTService(cfg)
|
|
|
|
// invalid token
|
|
_, err := svc.ValidateAccessToken("invalid.token.here")
|
|
assert.Error(t, err)
|
|
|
|
// wrong secret
|
|
cfg2 := &config.JWTConfig{Secret: "othersecret", ExpiryHours: 1, RefreshExpiryHours: 24}
|
|
svc2 := NewJWTService(cfg2)
|
|
user := &model.User{Base: model.Base{ID: 1}, Provider: "email"}
|
|
pair, _ := svc.GenerateTokenPair(user, 10, "agent")
|
|
_, err = svc2.ValidateAccessToken(pair.AccessToken)
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
func TestJWTValidateAccessTokenRefreshTokenRejected(t *testing.T) {
|
|
cfg := &config.JWTConfig{Secret: "testsecret", ExpiryHours: 1, RefreshExpiryHours: 24}
|
|
svc := NewJWTService(cfg)
|
|
|
|
user := &model.User{Base: model.Base{ID: 1}, Provider: "email"}
|
|
pair, err := svc.GenerateTokenPair(user, 10, "agent")
|
|
require.NoError(t, err)
|
|
|
|
// refresh token should not validate as access token
|
|
_, err = svc.ValidateAccessToken(pair.RefreshToken)
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
func TestJWTValidateRefreshToken(t *testing.T) {
|
|
cfg := &config.JWTConfig{Secret: "testsecret", ExpiryHours: 1, RefreshExpiryHours: 24}
|
|
svc := NewJWTService(cfg)
|
|
|
|
user := &model.User{Base: model.Base{ID: 1}, Provider: "email"}
|
|
pair, err := svc.GenerateTokenPair(user, 10, "agent")
|
|
require.NoError(t, err)
|
|
|
|
claims, err := svc.ValidateRefreshToken(pair.RefreshToken)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, uint(1), claims.UserID)
|
|
assert.Equal(t, "email", claims.Provider)
|
|
}
|
|
|
|
func TestJWTValidateRefreshTokenInvalid(t *testing.T) {
|
|
cfg := &config.JWTConfig{Secret: "testsecret", ExpiryHours: 1, RefreshExpiryHours: 24}
|
|
svc := NewJWTService(cfg)
|
|
|
|
// invalid token
|
|
_, err := svc.ValidateRefreshToken("invalid.token.here")
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
func TestJWTValidateRefreshTokenAccessTokenRejected(t *testing.T) {
|
|
cfg := &config.JWTConfig{Secret: "testsecret", ExpiryHours: 1, RefreshExpiryHours: 24}
|
|
svc := NewJWTService(cfg)
|
|
|
|
user := &model.User{Base: model.Base{ID: 1}, Provider: "email"}
|
|
pair, err := svc.GenerateTokenPair(user, 10, "agent")
|
|
require.NoError(t, err)
|
|
|
|
// access token should not validate as refresh token
|
|
_, err = svc.ValidateRefreshToken(pair.AccessToken)
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
func TestJWTRefreshAccessToken(t *testing.T) {
|
|
cfg := &config.JWTConfig{Secret: "testsecret", ExpiryHours: 1, RefreshExpiryHours: 24}
|
|
svc := NewJWTService(cfg)
|
|
|
|
user := &model.User{Base: model.Base{ID: 1}, Provider: "email"}
|
|
pair, err := svc.GenerateTokenPair(user, 10, "agent")
|
|
require.NoError(t, err)
|
|
|
|
newPair, err := svc.RefreshAccessToken(pair.RefreshToken, 10, "administrator")
|
|
require.NoError(t, err)
|
|
assert.NotEmpty(t, newPair.AccessToken)
|
|
assert.NotEmpty(t, newPair.RefreshToken)
|
|
assert.NotEqual(t, pair.AccessToken, newPair.AccessToken)
|
|
}
|
|
|
|
func TestJWTRefreshAccessTokenInvalid(t *testing.T) {
|
|
cfg := &config.JWTConfig{Secret: "testsecret", ExpiryHours: 1, RefreshExpiryHours: 24}
|
|
svc := NewJWTService(cfg)
|
|
|
|
_, err := svc.RefreshAccessToken("invalid", 10, "agent")
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
// --- Permission Tests ---
|
|
|
|
func TestHasPermission(t *testing.T) {
|
|
// super admin has all permissions
|
|
assert.True(t, HasPermission(RoleSuperAdmin, PermAccountCreate))
|
|
assert.True(t, HasPermission(RoleSuperAdmin, PermPlatformManage))
|
|
|
|
// administrator
|
|
assert.True(t, HasPermission(RoleAdministrator, PermAccountRead))
|
|
assert.True(t, HasPermission(RoleAdministrator, PermInboxCreate))
|
|
assert.False(t, HasPermission(RoleAdministrator, PermPlatformManage))
|
|
|
|
// agent
|
|
assert.True(t, HasPermission(RoleAgent, PermInboxRead))
|
|
assert.True(t, HasPermission(RoleAgent, PermConversationRead))
|
|
assert.False(t, HasPermission(RoleAgent, PermInboxCreate))
|
|
assert.False(t, HasPermission(RoleAgent, PermAccountUpdate))
|
|
|
|
// unknown role
|
|
assert.False(t, HasPermission(Role("unknown"), PermAccountRead))
|
|
}
|
|
|
|
func TestGetPermissions(t *testing.T) {
|
|
// super admin returns all permissions
|
|
perms := GetPermissions(RoleSuperAdmin)
|
|
assert.NotEmpty(t, perms)
|
|
|
|
// administrator
|
|
perms = GetPermissions(RoleAdministrator)
|
|
assert.NotEmpty(t, perms)
|
|
assert.Contains(t, perms, PermAccountRead)
|
|
|
|
// agent
|
|
perms = GetPermissions(RoleAgent)
|
|
assert.NotEmpty(t, perms)
|
|
assert.Contains(t, perms, PermInboxRead)
|
|
|
|
// unknown role
|
|
perms = GetPermissions(Role("unknown"))
|
|
assert.Empty(t, perms)
|
|
}
|
|
|
|
// --- Policy Tests ---
|
|
|
|
func TestPermissionLevelIsValid(t *testing.T) {
|
|
assert.True(t, PermissionFull.IsValid())
|
|
assert.True(t, PermissionRead.IsValid())
|
|
assert.True(t, PermissionNone.IsValid())
|
|
assert.False(t, PermissionLevel("invalid").IsValid())
|
|
}
|
|
|
|
func TestPermissionLevelCanWrite(t *testing.T) {
|
|
assert.True(t, PermissionFull.CanWrite())
|
|
assert.False(t, PermissionRead.CanWrite())
|
|
assert.False(t, PermissionNone.CanWrite())
|
|
}
|
|
|
|
func TestPermissionLevelCanRead(t *testing.T) {
|
|
assert.True(t, PermissionFull.CanRead())
|
|
assert.True(t, PermissionRead.CanRead())
|
|
assert.False(t, PermissionNone.CanRead())
|
|
}
|
|
|
|
func TestPermissionMatrixMapToJSON(t *testing.T) {
|
|
m := PermissionMatrixMap{
|
|
DimensionConversationManage: PermissionFull,
|
|
DimensionContactManage: PermissionRead,
|
|
}
|
|
data, err := m.ToJSON()
|
|
require.NoError(t, err)
|
|
assert.NotEmpty(t, data)
|
|
}
|
|
|
|
func TestPermissionMatrixFromJSON(t *testing.T) {
|
|
// valid
|
|
data := []byte(`{"conversation_manage":"full","contact_manage":"read"}`)
|
|
m, err := PermissionMatrixFromJSON(data)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, PermissionFull, m[DimensionConversationManage])
|
|
assert.Equal(t, PermissionRead, m[DimensionContactManage])
|
|
|
|
// invalid level
|
|
data = []byte(`{"conversation_manage":"invalid"}`)
|
|
_, err = PermissionMatrixFromJSON(data)
|
|
assert.Error(t, err)
|
|
|
|
// invalid JSON
|
|
_, err = PermissionMatrixFromJSON([]byte(`{invalid}`))
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
func TestNewPolicyContextAdministrator(t *testing.T) {
|
|
pc := NewPolicyContext(1, 10, "administrator", 0, nil)
|
|
assert.NotNil(t, pc)
|
|
assert.True(t, pc.IsAdministrator())
|
|
assert.False(t, pc.IsAgent())
|
|
assert.False(t, pc.IsCustomRole())
|
|
assert.Equal(t, AdministratorPermissions, pc.Permissions)
|
|
}
|
|
|
|
func TestNewPolicyContextAgent(t *testing.T) {
|
|
pc := NewPolicyContext(1, 10, "agent", 0, nil)
|
|
assert.NotNil(t, pc)
|
|
assert.False(t, pc.IsAdministrator())
|
|
assert.True(t, pc.IsAgent())
|
|
assert.False(t, pc.IsCustomRole())
|
|
assert.Equal(t, AgentDefaultPermissions, pc.Permissions)
|
|
}
|
|
|
|
func TestNewPolicyContextAgentWithCustomRole(t *testing.T) {
|
|
pc := NewPolicyContext(1, 10, "agent", 5, nil)
|
|
assert.NotNil(t, pc)
|
|
assert.True(t, pc.IsCustomRole())
|
|
assert.False(t, pc.IsAgent())
|
|
}
|
|
|
|
func TestNewPolicyContextCustomRole(t *testing.T) {
|
|
perms := PermissionMatrixMap{
|
|
DimensionConversationManage: PermissionFull,
|
|
}
|
|
pc := NewPolicyContext(1, 10, "custom_role", 5, perms)
|
|
assert.NotNil(t, pc)
|
|
assert.True(t, pc.IsCustomRole())
|
|
assert.Equal(t, perms, pc.Permissions)
|
|
}
|
|
|
|
func TestNewPolicyContextCustomRoleNilPerms(t *testing.T) {
|
|
pc := NewPolicyContext(1, 10, "custom_role", 5, nil)
|
|
assert.NotNil(t, pc)
|
|
assert.Equal(t, AgentDefaultPermissions, pc.Permissions)
|
|
}
|
|
|
|
func TestPolicyContextCanAdministrator(t *testing.T) {
|
|
pc := NewPolicyContext(1, 10, "administrator", 0, nil)
|
|
// Administrator can do everything
|
|
assert.True(t, pc.Can("manage", "conversation"))
|
|
assert.True(t, pc.Can("read", "contact"))
|
|
assert.True(t, pc.Can("delete", "conversation"))
|
|
assert.True(t, pc.Can("create", "inbox"))
|
|
}
|
|
|
|
func TestPolicyContextCanSuperAdmin(t *testing.T) {
|
|
pc := &PolicyContext{UserID: 1, AccountID: 10, Role: "super_admin"}
|
|
assert.True(t, pc.Can("manage", "conversation"))
|
|
assert.True(t, pc.Can("delete", "conversation"))
|
|
}
|
|
|
|
func TestPolicyContextCanAgent(t *testing.T) {
|
|
pc := NewPolicyContext(1, 10, "agent", 0, nil)
|
|
// Agent has read on conversations, no write
|
|
assert.True(t, pc.Can("read", "conversation"))
|
|
assert.False(t, pc.Can("manage", "conversation"))
|
|
assert.False(t, pc.Can("delete", "conversation"))
|
|
// Agent can create messages (reply)
|
|
assert.True(t, pc.Can("create", "message"))
|
|
}
|
|
|
|
func TestPolicyContextCanCustomRole(t *testing.T) {
|
|
perms := PermissionMatrixMap{
|
|
DimensionConversationManage: PermissionFull,
|
|
DimensionContactManage: PermissionRead,
|
|
DimensionConversationDelete: PermissionFull,
|
|
}
|
|
pc := NewPolicyContext(1, 10, "custom_role", 5, perms)
|
|
assert.True(t, pc.Can("manage", "conversation"))
|
|
assert.True(t, pc.Can("delete", "conversation"))
|
|
assert.True(t, pc.Can("read", "contact"))
|
|
assert.False(t, pc.Can("manage", "contact"))
|
|
}
|
|
|
|
func TestPolicyContextCanUnknownResource(t *testing.T) {
|
|
pc := NewPolicyContext(1, 10, "custom_role", 5, PermissionMatrixMap{})
|
|
assert.False(t, pc.Can("manage", "unknown_resource"))
|
|
}
|
|
|
|
func TestPolicyContextCanNilPermissionsAgent(t *testing.T) {
|
|
pc := &PolicyContext{UserID: 1, AccountID: 10, Role: "agent"}
|
|
// With nil permissions, agent should still get defaults applied in Can
|
|
assert.True(t, pc.Can("read", "conversation"))
|
|
}
|
|
|
|
func TestPolicyContextCanNilPermissionsUnknown(t *testing.T) {
|
|
pc := &PolicyContext{UserID: 1, AccountID: 10, Role: "unknown"}
|
|
assert.False(t, pc.Can("read", "conversation"))
|
|
}
|
|
|
|
func TestPolicyContextGetPermissionLevel(t *testing.T) {
|
|
pc := NewPolicyContext(1, 10, "agent", 0, nil)
|
|
assert.Equal(t, PermissionRead, pc.GetPermissionLevel(DimensionConversationManage))
|
|
assert.Equal(t, PermissionNone, pc.GetPermissionLevel(DimensionReportManage))
|
|
}
|
|
|
|
func TestPolicyContextHasFeatureAccess(t *testing.T) {
|
|
pc := NewPolicyContext(1, 10, "agent", 0, nil)
|
|
// Agent has read on conversation_manage
|
|
assert.True(t, pc.HasFeatureAccess(DimensionConversationManage, PermissionRead))
|
|
assert.False(t, pc.HasFeatureAccess(DimensionConversationManage, PermissionFull))
|
|
// None required always true
|
|
assert.True(t, pc.HasFeatureAccess(DimensionReportManage, PermissionNone))
|
|
// Report manage is none for agents
|
|
assert.False(t, pc.HasFeatureAccess(DimensionReportManage, PermissionRead))
|
|
|
|
// Invalid required level
|
|
assert.False(t, pc.HasFeatureAccess(DimensionConversationManage, PermissionLevel("invalid")))
|
|
}
|
|
|
|
// --- Session Store Tests ---
|
|
|
|
func TestNewSessionStore(t *testing.T) {
|
|
cfg := &config.SessionConfig{ExpirySeconds: 3600, TokenLength: 32}
|
|
store := NewSessionStore(cfg)
|
|
assert.NotNil(t, store)
|
|
assert.Equal(t, 0, store.Count())
|
|
}
|
|
|
|
func TestSessionStoreCreate(t *testing.T) {
|
|
cfg := &config.SessionConfig{ExpirySeconds: 3600, TokenLength: 32}
|
|
store := NewSessionStore(cfg)
|
|
|
|
session, err := store.Create(1, 10, "agent", "email")
|
|
require.NoError(t, err)
|
|
assert.NotEmpty(t, session.ID)
|
|
assert.Equal(t, uint(1), session.UserID)
|
|
assert.Equal(t, uint(10), session.AccountID)
|
|
assert.Equal(t, "agent", session.Role)
|
|
assert.Equal(t, "email", session.Provider)
|
|
assert.True(t, session.ExpiresAt.After(time.Now()))
|
|
assert.Equal(t, 1, store.Count())
|
|
}
|
|
|
|
func TestSessionStoreGet(t *testing.T) {
|
|
cfg := &config.SessionConfig{ExpirySeconds: 3600, TokenLength: 32}
|
|
store := NewSessionStore(cfg)
|
|
|
|
session, err := store.Create(1, 10, "agent", "email")
|
|
require.NoError(t, err)
|
|
|
|
got, err := store.Get(session.ID)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, session.ID, got.ID)
|
|
assert.Equal(t, uint(1), got.UserID)
|
|
}
|
|
|
|
func TestSessionStoreGetNotFound(t *testing.T) {
|
|
cfg := &config.SessionConfig{ExpirySeconds: 3600, TokenLength: 32}
|
|
store := NewSessionStore(cfg)
|
|
|
|
_, err := store.Get("nonexistent")
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
func TestSessionStoreGetExpired(t *testing.T) {
|
|
cfg := &config.SessionConfig{ExpirySeconds: 0, TokenLength: 32} // 0 means immediate expiry
|
|
store := NewSessionStore(cfg)
|
|
|
|
// Manually set very short expiry
|
|
cfg.ExpirySeconds = 1
|
|
session, err := store.Create(1, 10, "agent", "email")
|
|
require.NoError(t, err)
|
|
|
|
// Wait for expiry
|
|
time.Sleep(2 * time.Second)
|
|
_, err = store.Get(session.ID)
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
func TestSessionStoreDelete(t *testing.T) {
|
|
cfg := &config.SessionConfig{ExpirySeconds: 3600, TokenLength: 32}
|
|
store := NewSessionStore(cfg)
|
|
|
|
session, _ := store.Create(1, 10, "agent", "email")
|
|
err := store.Delete(session.ID)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, 0, store.Count())
|
|
|
|
// delete non-existent - no error
|
|
err = store.Delete("nonexistent")
|
|
assert.NoError(t, err)
|
|
}
|
|
|
|
func TestSessionStoreDeleteByUserID(t *testing.T) {
|
|
cfg := &config.SessionConfig{ExpirySeconds: 3600, TokenLength: 32}
|
|
store := NewSessionStore(cfg)
|
|
|
|
store.Create(1, 10, "agent", "email")
|
|
store.Create(1, 10, "agent", "email")
|
|
store.Create(2, 10, "agent", "email")
|
|
|
|
count := store.DeleteByUserID(1)
|
|
assert.Equal(t, 2, count)
|
|
assert.Equal(t, 1, store.Count())
|
|
}
|
|
|
|
func TestSessionStoreRefresh(t *testing.T) {
|
|
cfg := &config.SessionConfig{ExpirySeconds: 3600, TokenLength: 32}
|
|
store := NewSessionStore(cfg)
|
|
|
|
session, err := store.Create(1, 10, "agent", "email")
|
|
require.NoError(t, err)
|
|
originalExpiry := session.ExpiresAt
|
|
|
|
refreshed, err := store.Refresh(session.ID)
|
|
require.NoError(t, err)
|
|
assert.True(t, refreshed.ExpiresAt.After(originalExpiry) || refreshed.ExpiresAt.Equal(originalExpiry))
|
|
}
|
|
|
|
func TestSessionStoreRefreshNotFound(t *testing.T) {
|
|
cfg := &config.SessionConfig{ExpirySeconds: 3600, TokenLength: 32}
|
|
store := NewSessionStore(cfg)
|
|
|
|
_, err := store.Refresh("nonexistent")
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
func TestSessionStoreRefreshExpired(t *testing.T) {
|
|
cfg := &config.SessionConfig{ExpirySeconds: 1, TokenLength: 32}
|
|
store := NewSessionStore(cfg)
|
|
|
|
session, err := store.Create(1, 10, "agent", "email")
|
|
require.NoError(t, err)
|
|
|
|
time.Sleep(2 * time.Second)
|
|
_, err = store.Refresh(session.ID)
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
func TestSessionStoreSetData(t *testing.T) {
|
|
cfg := &config.SessionConfig{ExpirySeconds: 3600, TokenLength: 32}
|
|
store := NewSessionStore(cfg)
|
|
|
|
session, _ := store.Create(1, 10, "agent", "email")
|
|
|
|
err := store.SetData(session.ID, "key1", "value1")
|
|
require.NoError(t, err)
|
|
|
|
val, err := store.GetData(session.ID, "key1")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "value1", val)
|
|
}
|
|
|
|
func TestSessionStoreSetDataNotFound(t *testing.T) {
|
|
cfg := &config.SessionConfig{ExpirySeconds: 3600, TokenLength: 32}
|
|
store := NewSessionStore(cfg)
|
|
|
|
err := store.SetData("nonexistent", "key1", "value1")
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
func TestSessionStoreGetDataNotFound(t *testing.T) {
|
|
cfg := &config.SessionConfig{ExpirySeconds: 3600, TokenLength: 32}
|
|
store := NewSessionStore(cfg)
|
|
|
|
// session not found
|
|
_, err := store.GetData("nonexistent", "key1")
|
|
assert.Error(t, err)
|
|
|
|
// key not found
|
|
session, _ := store.Create(1, 10, "agent", "email")
|
|
_, err = store.GetData(session.ID, "nonexistent_key")
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
func TestSessionStoreCleanupExpired(t *testing.T) {
|
|
cfg := &config.SessionConfig{ExpirySeconds: 1, TokenLength: 32}
|
|
store := NewSessionStore(cfg)
|
|
|
|
store.Create(1, 10, "agent", "email")
|
|
store.Create(2, 10, "agent", "email")
|
|
|
|
time.Sleep(2 * time.Second)
|
|
count := store.CleanupExpired()
|
|
assert.Equal(t, 2, count)
|
|
assert.Equal(t, 0, store.Count())
|
|
}
|
|
|
|
func TestSessionStoreCount(t *testing.T) {
|
|
cfg := &config.SessionConfig{ExpirySeconds: 3600, TokenLength: 32}
|
|
store := NewSessionStore(cfg)
|
|
|
|
assert.Equal(t, 0, store.Count())
|
|
store.Create(1, 10, "agent", "email")
|
|
assert.Equal(t, 1, store.Count())
|
|
store.Create(2, 10, "agent", "email")
|
|
assert.Equal(t, 2, store.Count())
|
|
}
|
|
|
|
func TestGenerateSessionID(t *testing.T) {
|
|
id, err := generateSessionID(32)
|
|
require.NoError(t, err)
|
|
assert.Len(t, id, 64) // hex encoding doubles length
|
|
|
|
// different calls produce different IDs
|
|
id2, _ := generateSessionID(32)
|
|
assert.NotEqual(t, id, id2)
|
|
}
|
|
|
|
// --- Refresh Token Store Tests ---
|
|
|
|
func TestRefreshTokenStoreStoreValidate(t *testing.T) {
|
|
store := NewRefreshTokenStore(nil, &config.JWTConfig{RefreshExpiryHours: 24})
|
|
ctx := context.Background()
|
|
|
|
err := store.Store(ctx, 1, "token123")
|
|
require.NoError(t, err)
|
|
|
|
valid, err := store.Validate(ctx, 1, "token123")
|
|
require.NoError(t, err)
|
|
assert.True(t, valid)
|
|
|
|
valid, err = store.Validate(ctx, 1, "wrong")
|
|
require.NoError(t, err)
|
|
assert.False(t, valid)
|
|
|
|
valid, err = store.Validate(ctx, 999, "token123")
|
|
require.NoError(t, err)
|
|
assert.False(t, valid)
|
|
}
|
|
|
|
func TestRefreshTokenStoreRevoke(t *testing.T) {
|
|
store := NewRefreshTokenStore(nil, &config.JWTConfig{RefreshExpiryHours: 24})
|
|
ctx := context.Background()
|
|
|
|
store.Store(ctx, 1, "token123")
|
|
err := store.Revoke(ctx, 1)
|
|
require.NoError(t, err)
|
|
|
|
valid, _ := store.Validate(ctx, 1, "token123")
|
|
assert.False(t, valid)
|
|
}
|
|
|
|
func TestRefreshTokenStoreRotate(t *testing.T) {
|
|
store := NewRefreshTokenStore(nil, &config.JWTConfig{RefreshExpiryHours: 24})
|
|
ctx := context.Background()
|
|
|
|
store.Store(ctx, 1, "old_token")
|
|
err := store.Rotate(ctx, 1, "new_token")
|
|
require.NoError(t, err)
|
|
|
|
valid, _ := store.Validate(ctx, 1, "new_token")
|
|
assert.True(t, valid)
|
|
}
|
|
|
|
func TestRefreshTokenStoreRotateForClient(t *testing.T) {
|
|
store := NewRefreshTokenStore(nil, &config.JWTConfig{RefreshExpiryHours: 24})
|
|
ctx := context.Background()
|
|
|
|
err := store.RotateForClient(ctx, 1, "client1", "new_token")
|
|
require.NoError(t, err)
|
|
|
|
valid, _ := store.ValidateForClient(ctx, 1, "client1", "new_token")
|
|
assert.True(t, valid)
|
|
}
|
|
|
|
func TestRefreshTokenStoreHasClient(t *testing.T) {
|
|
store := NewRefreshTokenStore(nil, &config.JWTConfig{RefreshExpiryHours: 24})
|
|
ctx := context.Background()
|
|
|
|
has, err := store.HasClient(ctx, 1, "client1")
|
|
require.NoError(t, err)
|
|
assert.False(t, has)
|
|
|
|
store.StoreForClient(ctx, 1, "client1", "token")
|
|
has, err = store.HasClient(ctx, 1, "client1")
|
|
require.NoError(t, err)
|
|
assert.True(t, has)
|
|
|
|
has, err = store.HasClient(ctx, 1, "client2")
|
|
require.NoError(t, err)
|
|
assert.False(t, has)
|
|
}
|
|
|
|
func TestRefreshTokenStoreExpiredToken(t *testing.T) {
|
|
store := NewRefreshTokenStore(nil, &config.JWTConfig{RefreshExpiryHours: 24})
|
|
ctx := context.Background()
|
|
|
|
store.Store(ctx, 1, "token123")
|
|
|
|
// Simulate expiry by modifying the stored entry
|
|
store.mu.Lock()
|
|
for k, v := range store.mem {
|
|
v.expiresAt = time.Now().Add(-1 * time.Hour)
|
|
store.mem[k] = v
|
|
}
|
|
store.mu.Unlock()
|
|
|
|
valid, _ := store.Validate(ctx, 1, "token123")
|
|
assert.False(t, valid)
|
|
}
|
|
|
|
// --- Webhook Token Registry Tests ---
|
|
|
|
func TestNewWebhookTokenRegistry(t *testing.T) {
|
|
reg := NewWebhookTokenRegistry()
|
|
assert.NotNil(t, reg)
|
|
}
|
|
|
|
func TestWebhookTokenRegistryRegisterLookup(t *testing.T) {
|
|
reg := NewWebhookTokenRegistry()
|
|
entry := WebhookTokenEntry{
|
|
InboxID: 1,
|
|
AccountID: 10,
|
|
Secret: "mysecret",
|
|
Identifier: "bot123",
|
|
}
|
|
|
|
reg.Register("telegram", "bot123", entry)
|
|
|
|
got, found := reg.Lookup("telegram", "bot123")
|
|
require.True(t, found)
|
|
assert.Equal(t, uint(1), got.InboxID)
|
|
assert.Equal(t, "mysecret", got.Secret)
|
|
}
|
|
|
|
func TestWebhookTokenRegistryLookupNotFound(t *testing.T) {
|
|
reg := NewWebhookTokenRegistry()
|
|
|
|
// unknown channel type
|
|
_, found := reg.Lookup("telegram", "bot123")
|
|
assert.False(t, found)
|
|
|
|
// register then lookup different identifier
|
|
reg.Register("telegram", "bot123", WebhookTokenEntry{Identifier: "bot123"})
|
|
_, found = reg.Lookup("telegram", "bot456")
|
|
assert.False(t, found)
|
|
}
|
|
|
|
func TestWebhookTokenRegistryUnregister(t *testing.T) {
|
|
reg := NewWebhookTokenRegistry()
|
|
reg.Register("telegram", "bot123", WebhookTokenEntry{Identifier: "bot123"})
|
|
reg.Unregister("telegram", "bot123")
|
|
|
|
_, found := reg.Lookup("telegram", "bot123")
|
|
assert.False(t, found)
|
|
}
|
|
|
|
func TestWebhookTokenRegistryUnregisterUnknownChannel(t *testing.T) {
|
|
reg := NewWebhookTokenRegistry()
|
|
// should not panic
|
|
reg.Unregister("unknown", "bot123")
|
|
}
|
|
|
|
func TestWebhookTokenRegistryValidate(t *testing.T) {
|
|
reg := NewWebhookTokenRegistry()
|
|
|
|
// unknown identifier
|
|
req := httptest.NewRequest(http.MethodGet, "/webhook", nil)
|
|
valid, err := reg.Validate("telegram", "unknown", req)
|
|
require.NoError(t, err)
|
|
assert.False(t, valid)
|
|
|
|
// telegram with secret
|
|
reg.Register("telegram", "bot123", WebhookTokenEntry{Secret: "mysecret", Identifier: "bot123"})
|
|
req.Header.Set("X-Telegram-Bot-Api-Secret-Token", "mysecret")
|
|
valid, err = reg.Validate("telegram", "bot123", req)
|
|
require.NoError(t, err)
|
|
assert.True(t, valid)
|
|
|
|
// telegram with wrong secret
|
|
req.Header.Set("X-Telegram-Bot-Api-Secret-Token", "wrong")
|
|
valid, err = reg.Validate("telegram", "bot123", req)
|
|
require.NoError(t, err)
|
|
assert.False(t, valid)
|
|
|
|
// telegram with no secret configured
|
|
reg.Register("telegram", "bot456", WebhookTokenEntry{Secret: "", Identifier: "bot456"})
|
|
valid, err = reg.Validate("telegram", "bot456", req)
|
|
require.NoError(t, err)
|
|
assert.True(t, valid)
|
|
|
|
// web_widget always valid
|
|
reg.Register("web_widget", "widget1", WebhookTokenEntry{Identifier: "widget1"})
|
|
valid, err = reg.Validate("web_widget", "widget1", req)
|
|
require.NoError(t, err)
|
|
assert.True(t, valid)
|
|
|
|
// facebook
|
|
reg.Register("facebook", "fb1", WebhookTokenEntry{Identifier: "fb1"})
|
|
valid, err = reg.Validate("facebook", "fb1", req)
|
|
require.NoError(t, err)
|
|
assert.True(t, valid)
|
|
|
|
// whatsapp
|
|
reg.Register("whatsapp", "wa1", WebhookTokenEntry{Identifier: "wa1"})
|
|
valid, err = reg.Validate("whatsapp", "wa1", req)
|
|
require.NoError(t, err)
|
|
assert.True(t, valid)
|
|
|
|
// unknown channel type
|
|
reg.Register("custom", "c1", WebhookTokenEntry{Identifier: "c1"})
|
|
valid, err = reg.Validate("custom", "c1", req)
|
|
require.NoError(t, err)
|
|
assert.True(t, valid)
|
|
}
|
|
|
|
func TestWebhookTokenRegistryGetAllIdentifiers(t *testing.T) {
|
|
reg := NewWebhookTokenRegistry()
|
|
|
|
// empty
|
|
entries := reg.GetAllIdentifiers("telegram")
|
|
assert.Empty(t, entries)
|
|
|
|
// with entries
|
|
reg.Register("telegram", "bot1", WebhookTokenEntry{Identifier: "bot1"})
|
|
reg.Register("telegram", "bot2", WebhookTokenEntry{Identifier: "bot2"})
|
|
|
|
entries = reg.GetAllIdentifiers("telegram")
|
|
assert.Len(t, entries, 2)
|
|
|
|
// unknown channel type
|
|
entries = reg.GetAllIdentifiers("unknown")
|
|
assert.Empty(t, entries)
|
|
}
|
|
|
|
// --- SSO Session Store Tests ---
|
|
|
|
func TestNewSSOSessionStore(t *testing.T) {
|
|
store := NewSSOSessionStore(nil, 0)
|
|
assert.NotNil(t, store)
|
|
assert.Equal(t, 24*time.Hour, store.SessionTTL())
|
|
|
|
store2 := NewSSOSessionStore(nil, 48*time.Hour)
|
|
assert.Equal(t, 48*time.Hour, store2.SessionTTL())
|
|
}
|
|
|
|
func TestSSOSessionStoreCreate(t *testing.T) {
|
|
store := NewSSOSessionStore(nil, 24*time.Hour)
|
|
ctx := context.Background()
|
|
|
|
data := &SSOSessionData{
|
|
SessionID: "sso_test123",
|
|
UserID: 1,
|
|
Provider: "oidc",
|
|
}
|
|
|
|
id, err := store.Create(ctx, data)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "sso_test123", id)
|
|
}
|
|
|
|
func TestSSOSessionStoreGet(t *testing.T) {
|
|
store := NewSSOSessionStore(nil, 24*time.Hour)
|
|
ctx := context.Background()
|
|
|
|
result, err := store.Get(ctx, "sso_test123")
|
|
require.NoError(t, err)
|
|
assert.Nil(t, result)
|
|
}
|
|
|
|
func TestSSOSessionStoreGetByUser(t *testing.T) {
|
|
store := NewSSOSessionStore(nil, 24*time.Hour)
|
|
ctx := context.Background()
|
|
|
|
result, err := store.GetByUser(ctx, 1)
|
|
require.NoError(t, err)
|
|
assert.Nil(t, result)
|
|
}
|
|
|
|
func TestSSOSessionStoreGetByIdP(t *testing.T) {
|
|
store := NewSSOSessionStore(nil, 24*time.Hour)
|
|
ctx := context.Background()
|
|
|
|
result, err := store.GetByIdP(ctx, "idp_entity")
|
|
require.NoError(t, err)
|
|
assert.Nil(t, result)
|
|
}
|
|
|
|
func TestSSOSessionStoreTerminate(t *testing.T) {
|
|
store := NewSSOSessionStore(nil, 24*time.Hour)
|
|
ctx := context.Background()
|
|
|
|
terminated, err := store.Terminate(ctx, "sso_test123")
|
|
require.NoError(t, err)
|
|
assert.True(t, terminated)
|
|
}
|
|
|
|
func TestSSOSessionStoreTerminateUserSessions(t *testing.T) {
|
|
store := NewSSOSessionStore(nil, 24*time.Hour)
|
|
ctx := context.Background()
|
|
|
|
count, err := store.TerminateUserSessions(ctx, 1)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, 0, count)
|
|
}
|
|
|
|
func TestSSOSessionStoreTerminateIdPSessions(t *testing.T) {
|
|
store := NewSSOSessionStore(nil, 24*time.Hour)
|
|
ctx := context.Background()
|
|
|
|
count, err := store.TerminateIdPSessions(ctx, "idp_entity")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, 0, count)
|
|
}
|
|
|
|
func TestSSOSessionStoreRefresh(t *testing.T) {
|
|
store := NewSSOSessionStore(nil, 24*time.Hour)
|
|
ctx := context.Background()
|
|
|
|
err := store.Refresh(ctx, "sso_test123", 48*time.Hour)
|
|
require.NoError(t, err)
|
|
}
|
|
|
|
func TestSSOSessionStoreCountByUser(t *testing.T) {
|
|
store := NewSSOSessionStore(nil, 24*time.Hour)
|
|
ctx := context.Background()
|
|
|
|
count, err := store.CountByUser(ctx, 1)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, int64(0), count)
|
|
}
|
|
|
|
func TestSSOSessionStoreExists(t *testing.T) {
|
|
store := NewSSOSessionStore(nil, 24*time.Hour)
|
|
ctx := context.Background()
|
|
|
|
exists, err := store.Exists(ctx, "sso_test123")
|
|
require.NoError(t, err)
|
|
assert.False(t, exists)
|
|
}
|
|
|
|
// --- Platform Auth Helper Tests ---
|
|
|
|
func TestAuthTokenPrefix(t *testing.T) {
|
|
assert.Equal(t, "12345678", authTokenPrefix("12345678901234"))
|
|
assert.Equal(t, "short", authTokenPrefix("short"))
|
|
assert.Equal(t, "", authTokenPrefix(""))
|
|
}
|
|
|
|
func TestHashAuthToken(t *testing.T) {
|
|
hash := hashAuthToken("testkey")
|
|
assert.NotEmpty(t, hash)
|
|
assert.Len(t, hash, 64) // SHA-256 hex = 64 chars
|
|
|
|
// same input produces same hash
|
|
hash2 := hashAuthToken("testkey")
|
|
assert.Equal(t, hash, hash2)
|
|
|
|
// different input produces different hash
|
|
hash3 := hashAuthToken("otherkey")
|
|
assert.NotEqual(t, hash, hash3)
|
|
}
|
|
|
|
func TestGeneratePlatformAPIKey(t *testing.T) {
|
|
key := generatePlatformAPIKey()
|
|
assert.Contains(t, key, "gochat_pa_")
|
|
assert.True(t, len(key) > 40)
|
|
|
|
// different calls produce different keys
|
|
key2 := generatePlatformAPIKey()
|
|
assert.NotEqual(t, key, key2)
|
|
}
|
|
|
|
func TestGenerateAgentBotToken(t *testing.T) {
|
|
token := generateAgentBotToken()
|
|
assert.Contains(t, token, "gochat_ab_")
|
|
assert.True(t, len(token) > 40)
|
|
|
|
// different calls produce different tokens
|
|
token2 := generateAgentBotToken()
|
|
assert.NotEqual(t, token, token2)
|
|
}
|
|
|
|
func TestAgentBotTableName(t *testing.T) {
|
|
assert.Equal(t, "agent_bots", AgentBot{}.TableName())
|
|
}
|
|
|
|
// --- OIDC fetchUserInfo test ---
|
|
|
|
func TestOIDCFetchUserInfo(t *testing.T) {
|
|
// Create a test HTTP server that returns userinfo
|
|
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
assert.Equal(t, "Bearer mytoken", r.Header.Get("Authorization"))
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.Write([]byte(`{"sub":"user123","email":"test@example.com","name":"Test User"}`))
|
|
}))
|
|
defer ts.Close()
|
|
|
|
svc := &OIDCService{
|
|
httpClient: ts.Client(),
|
|
}
|
|
|
|
claims, err := svc.fetchUserInfo(context.Background(), "mytoken", ts.URL)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "user123", claims["sub"])
|
|
assert.Equal(t, "test@example.com", claims["email"])
|
|
}
|
|
|
|
func TestOIDCFetchUserInfoError(t *testing.T) {
|
|
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
w.WriteHeader(http.StatusUnauthorized)
|
|
w.Write([]byte(`{"error":"invalid_token"}`))
|
|
}))
|
|
defer ts.Close()
|
|
|
|
svc := &OIDCService{
|
|
httpClient: ts.Client(),
|
|
}
|
|
|
|
_, err := svc.fetchUserInfo(context.Background(), "mytoken", ts.URL)
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
func TestOIDCFetchUserInfoInvalidJSON(t *testing.T) {
|
|
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.Write([]byte(`{invalid json`))
|
|
}))
|
|
defer ts.Close()
|
|
|
|
svc := &OIDCService{
|
|
httpClient: ts.Client(),
|
|
}
|
|
|
|
_, err := svc.fetchUserInfo(context.Background(), "mytoken", ts.URL)
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
func TestOIDCGetAccountSettingsNilDB(t *testing.T) {
|
|
svc := &OIDCService{
|
|
cfg: &config.OIDCConfig{},
|
|
db: nil,
|
|
}
|
|
|
|
settings, err := svc.GetAccountSettings(1)
|
|
assert.Error(t, err)
|
|
assert.Nil(t, settings)
|
|
}
|