293 lines
10 KiB
Python
293 lines
10 KiB
Python
from __future__ import annotations
|
|
|
|
import uuid
|
|
from datetime import UTC, datetime, timedelta
|
|
|
|
import pytest
|
|
from sqlalchemy.orm import Session
|
|
|
|
from app.models.user import RefreshToken, Role, User
|
|
from app.repositories.user_repository import (
|
|
AuditLogRepository,
|
|
RefreshTokenRepository,
|
|
RoleRepository,
|
|
UserRepository,
|
|
)
|
|
from tests.conftest import create_test_role, create_test_user
|
|
|
|
|
|
class TestUserRepository:
|
|
"""Tests for UserRepository CRUD operations."""
|
|
|
|
def test_create_user(self, db: Session) -> None:
|
|
repo = UserRepository(db)
|
|
user = repo.create_user(
|
|
username="repo_user",
|
|
email="repo@test.com",
|
|
hashed_password="$2b$12$fakehash",
|
|
full_name="Repo User",
|
|
)
|
|
assert user.id is not None
|
|
assert user.username == "repo_user"
|
|
assert user.email == "repo@test.com"
|
|
assert user.is_active is True
|
|
assert user.is_superuser is False
|
|
|
|
def test_get_by_username(self, db: Session) -> None:
|
|
user = create_test_user(db, username="findme")
|
|
repo = UserRepository(db)
|
|
found = repo.get_by_username("findme")
|
|
assert found is not None
|
|
assert found.id == user.id
|
|
|
|
def test_get_by_username_not_found(self, db: Session) -> None:
|
|
repo = UserRepository(db)
|
|
assert repo.get_by_username("nonexistent") is None
|
|
|
|
def test_get_by_email(self, db: Session) -> None:
|
|
user = create_test_user(db, email="email@test.com")
|
|
repo = UserRepository(db)
|
|
found = repo.get_by_email("email@test.com")
|
|
assert found is not None
|
|
assert found.id == user.id
|
|
|
|
def test_get_by_email_not_found(self, db: Session) -> None:
|
|
repo = UserRepository(db)
|
|
assert repo.get_by_email("nobody@test.com") is None
|
|
|
|
def test_get_by_id(self, db: Session) -> None:
|
|
user = create_test_user(db)
|
|
repo = UserRepository(db)
|
|
found = repo.get_by_id(user.id)
|
|
assert found is not None
|
|
assert found.username == user.username
|
|
|
|
def test_get_by_id_not_found(self, db: Session) -> None:
|
|
repo = UserRepository(db)
|
|
assert repo.get_by_id(uuid.uuid4()) is None
|
|
|
|
def test_update_last_login(self, db: Session) -> None:
|
|
user = create_test_user(db)
|
|
assert user.last_login is None
|
|
repo = UserRepository(db)
|
|
updated = repo.update_last_login(user)
|
|
assert updated.last_login is not None
|
|
|
|
def test_get_active_users(self, db: Session) -> None:
|
|
create_test_user(db, username="active1", is_active=True)
|
|
create_test_user(db, username="active2", is_active=True)
|
|
create_test_user(db, username="inactive1", is_active=False)
|
|
|
|
repo = UserRepository(db)
|
|
active = repo.get_active_users()
|
|
usernames = [u.username for u in active]
|
|
assert "active1" in usernames
|
|
assert "active2" in usernames
|
|
assert "inactive1" not in usernames
|
|
|
|
def test_create_user_with_roles(self, db: Session) -> None:
|
|
create_test_role(db, name="admin")
|
|
create_test_role(db, name="user")
|
|
|
|
repo = UserRepository(db)
|
|
user = repo.create_user(
|
|
username="roled_user",
|
|
email="roled@test.com",
|
|
hashed_password="$2b$12$fakehash",
|
|
role_names=["admin", "user"],
|
|
)
|
|
role_names = [r.name for r in user.roles]
|
|
assert "admin" in role_names
|
|
assert "user" in role_names
|
|
|
|
def test_assign_roles(self, db: Session) -> None:
|
|
create_test_role(db, name="viewer")
|
|
user = create_test_user(db)
|
|
repo = UserRepository(db)
|
|
updated = repo.assign_roles(user, ["viewer"])
|
|
role_names = [r.name for r in updated.roles]
|
|
assert "viewer" in role_names
|
|
|
|
def test_delete_user(self, db: Session) -> None:
|
|
user = create_test_user(db)
|
|
repo = UserRepository(db)
|
|
assert repo.delete_by_id(user.id) is True
|
|
assert repo.get_by_id(user.id) is None
|
|
|
|
def test_delete_nonexistent_user(self, db: Session) -> None:
|
|
repo = UserRepository(db)
|
|
assert repo.delete_by_id(uuid.uuid4()) is False
|
|
|
|
def test_exists(self, db: Session) -> None:
|
|
user = create_test_user(db)
|
|
repo = UserRepository(db)
|
|
assert repo.exists(user.id) is True
|
|
assert repo.exists(uuid.uuid4()) is False
|
|
|
|
def test_count(self, db: Session) -> None:
|
|
create_test_user(db, username="count1")
|
|
create_test_user(db, username="count2")
|
|
repo = UserRepository(db)
|
|
assert repo.count() >= 2
|
|
|
|
|
|
class TestRoleRepository:
|
|
"""Tests for RoleRepository."""
|
|
|
|
def test_create_role(self, db: Session) -> None:
|
|
repo = RoleRepository(db)
|
|
role = repo.create_role(name="editor", description="Can edit documents")
|
|
assert role.id is not None
|
|
assert role.name == "editor"
|
|
|
|
def test_get_by_name(self, db: Session) -> None:
|
|
create_test_role(db, name="tester")
|
|
repo = RoleRepository(db)
|
|
found = repo.get_by_name("tester")
|
|
assert found is not None
|
|
assert found.name == "tester"
|
|
|
|
def test_get_by_name_not_found(self, db: Session) -> None:
|
|
repo = RoleRepository(db)
|
|
assert repo.get_by_name("nonexistent_role") is None
|
|
|
|
def test_get_all_roles(self, db: Session) -> None:
|
|
create_test_role(db, name="role_a")
|
|
create_test_role(db, name="role_b")
|
|
repo = RoleRepository(db)
|
|
roles = repo.get_all_roles()
|
|
names = [r.name for r in roles]
|
|
assert "role_a" in names
|
|
assert "role_b" in names
|
|
|
|
|
|
class TestRefreshTokenRepository:
|
|
"""Tests for RefreshTokenRepository."""
|
|
|
|
def test_create_token(self, db: Session) -> None:
|
|
user = create_test_user(db)
|
|
repo = RefreshTokenRepository(db)
|
|
expires = datetime.now(UTC) + timedelta(days=7)
|
|
token = repo.create_token(
|
|
user_id=user.id,
|
|
token="test_refresh_token_abc",
|
|
expires_at=expires,
|
|
)
|
|
assert token.id is not None
|
|
assert token.user_id == user.id
|
|
assert token.revoked is False
|
|
|
|
def test_get_by_token(self, db: Session) -> None:
|
|
user = create_test_user(db)
|
|
repo = RefreshTokenRepository(db)
|
|
expires = datetime.now(UTC) + timedelta(days=7)
|
|
repo.create_token(user_id=user.id, token="findable_token", expires_at=expires)
|
|
|
|
found = repo.get_by_token("findable_token")
|
|
assert found is not None
|
|
assert found.user_id == user.id
|
|
|
|
def test_get_by_token_not_found(self, db: Session) -> None:
|
|
repo = RefreshTokenRepository(db)
|
|
assert repo.get_by_token("nonexistent_token") is None
|
|
|
|
def test_get_by_revoked_token_returns_none(self, db: Session) -> None:
|
|
user = create_test_user(db)
|
|
repo = RefreshTokenRepository(db)
|
|
expires = datetime.now(UTC) + timedelta(days=7)
|
|
repo.create_token(user_id=user.id, token="revoked_token", expires_at=expires)
|
|
repo.revoke_token("revoked_token")
|
|
db.flush()
|
|
|
|
assert repo.get_by_token("revoked_token") is None
|
|
|
|
def test_revoke_token(self, db: Session) -> None:
|
|
user = create_test_user(db)
|
|
repo = RefreshTokenRepository(db)
|
|
expires = datetime.now(UTC) + timedelta(days=7)
|
|
repo.create_token(user_id=user.id, token="to_revoke", expires_at=expires)
|
|
|
|
assert repo.revoke_token("to_revoke") is True
|
|
assert repo.get_by_token("to_revoke") is None
|
|
|
|
def test_revoke_nonexistent_token(self, db: Session) -> None:
|
|
repo = RefreshTokenRepository(db)
|
|
assert repo.revoke_token("does_not_exist") is False
|
|
|
|
def test_revoke_all_user_tokens(self, db: Session) -> None:
|
|
user = create_test_user(db)
|
|
repo = RefreshTokenRepository(db)
|
|
expires = datetime.now(UTC) + timedelta(days=7)
|
|
repo.create_token(user_id=user.id, token="token_1", expires_at=expires)
|
|
repo.create_token(user_id=user.id, token="token_2", expires_at=expires)
|
|
repo.create_token(user_id=user.id, token="token_3", expires_at=expires)
|
|
|
|
count = repo.revoke_all_user_tokens(user.id)
|
|
assert count == 3
|
|
assert repo.get_by_token("token_1") is None
|
|
assert repo.get_by_token("token_2") is None
|
|
assert repo.get_by_token("token_3") is None
|
|
|
|
def test_cleanup_expired_tokens(self, db: Session) -> None:
|
|
user = create_test_user(db)
|
|
repo = RefreshTokenRepository(db)
|
|
|
|
# Create expired token
|
|
expired = datetime.now(UTC) - timedelta(days=1)
|
|
repo.create_token(user_id=user.id, token="expired_tok", expires_at=expired)
|
|
|
|
# Create valid token
|
|
valid = datetime.now(UTC) + timedelta(days=7)
|
|
repo.create_token(user_id=user.id, token="valid_tok", expires_at=valid)
|
|
|
|
count = repo.cleanup_expired_tokens()
|
|
assert count >= 1
|
|
|
|
|
|
class TestAuditLogRepository:
|
|
"""Tests for AuditLogRepository."""
|
|
|
|
def test_log_action(self, db: Session) -> None:
|
|
user = create_test_user(db)
|
|
repo = AuditLogRepository(db)
|
|
log = repo.log_action(
|
|
action="login",
|
|
resource_type="auth",
|
|
user_id=user.id,
|
|
ip_address="127.0.0.1",
|
|
user_agent="TestAgent/1.0",
|
|
)
|
|
assert log.id is not None
|
|
assert log.action == "login"
|
|
assert log.resource_type == "auth"
|
|
|
|
def test_log_action_without_user(self, db: Session) -> None:
|
|
repo = AuditLogRepository(db)
|
|
log = repo.log_action(
|
|
action="anonymous_access",
|
|
resource_type="documents",
|
|
)
|
|
assert log.id is not None
|
|
assert log.user_id is None
|
|
|
|
def test_get_user_logs(self, db: Session) -> None:
|
|
user = create_test_user(db)
|
|
repo = AuditLogRepository(db)
|
|
repo.log_action(action="view", resource_type="documents", user_id=user.id)
|
|
repo.log_action(action="edit", resource_type="templates", user_id=user.id)
|
|
|
|
logs = repo.get_user_logs(user.id)
|
|
assert len(logs) >= 2
|
|
|
|
def test_get_resource_logs(self, db: Session) -> None:
|
|
repo = AuditLogRepository(db)
|
|
resource_id = str(uuid.uuid4())
|
|
repo.log_action(action="create", resource_type="documents", resource_id=resource_id)
|
|
repo.log_action(action="update", resource_type="documents", resource_id=resource_id)
|
|
|
|
logs = repo.get_resource_logs("documents", resource_id)
|
|
assert len(logs) >= 2
|
|
for log in logs:
|
|
assert log.resource_type == "documents"
|
|
assert log.resource_id == resource_id
|