Files
cmbone/internal/services/store_test.go
T

159 lines
4.2 KiB
Go
Raw Normal View History

2026-07-08 15:13:26 +08:00
package services
import (
"database/sql"
2026-07-08 23:37:19 +08:00
"errors"
2026-07-08 15:13:26 +08:00
"testing"
)
func newTestStore(t *testing.T) *SuiStore {
t.Helper()
db, err := sql.Open("sqlite", ":memory:")
if err != nil {
t.Fatalf("open sqlite: %v", err)
}
t.Cleanup(func() {
_ = db.Close()
})
if err := Migrate(db); err != nil {
t.Fatalf("migrate: %v", err)
}
return &SuiStore{DB: db}
}
func TestMigrateCreatesDefaultDataOnce(t *testing.T) {
store := newTestStore(t)
if err := Migrate(store.DB); err != nil {
t.Fatalf("second migrate: %v", err)
}
var hotkeyCount int
if err := store.DB.QueryRow("SELECT COUNT(*) FROM hotkeys").Scan(&hotkeyCount); err != nil {
t.Fatalf("count hotkeys: %v", err)
}
if hotkeyCount != 2 {
t.Fatalf("hotkey count = %d, want 2", hotkeyCount)
}
var lang string
if err := store.DB.QueryRow("SELECT value FROM appconfig WHERE key = 'language'").Scan(&lang); err != nil {
t.Fatalf("query language: %v", err)
}
if lang != "zh" {
t.Fatalf("language = %q, want zh", lang)
}
}
func TestAppConfigServiceLanguage(t *testing.T) {
store := newTestStore(t)
service := NewAppConfigService(store)
if got := service.GetLanguage(); got != "zh" {
t.Fatalf("default language = %q, want zh", got)
}
if err := service.SetLanguage("en"); err != nil {
t.Fatalf("set language: %v", err)
}
if got := service.GetLanguage(); got != "en" {
t.Fatalf("updated language = %q, want en", got)
}
}
2026-07-08 23:37:19 +08:00
func TestAuthServiceLocalLoginLifecycle(t *testing.T) {
store := newTestStore(t)
service := NewAuthService(store)
session, err := service.Login("local", "admin", "admin123")
if err != nil {
t.Fatalf("login: %v", err)
}
if !session.Active {
t.Fatal("session is not active")
}
if session.Username != "admin" {
t.Fatalf("username = %q, want admin", session.Username)
}
if session.Provider != "local" {
t.Fatalf("provider = %q, want local", session.Provider)
}
if session.Token == "" {
t.Fatal("session token is empty")
}
current, err := service.GetCurrentSession()
if err != nil {
t.Fatalf("get current session: %v", err)
}
if current.ID != session.ID {
t.Fatalf("current session id = %d, want %d", current.ID, session.ID)
}
var loginAuditCount int
if err := store.DB.QueryRow("SELECT COUNT(*) FROM audit_logs WHERE action = 'auth.login.success' AND actor = 'admin'").Scan(&loginAuditCount); err != nil {
t.Fatalf("count login audit logs: %v", err)
}
if loginAuditCount != 1 {
t.Fatalf("login audit count = %d, want 1", loginAuditCount)
}
if err := service.Logout(); err != nil {
t.Fatalf("logout: %v", err)
}
current, err = service.GetCurrentSession()
if err != nil {
t.Fatalf("get current session after logout: %v", err)
}
if current.Active {
t.Fatal("current session is active after logout")
}
var logoutAuditCount int
if err := store.DB.QueryRow("SELECT COUNT(*) FROM audit_logs WHERE action = 'auth.logout' AND actor = 'admin'").Scan(&logoutAuditCount); err != nil {
t.Fatalf("count logout audit logs: %v", err)
}
if logoutAuditCount != 1 {
t.Fatalf("logout audit count = %d, want 1", logoutAuditCount)
}
}
func TestAuthServiceLocalLoginFailureWritesAuditLog(t *testing.T) {
store := newTestStore(t)
service := NewAuthService(store)
_, err := service.Login("local", "admin", "bad-password")
if !errors.Is(err, errInvalidCredentials) {
t.Fatalf("login error = %v, want invalid credentials", err)
}
var auditCount int
if err := store.DB.QueryRow("SELECT COUNT(*) FROM audit_logs WHERE action = 'auth.login.failure' AND actor = 'admin'").Scan(&auditCount); err != nil {
t.Fatalf("count audit logs: %v", err)
}
if auditCount != 1 {
t.Fatalf("audit count = %d, want 1", auditCount)
}
}
func TestAuthServiceRemoteProviderReserved(t *testing.T) {
store := newTestStore(t)
service := NewAuthService(store)
_, err := service.Login("remote", "admin", "admin123")
if !errors.Is(err, errRemoteProviderNotReady) {
t.Fatalf("login error = %v, want remote provider not ready", err)
}
var auditCount int
if err := store.DB.QueryRow("SELECT COUNT(*) FROM audit_logs WHERE action = 'auth.login.failure' AND target_id = 'remote'").Scan(&auditCount); err != nil {
t.Fatalf("count audit logs: %v", err)
}
if auditCount != 1 {
t.Fatalf("audit count = %d, want 1", auditCount)
}
}