diff --git a/backend/internal/store/user.go b/backend/internal/store/user.go index 40cb1fd..b743875 100644 --- a/backend/internal/store/user.go +++ b/backend/internal/store/user.go @@ -68,35 +68,35 @@ func (s *Store) GetPublicUser(ctx context.Context, id uuid.UUID) (*m.PublicUser, return &u, nil } -func (s *Store) Upsert(ctx context.Context, subjectID, email, name string) (uuid.UUID, error) { +func (s *Store) Upsert(ctx context.Context, subjectID, email, name string) (*m.User, error) { rootSpan := otel.GetRootSpan(ctx) - var id uuid.UUID + var user m.User // ON CONFLICT (subject_id) tells Postgres: "If this OIDC user already exists, // just update their info and last_login instead of throwing an error." query := ` INSERT INTO users (subject_id, email, display_name, last_login) VALUES ($1, $2, $3, NOW()) - ON CONFLICT (subject_id) DO UPDATE - SET last_login = NOW(), + ON CONFLICT (subject_id) DO UPDATE + SET last_login = NOW(), display_name = EXCLUDED.display_name, email = EXCLUDED.email - RETURNING id; + RETURNING id, subject_id, email, display_name; ` - err := s.db.QueryRowContext(ctx, query, subjectID, email, name).Scan(&id) + err := s.db.QueryRowxContext(ctx, query, subjectID, email, name).StructScan(&user) if err != nil { rootSpan.RecordError(err) rootSpan.SetStatus(codes.Error, "user_upsert_failed") - return uuid.Nil, fmt.Errorf("failed to upsert user: %w", err) + return nil, fmt.Errorf("failed to upsert user: %w", err) } rootSpan.SetAttributes( - attribute.String("user.id", id.String()), + attribute.String("user.id", user.ID.String()), ) - return id, nil + return &user, nil } func (s *Store) ListUsers(ctx context.Context) ([]models.PublicUser, error) { diff --git a/backend/internal/store/user_test.go b/backend/internal/store/user_test.go new file mode 100644 index 0000000..78f11d6 --- /dev/null +++ b/backend/internal/store/user_test.go @@ -0,0 +1,108 @@ +package store + +import ( + "testing" + + "github.com/google/uuid" +) + +func TestUpsertUser(t *testing.T) { + clearDatabase(t) + s := NewStore(globalTestDB) + + user, err := s.Upsert(globalTestCtx, "00001", "test@test.com", "Test User") + t.Run("new user", func(t *testing.T) { + if err != nil { + t.Errorf("failed to upsert user: %s", err) + } + }) + + t.Run("duplicate user", func(t *testing.T) { + duplicate_user, err := s.Upsert(globalTestCtx, "00001", "test2@test.com", "Test User123") + if err != nil { + t.Errorf("failed to upsert user (conflict): %s", err) + } + + if user.ID != duplicate_user.ID { + t.Errorf("expected duplicate upsert to return same user_id, got %s and %s", user.ID, duplicate_user.ID) + } + + if user.Email == duplicate_user.Email { + t.Errorf("expected duplicate upsert to update email, got %s and %s", user.Email, duplicate_user.Email) + } + + if user.DisplayName == duplicate_user.DisplayName { + t.Errorf("expected duplicate upsert to update display_name, got %s and %s", user.DisplayName, duplicate_user.DisplayName) + } + }) +} + +func TestGetUser(t *testing.T) { + clearDatabase(t) + s := NewStore(globalTestDB) + user_id, _ := s.Upsert(globalTestCtx, "00001", "test@test.com", "Test User") + + t.Run("existing user", func(t *testing.T) { + _, err := s.GetUser(globalTestCtx, user_id.ID) + if err != nil { + t.Errorf("failed to get user: %s", err) + } + }) + + t.Run("non-existent user", func(t *testing.T) { + user, err := s.GetUser(globalTestCtx, uuid.New()) + if err == nil || user != nil { + t.Errorf("expected error for non-existent user, got: %v", err) + } + }) +} + +func TestGetPublicUser(t *testing.T) { + clearDatabase(t) + s := NewStore(globalTestDB) + user_id, _ := s.Upsert(globalTestCtx, "00001", "test@test.com", "Test User") + + t.Run("existing user", func(t *testing.T) { + _, err := s.GetPublicUser(globalTestCtx, user_id.ID) + if err != nil { + t.Errorf("failed to get user: %s", err) + } + }) + + t.Run("non-existent user", func(t *testing.T) { + user, err := s.GetPublicUser(globalTestCtx, uuid.New()) + if err == nil || user != nil { + t.Errorf("expected error for non-existent user, got: %v", err) + } + }) +} + +func TestListUsers(t *testing.T) { + clearDatabase(t) + s := NewStore(globalTestDB) + + t.Run("empty user list", func(t *testing.T) { + users, err := s.ListUsers(globalTestCtx) + if users != nil { + t.Errorf("user list should be empty. probably poisoned tests: %v", err) + } + }) + + t.Run("filled user list", func(t *testing.T) { + s.Upsert(globalTestCtx, "1", "test1@test.com", "Test1") + s.Upsert(globalTestCtx, "2", "test2@test.com", "Test2") + s.Upsert(globalTestCtx, "3", "test3@test.com", "Test3") + + users, err := s.ListUsers(globalTestCtx) + if err != nil { + t.Errorf("Error while querying user list: %v", err) + } + if users == nil { + t.Error("user list should not be empty.") + } + + if users[0].DisplayName != "Test1" { + t.Errorf("user name should be Test1 but is %v", users[0].DisplayName) + } + }) +}