diff --git a/tests/unit/test_email_service_smtp.py b/tests/unit/test_email_service_smtp.py new file mode 100755 index 000000000..718edecc0 --- /dev/null +++ b/tests/unit/test_email_service_smtp.py @@ -0,0 +1,272 @@ +"""Email Service (SMTP) 单元测试""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest + +from packages.adapters.smtp.email_service import ( + EmailService, + NoopEmailService, + get_email_service, +) +from packages.domain.auth.email_service import EmailConfig + + +@pytest.fixture +def email_config(): + return EmailConfig( + smtp_host="smtp.example.com", + smtp_port=587, + from_email="noreply@example.com", + from_name="小虾 SaaS", + smtp_user="user", + smtp_password="pass", + use_tls=True, + ) + + +@pytest.fixture +def email_service(email_config): + return EmailService(email_config) + + +class TestNoopEmailService: + """NoopEmailService 测试""" + + def test_send_verification_email_returns_false(self): + """验证邮件返回失败""" + svc = NoopEmailService() + success, msg = svc.send_verification_email( + to_email="test@example.com", + username="testuser", + verification_url="https://example.com/verify", + ) + assert success is False + assert "disabled" in msg.lower() + + def test_send_password_reset_email_returns_false(self): + """密码重置邮件返回失败""" + svc = NoopEmailService() + success, msg = svc.send_password_reset_email( + to_email="test@example.com", + username="testuser", + reset_url="https://example.com/reset", + ) + assert success is False + assert "disabled" in msg.lower() + + +class TestEmailServiceInit: + """初始化测试""" + + def test_init_with_config(self, email_config): + """使用指定配置初始化""" + svc = EmailService(email_config) + assert svc.config is email_config + + def test_init_without_config(self): + """不指定配置使用默认 EmailConfig""" + svc = EmailService() + assert svc.config is not None + assert isinstance(svc.config, EmailConfig) + + +class TestSendEmail: + """send_email 方法测试""" + + def test_send_success(self, email_service): + """发送成功""" + with patch("packages.adapters.smtp.email_service.smtplib.SMTP") as mock_smtp: + mock_server = MagicMock() + mock_smtp.return_value.__enter__.return_value = mock_server + + success, error = email_service.send_email( + to_email="user@example.com", + subject="测试主题", + html_body="
测试内容
", + ) + + assert success is True + assert error is None + mock_server.starttls.assert_called_once() + mock_server.login.assert_called_once_with("user", "pass") + mock_server.sendmail.assert_called_once() + + def test_send_without_tls(self, email_config): + """不使用 TLS""" + email_config.use_tls = False + svc = EmailService(email_config) + + with patch("packages.adapters.smtp.email_service.smtplib.SMTP") as mock_smtp: + mock_server = MagicMock() + mock_smtp.return_value.__enter__.return_value = mock_server + + svc.send_email(to_email="u@e.com", subject="s", html_body="body") + mock_server.starttls.assert_not_called() + + def test_send_without_auth(self, email_config): + """不配置用户名密码时不登录""" + email_config.smtp_user = "" + email_config.smtp_password = "" + svc = EmailService(email_config) + + with patch("packages.adapters.smtp.email_service.smtplib.SMTP") as mock_smtp: + mock_server = MagicMock() + mock_smtp.return_value.__enter__.return_value = mock_server + + svc.send_email(to_email="u@e.com", subject="s", html_body="body") + mock_server.login.assert_not_called() + + def test_send_with_cc_and_bcc(self, email_service): + """发送带抄送和密送""" + with patch("packages.adapters.smtp.email_service.smtplib.SMTP") as mock_smtp: + mock_server = MagicMock() + mock_smtp.return_value.__enter__.return_value = mock_server + + email_service.send_email( + to_email="to@example.com", + subject="s", + html_body="body", + cc=["cc1@example.com", "cc2@example.com"], + bcc=["bcc@example.com"], + ) + + # 验证 recipients 包含所有收件人 + call_args = mock_server.sendmail.call_args + recipients = call_args[0][1] + assert "to@example.com" in recipients + assert "cc1@example.com" in recipients + assert "cc2@example.com" in recipients + assert "bcc@example.com" in recipients + + def test_send_with_text_body(self, email_service): + """带纯文本正文""" + with patch("packages.adapters.smtp.email_service.smtplib.SMTP") as mock_smtp: + mock_server = MagicMock() + mock_smtp.return_value.__enter__.return_value = mock_server + + email_service.send_email( + to_email="u@e.com", + subject="s", + html_body="html
", + text_body="plain text", + ) + + mock_server.sendmail.assert_called_once() + + def test_send_failure_returns_false(self, email_service): + """发送失败返回 False 和错误信息""" + with patch("packages.adapters.smtp.email_service.smtplib.SMTP") as mock_smtp: + mock_server = MagicMock() + mock_server.sendmail.side_effect = Exception("Connection refused") + mock_smtp.return_value.__enter__.return_value = mock_server + + success, error = email_service.send_email( + to_email="u@e.com", subject="s", html_body="body" + ) + + assert success is False + assert "Connection refused" in error + + def test_from_header_exists(self, email_service): + """From 头存在""" + with patch("packages.adapters.smtp.email_service.smtplib.SMTP") as mock_smtp: + mock_server = MagicMock() + mock_smtp.return_value.__enter__.return_value = mock_server + + email_service.send_email(to_email="u@e.com", subject="s", html_body="body") + + call_args = mock_server.sendmail.call_args + msg_str = call_args[0][2] + assert "From:" in msg_str + assert "To: u@e.com" in msg_str + + +class TestSendVerificationEmail: + """发送验证邮件测试""" + + def test_verification_email_contains_url(self, email_service): + """验证邮件包含验证链接""" + with patch.object(email_service, "send_email", return_value=(True, None)) as mock_send: + email_service.send_verification_email( + to_email="user@example.com", + username="testuser", + verification_url="https://app.example.com/verify?token=abc123", + ) + + mock_send.assert_called_once() + call_args = mock_send.call_args + # 验证主题 + assert "验证" in call_args[0][1] + # HTML 正文包含用户名和链接 + assert "testuser" in call_args[0][2] + assert "https://app.example.com/verify?token=abc123" in call_args[0][2] + + def test_verification_email_has_text_body(self, email_service): + """验证邮件有纯文本版""" + with patch.object(email_service, "send_email", return_value=(True, None)) as mock_send: + email_service.send_verification_email( + to_email="u@e.com", + username="u", + verification_url="https://example.com/v", + ) + + call_args = mock_send.call_args + # 第四个参数是 text_body + assert call_args[0][3] is not None + assert len(call_args[0][3]) > 0 + + +class TestSendPasswordResetEmail: + """发送密码重置邮件测试""" + + def test_reset_email_contains_url(self, email_service): + """重置邮件包含重置链接""" + with patch.object(email_service, "send_email", return_value=(True, None)) as mock_send: + email_service.send_password_reset_email( + to_email="user@example.com", + username="testuser", + reset_url="https://app.example.com/reset?token=xyz", + ) + + mock_send.assert_called_once() + call_args = mock_send.call_args + assert "重置" in call_args[0][1] + assert "testuser" in call_args[0][2] + assert "https://app.example.com/reset?token=xyz" in call_args[0][2] + + def test_reset_email_has_text_body(self, email_service): + """重置邮件有纯文本版""" + with patch.object(email_service, "send_email", return_value=(True, None)) as mock_send: + email_service.send_password_reset_email( + to_email="u@e.com", + username="u", + reset_url="https://example.com/r", + ) + + call_args = mock_send.call_args + assert call_args[0][3] is not None + assert len(call_args[0][3]) > 0 + + +class TestGetEmailService: + """工厂函数测试""" + + def test_disabled_returns_noop(self): + """禁用时返回 NoopEmailService""" + svc = get_email_service(enabled=False) + assert isinstance(svc, NoopEmailService) + + def test_enabled_returns_email_service(self, email_config): + """启用时返回 EmailService""" + svc = get_email_service(config=email_config, enabled=True) + assert isinstance(svc, EmailService) + + def test_singleton_default(self): + """默认情况下是单例""" + svc1 = get_email_service() + svc2 = get_email_service() + # 两个都可能是 Noop 或 EmailService,取决于环境 + assert type(svc1) == type(svc2) diff --git a/tests/unit/test_redis_session_store.py b/tests/unit/test_redis_session_store.py new file mode 100755 index 000000000..c26223952 --- /dev/null +++ b/tests/unit/test_redis_session_store.py @@ -0,0 +1,378 @@ +"""Redis Session Store 单元测试""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest + +from packages.adapters.redis.session_store import ( + NoopSessionStore, + RedisConfig, + SessionStore, + get_session_store, +) + + +@pytest.fixture +def mock_redis(): + return MagicMock() + + +@pytest.fixture +def session_store(mock_redis): + return SessionStore(redis_client=mock_redis) + + +class TestRedisConfig: + """RedisConfig 默认值测试""" + + def test_default_values(self): + """默认配置""" + cfg = RedisConfig() + assert cfg.HOST == "localhost" + assert cfg.PORT == 6379 + assert cfg.DB == 0 + assert cfg.PASSWORD is None + assert cfg.DECODE_RESPONSES is True + + +class TestNoopSessionStore: + """NoopSessionStore 测试""" + + def test_save_session_returns_false(self): + """保存返回 False""" + store = NoopSessionStore() + assert store.save_session() is False + + def test_get_session_returns_none(self): + """获取返回 None""" + store = NoopSessionStore() + assert store.get_session("sess_123") is None + + def test_get_session_by_refresh_token_returns_none(self): + """通过 refresh_token 获取返回 None""" + store = NoopSessionStore() + assert store.get_session_by_refresh_token("tok_123") is None + + def test_get_refresh_token_returns_none(self): + """获取 refresh_token 返回 None""" + store = NoopSessionStore() + assert store.get_refresh_token("sess_123") is None + + def test_update_last_active_returns_false(self): + """更新活跃时间返回 False""" + store = NoopSessionStore() + assert store.update_last_active("sess_123") is False + + def test_delete_session_returns_false(self): + """删除返回 False""" + store = NoopSessionStore() + assert store.delete_session("sess_123") is False + + def test_get_user_sessions_returns_empty(self): + """用户 session 列表为空""" + store = NoopSessionStore() + assert store.get_user_sessions("user_1") == [] + + def test_delete_all_user_sessions_returns_zero(self): + """删除所有返回 0""" + store = NoopSessionStore() + assert store.delete_all_user_sessions("user_1") == 0 + + def test_session_exists_returns_false(self): + """存在性检查返回 False""" + store = NoopSessionStore() + assert store.session_exists("sess_123") is False + + +class TestSessionStoreInit: + """SessionStore 初始化测试""" + + def test_init_with_redis_client(self, mock_redis): + """使用注入的 redis client""" + store = SessionStore(redis_client=mock_redis) + assert store.redis is mock_redis + + def test_init_with_config(self): + """使用配置创建 redis client""" + cfg = RedisConfig() + cfg.HOST = "redis.example.com" + cfg.PORT = 6380 + + with patch("packages.adapters.redis.session_store.redis.Redis") as mock_redis_cls: + store = SessionStore(config=cfg) + mock_redis_cls.assert_called_once_with( + host="redis.example.com", + port=6380, + db=0, + password=None, + decode_responses=True, + ) + + +class TestSessionStoreSave: + """save_session 测试""" + + def test_save_success(self, session_store, mock_redis): + """保存成功""" + result = session_store.save_session( + session_id="sess_001", + user_id="user_001", + refresh_token="refresh_abc", + device_info="Chrome/Windows", + ip_address="192.168.1.1", + ) + + assert result is True + # 验证 session 数据保存 + mock_redis.setex.assert_any_call( + "session:sess_001", 30 * 24 * 60 * 60, mock_redis.setex.call_args_list[0][0][2] + ) + # 验证 refresh_token 保存 + mock_redis.setex.assert_any_call( + "refresh_token:sess_001", 30 * 24 * 60 * 60, "refresh_abc" + ) + # 验证反向映射 + mock_redis.setex.assert_any_call( + "refresh_token_map:refresh_abc", 30 * 24 * 60 * 60, "sess_001" + ) + # 验证用户集合 + mock_redis.sadd.assert_called_once_with("user_sessions:user_001", "sess_001") + mock_redis.expire.assert_called_once() + + def test_save_custom_expiry(self, session_store, mock_redis): + """自定义过期时间""" + session_store.save_session( + session_id="sess_001", + user_id="user_001", + refresh_token="tok", + device_info="d", + ip_address="1.1.1.1", + expires_in_seconds=3600, + ) + + # 验证 TTL 为 3600 + call_args = mock_redis.setex.call_args_list[0] + assert call_args[0][1] == 3600 + + def test_save_returns_false_on_error(self, session_store, mock_redis): + """Redis 异常时返回 False""" + mock_redis.setex.side_effect = Exception("Connection error") + + result = session_store.save_session( + session_id="s1", user_id="u1", refresh_token="t1", + device_info="d", ip_address="1.1.1.1", + ) + assert result is False + + +class TestSessionStoreGet: + """get_session 测试""" + + def test_get_existing_session(self, session_store, mock_redis): + """获取存在的 session""" + import json + session_data = { + "session_id": "sess_001", + "user_id": "user_001", + "device_info": "Chrome", + "ip_address": "1.1.1.1", + } + mock_redis.get.return_value = json.dumps(session_data) + + result = session_store.get_session("sess_001") + assert result is not None + assert result["user_id"] == "user_001" + assert result["session_id"] == "sess_001" + mock_redis.get.assert_called_once_with("session:sess_001") + + def test_get_nonexistent_session(self, session_store, mock_redis): + """获取不存在的 session 返回 None""" + mock_redis.get.return_value = None + + result = session_store.get_session("nonexistent") + assert result is None + + def test_get_returns_none_on_error(self, session_store, mock_redis): + """Redis 异常返回 None""" + mock_redis.get.side_effect = Exception("error") + result = session_store.get_session("s1") + assert result is None + + +class TestSessionStoreGetByRefreshToken: + """get_session_by_refresh_token 测试""" + + def test_get_by_refresh_token_success(self, session_store, mock_redis): + """通过 refresh_token 获取成功""" + import json + session_data = {"session_id": "sess_001", "user_id": "user_001"} + + # 第一次调用(反向映射)返回 session_id + # 第二次调用(session数据)返回 json + mock_redis.get.side_effect = ["sess_001", json.dumps(session_data)] + + result = session_store.get_session_by_refresh_token("refresh_abc") + assert result is not None + assert result["session_id"] == "sess_001" + + def test_get_by_refresh_token_not_found(self, session_store, mock_redis): + """refresh_token 不存在返回 None""" + mock_redis.get.return_value = None + + result = session_store.get_session_by_refresh_token("invalid") + assert result is None + + +class TestSessionStoreGetRefreshToken: + """get_refresh_token 测试""" + + def test_get_refresh_token_success(self, session_store, mock_redis): + """获取 refresh_token 成功""" + mock_redis.get.return_value = "refresh_abc" + + result = session_store.get_refresh_token("sess_001") + assert result == "refresh_abc" + mock_redis.get.assert_called_once_with("refresh_token:sess_001") + + def test_get_refresh_token_not_found(self, session_store, mock_redis): + """不存在返回 None""" + mock_redis.get.return_value = None + assert session_store.get_refresh_token("sess_001") is None + + +class TestSessionStoreUpdateLastActive: + """update_last_active 测试""" + + def test_update_success(self, session_store, mock_redis): + """更新成功""" + import json + session_data = { + "session_id": "sess_001", + "user_id": "user_001", + "last_active_at": "2024-01-01T00:00:00+00:00", + } + mock_redis.get.return_value = json.dumps(session_data) + mock_redis.ttl.return_value = 1800 + + result = session_store.update_last_active("sess_001") + assert result is True + mock_redis.setex.assert_called_once() + + def test_update_session_not_found(self, session_store, mock_redis): + """session 不存在返回 False""" + mock_redis.get.return_value = None + result = session_store.update_last_active("nonexistent") + assert result is False + + def test_update_expired_session(self, session_store, mock_redis): + """已过期的 session 返回 False""" + import json + session_data = {"session_id": "s1", "user_id": "u1"} + mock_redis.get.return_value = json.dumps(session_data) + mock_redis.ttl.return_value = -2 # 已过期 + + result = session_store.update_last_active("s1") + assert result is False + + +class TestSessionStoreDelete: + """delete_session 测试""" + + def test_delete_success(self, session_store, mock_redis): + """删除成功""" + import json + session_data = {"session_id": "sess_001", "user_id": "user_001"} + mock_redis.get.side_effect = [ + json.dumps(session_data), # get_session + "refresh_abc", # get refresh_token + ] + + result = session_store.delete_session("sess_001") + assert result is True + # 删除 session、refresh_token、反向映射、从用户集合移除 + assert mock_redis.delete.call_count >= 3 + mock_redis.srem.assert_called_once_with("user_sessions:user_001", "sess_001") + + def test_delete_not_found(self, session_store, mock_redis): + """删除不存在的 session 返回 False""" + mock_redis.get.return_value = None + result = session_store.delete_session("nonexistent") + assert result is False + + +class TestSessionStoreUserSessions: + """用户 Session 列表测试""" + + def test_get_user_sessions(self, session_store, mock_redis): + """获取用户所有 session""" + import json + mock_redis.smembers.return_value = {"sess_001", "sess_002"} + + session1 = json.dumps({"session_id": "sess_001", "user_id": "u1"}) + session2 = json.dumps({"session_id": "sess_002", "user_id": "u1"}) + mock_redis.get.side_effect = [session1, session2] + + result = session_store.get_user_sessions("user_001") + assert len(result) == 2 + + def test_get_user_sessions_empty(self, session_store, mock_redis): + """用户无 session""" + mock_redis.smembers.return_value = set() + result = session_store.get_user_sessions("user_001") + assert result == [] + + def test_delete_all_user_sessions(self, session_store, mock_redis): + """删除用户所有 session""" + import json + mock_redis.smembers.return_value = {"sess_001", "sess_002"} + + session1 = json.dumps({"session_id": "sess_001", "user_id": "u1"}) + session2 = json.dumps({"session_id": "sess_002", "user_id": "u1"}) + # get 调用顺序: + # 1-2: get_user_sessions 中两个 session 的 get + # 3-4: delete sess_001 (get_session + get refresh_token) + # 5-6: delete sess_002 (get_session + get refresh_token) + mock_redis.get.side_effect = [ + session1, session2, # get_user_sessions + session1, "tok1", # delete sess_001 + session2, "tok2", # delete sess_002 + ] + + count = session_store.delete_all_user_sessions("user_001") + assert count == 2 + + def test_delete_all_empty_user(self, session_store, mock_redis): + """删除无 session 的用户""" + mock_redis.smembers.return_value = set() + count = session_store.delete_all_user_sessions("user_001") + assert count == 0 + + +class TestSessionStoreExists: + """session_exists 测试""" + + def test_exists_true(self, session_store, mock_redis): + """存在返回 True""" + mock_redis.exists.return_value = 1 + assert session_store.session_exists("sess_001") is True + + def test_exists_false(self, session_store, mock_redis): + """不存在返回 False""" + mock_redis.exists.return_value = 0 + assert session_store.session_exists("sess_001") is False + + def test_exists_error_returns_false(self, session_store, mock_redis): + """异常返回 False""" + mock_redis.exists.side_effect = Exception("error") + assert session_store.session_exists("sess_001") is False + + +class TestGetSessionStore: + """工厂函数测试""" + + def test_disabled_returns_noop(self): + """禁用返回 NoopSessionStore""" + store = get_session_store(enabled=False) + assert isinstance(store, NoopSessionStore)