diff --git a/digital-avatar-app/backend/services/takeover_service.py b/digital-avatar-app/backend/services/takeover_service.py index 2316c61..8a0d6ff 100644 --- a/digital-avatar-app/backend/services/takeover_service.py +++ b/digital-avatar-app/backend/services/takeover_service.py @@ -7,7 +7,7 @@ from typing import Optional import httpx from sqlalchemy.orm import Session -from models import Authorization +from models import Avatar, Authorization from services.boxim_client import BoxIMClient logger = logging.getLogger(__name__) @@ -33,15 +33,18 @@ class TakeoverService: self, owner_huihui_id: str, from_user_id: str ) -> Optional[Authorization]: """Check whether takeover is enabled for the given target user.""" + avatar = self.db.query(Avatar).filter(Avatar.owner_id == owner_huihui_id).first() + if not avatar: + return None + auth = ( self.db.query(Authorization) + .filter(Authorization.avatar_id == avatar.id) .filter(Authorization.target_id == from_user_id) .filter(Authorization.takeover_enabled == True) .first() ) - if auth and auth.takeover_enabled: - return auth - return None + return auth if auth and auth.takeover_enabled else None async def generate_reply(self, avatar_id: str, message: str) -> str: """Call the avatar chat endpoint to generate a reply.""" @@ -63,7 +66,13 @@ class TakeoverService: async def execute_takeover(self, auth: Authorization, message: dict) -> bool: """Execute takeover: generate a reply and send it as the owner via IM.""" try: - owner_huihui_id = auth.owner_id if hasattr(auth, "owner_id") else "" + # Resolve owner through Avatar model + avatar = self.db.query(Avatar).filter(Avatar.id == auth.avatar_id).first() + if not avatar: + logger.warning(f"Avatar not found: {auth.avatar_id}") + return False + + owner_huihui_id = avatar.owner_id credentials = await self.boxim.get_credentials(owner_huihui_id) if not credentials: logger.warning(f"Cannot obtain IM credentials for owner: {owner_huihui_id}") @@ -91,12 +100,15 @@ class TakeoverService: if not self.redis: logger.warning("Redis not configured, degrading to immediate takeover") return + + avatar = self.db.query(Avatar).filter(Avatar.id == auth.avatar_id).first() + owner_huihui_id = avatar.owner_id if avatar else "" key = f"takeover:delayed:{auth.target_id}:{message.get('msg_id', '')}" value = json.dumps({ "avatar_id": auth.avatar_id, "from_accid": message.get("from_accid", ""), "content": message.get("content", ""), - "owner_huihui_id": auth.owner_id if hasattr(auth, "owner_id") else "", + "owner_huihui_id": owner_huihui_id, }) self.redis.setex(key, auth.takeover_delay_seconds + 10, value) logger.info(f"Message enqueued to delayed queue: {key}") @@ -104,8 +116,31 @@ class TakeoverService: async def process_delayed_queue(self): """Process expired messages from the delayed queue. - Redis handles expiration automatically; this is a placeholder - for a future scheduled worker that scans and dispatches. + Scans Redis keys matching the takeover:delayed: pattern and dispatches + each to execute_takeover after resolving the Authorization. """ if not self.redis: return + try: + pattern = "takeover:delayed:*" + keys = self.redis.keys(pattern) + for key in keys: + raw = self.redis.get(key) + if not raw: + continue + data = json.loads(raw) + auth = ( + self.db.query(Authorization) + .filter(Authorization.target_id == key.split(":")[2]) + .first() + ) + if auth: + message = { + "msg_id": key.split(":")[-1], + "from_accid": data.get("from_accid", ""), + "content": data.get("content", ""), + } + await self.execute_takeover(auth, message) + self.redis.delete(key) + except Exception as e: + logger.error(f"Failed to process delayed queue: {e}") diff --git a/digital-avatar-app/backend/tests/test_takeover_service.py b/digital-avatar-app/backend/tests/test_takeover_service.py index 49afc11..f726f27 100644 --- a/digital-avatar-app/backend/tests/test_takeover_service.py +++ b/digital-avatar-app/backend/tests/test_takeover_service.py @@ -2,7 +2,7 @@ import pytest from unittest.mock import AsyncMock, patch, MagicMock from services.takeover_service import TakeoverService -from models import Authorization +from models import Authorization, Avatar @pytest.fixture @@ -27,26 +27,110 @@ def mock_auth(): auth.takeover_delay_seconds = 30 auth.avatar_id = "avatar_123" auth.target_id = "target_user_123" - auth.owner_id = "owner_huihui_123" return auth -def test_check_takeover_enabled_returns_auth_when_enabled(mock_db, mock_auth, mock_boxim): - mock_db.query.return_value.filter.return_value.filter.return_value.first.return_value = mock_auth +@pytest.fixture +def mock_avatar(): + avatar = MagicMock(spec=Avatar) + avatar.id = "avatar_123" + avatar.owner_id = "owner_huihui_123" + return avatar + + +# --- check_takeover_enabled --- + + +def test_check_takeover_enabled_returns_auth_when_enabled(mock_db, mock_auth, mock_boxim, mock_avatar): + avatar_filter = MagicMock() + avatar_filter.first.return_value = mock_avatar + avatar_query = MagicMock() + avatar_query.filter.return_value = avatar_filter + + auth_filter = MagicMock() + auth_filter.filter.return_value = auth_filter + auth_filter.first.return_value = mock_auth + + def query_side_effect(model): + if model == Avatar: + return avatar_query + return auth_filter + + mock_db.query.side_effect = query_side_effect + service = TakeoverService(mock_db, mock_boxim) - result = service.check_takeover_enabled("owner_123", "target_123") + result = service.check_takeover_enabled("owner_huihui_123", "target_user_123") assert result == mock_auth -def test_check_takeover_enabled_returns_none_when_disabled(mock_db, mock_boxim): - disabled_auth = MagicMock(spec=Authorization) - disabled_auth.takeover_enabled = False - mock_db.query.return_value.filter.return_value.filter.return_value.first.return_value = disabled_auth +def test_check_takeover_enabled_returns_none_when_no_avatar(mock_db, mock_boxim): + avatar_filter = MagicMock() + avatar_filter.first.return_value = None + avatar_query = MagicMock() + avatar_query.filter.return_value = avatar_filter + mock_db.query.return_value = avatar_query + service = TakeoverService(mock_db, mock_boxim) result = service.check_takeover_enabled("owner_123", "target_123") assert result is None +def test_check_takeover_enabled_returns_none_when_disabled(mock_db, mock_boxim, mock_avatar): + avatar_filter = MagicMock() + avatar_filter.first.return_value = mock_avatar + avatar_query = MagicMock() + avatar_query.filter.return_value = avatar_filter + + disabled_auth = MagicMock(spec=Authorization) + disabled_auth.takeover_enabled = False + auth_filter = MagicMock() + auth_filter.filter.return_value = auth_filter + auth_filter.first.return_value = disabled_auth + + def query_side_effect(model): + if model == Avatar: + return avatar_query + return auth_filter + + mock_db.query.side_effect = query_side_effect + + service = TakeoverService(mock_db, mock_boxim) + result = service.check_takeover_enabled("owner_123", "target_123") + assert result is None + + +def test_check_takeover_enabled_filters_by_owner_and_target(mock_db, mock_boxim, mock_avatar, mock_auth): + """Verify that queries use the correct filter arguments.""" + avatar_filter = MagicMock() + avatar_filter.first.return_value = mock_avatar + avatar_query = MagicMock() + avatar_query.filter.return_value = avatar_filter + + auth_filter = MagicMock() + auth_filter.filter.return_value = auth_filter + auth_filter.first.return_value = mock_auth + + call_order = [] + + def query_side_effect(model): + if model == Avatar: + call_order.append("Avatar") + return avatar_query + call_order.append("Authorization") + return auth_filter + + mock_db.query.side_effect = query_side_effect + + service = TakeoverService(mock_db, mock_boxim) + service.check_takeover_enabled("owner_huihui_123", "target_user_123") + + assert "Avatar" in call_order + assert "Authorization" in call_order + + +# --- generate_reply --- + + @pytest.mark.asyncio async def test_generate_reply_returns_answer(mock_boxim): mock_db = MagicMock() @@ -60,54 +144,6 @@ async def test_generate_reply_returns_answer(mock_boxim): assert result == "Hello back" -@pytest.mark.asyncio -async def test_execute_takeover_success(mock_db, mock_boxim, mock_auth): - with patch("services.takeover_service.httpx.AsyncClient") as mock_client_class: - mock_response = MagicMock() - mock_response.json.return_value = {"code": 200, "data": {"answer": "Hello back"}} - mock_client_class.return_value.__aenter__.return_value.post.return_value = mock_response - - service = TakeoverService(mock_db, mock_boxim) - message = {"from_accid": "user_acc", "content": "Hello"} - - result = await service.execute_takeover(mock_auth, message) - - assert result is True - mock_boxim.get_credentials.assert_called_once() - mock_boxim.send_p2p_message.assert_called_once() - - -def test_enqueue_delayed_message_with_redis(mock_db, mock_boxim, mock_auth): - mock_redis = MagicMock() - service = TakeoverService(mock_db, mock_boxim, mock_redis) - message = {"msg_id": "msg_1", "from_accid": "user_acc", "content": "Hello"} - - service.enqueue_delayed_message(mock_auth, message) - - mock_redis.setex.assert_called_once() - - -def test_enqueue_delayed_message_without_redis_logs_warning(mock_db, mock_boxim, mock_auth): - """When Redis is not configured, enqueue_delayed_message should log a warning and not crash.""" - service = TakeoverService(mock_db, mock_boxim) - message = {"msg_id": "msg_1", "from_accid": "user_acc", "content": "Hello"} - - # Should not raise - service.enqueue_delayed_message(mock_auth, message) - - -@pytest.mark.asyncio -async def test_execute_takeover_fails_when_no_credentials(mock_db, mock_boxim, mock_auth): - """execute_takeover should return False when boxim.get_credentials returns None.""" - mock_boxim.get_credentials.return_value = None - service = TakeoverService(mock_db, mock_boxim) - message = {"from_accid": "user_acc", "content": "Hello"} - - result = await service.execute_takeover(mock_auth, message) - - assert result is False - - @pytest.mark.asyncio async def test_generate_reply_handles_empty_answer(mock_boxim): """generate_reply should return empty string when answer is missing.""" @@ -134,3 +170,149 @@ async def test_generate_reply_handles_error_code(mock_boxim): service = TakeoverService(mock_db, mock_boxim) result = await service.generate_reply("avatar_123", "Hello") assert result == "" + + +# --- execute_takeover --- + + +@pytest.mark.asyncio +async def test_execute_takeover_success(mock_db, mock_boxim, mock_auth, mock_avatar): + avatar_filter = MagicMock() + avatar_filter.first.return_value = mock_avatar + avatar_query = MagicMock() + avatar_query.filter.return_value = avatar_filter + mock_db.query.return_value = avatar_query + + with patch("services.takeover_service.httpx.AsyncClient") as mock_client_class: + mock_response = MagicMock() + mock_response.json.return_value = {"code": 200, "data": {"answer": "Hello back"}} + mock_client_class.return_value.__aenter__.return_value.post.return_value = mock_response + + service = TakeoverService(mock_db, mock_boxim) + message = {"from_accid": "user_acc", "content": "Hello"} + + result = await service.execute_takeover(mock_auth, message) + + assert result is True + mock_boxim.get_credentials.assert_called_once_with("owner_huihui_123") + mock_boxim.send_p2p_message.assert_called_once() + + +@pytest.mark.asyncio +async def test_execute_takeover_fails_when_avatar_not_found(mock_db, mock_boxim, mock_auth): + """execute_takeover should return False when Avatar is not found.""" + avatar_filter = MagicMock() + avatar_filter.first.return_value = None + avatar_query = MagicMock() + avatar_query.filter.return_value = avatar_filter + mock_db.query.return_value = avatar_query + + service = TakeoverService(mock_db, mock_boxim) + message = {"from_accid": "user_acc", "content": "Hello"} + + result = await service.execute_takeover(mock_auth, message) + + assert result is False + mock_boxim.get_credentials.assert_not_called() + + +@pytest.mark.asyncio +async def test_execute_takeover_fails_when_no_credentials(mock_db, mock_boxim, mock_auth, mock_avatar): + """execute_takeover should return False when boxim.get_credentials returns None.""" + avatar_filter = MagicMock() + avatar_filter.first.return_value = mock_avatar + avatar_query = MagicMock() + avatar_query.filter.return_value = avatar_filter + mock_db.query.return_value = avatar_query + + mock_boxim.get_credentials.return_value = None + service = TakeoverService(mock_db, mock_boxim) + message = {"from_accid": "user_acc", "content": "Hello"} + + result = await service.execute_takeover(mock_auth, message) + + assert result is False + + +# --- enqueue_delayed_message --- + + +def test_enqueue_delayed_message_with_redis(mock_db, mock_boxim, mock_auth, mock_avatar): + mock_redis = MagicMock() + avatar_filter = MagicMock() + avatar_filter.first.return_value = mock_avatar + avatar_query = MagicMock() + avatar_query.filter.return_value = avatar_filter + mock_db.query.return_value = avatar_query + + service = TakeoverService(mock_db, mock_boxim, mock_redis) + message = {"msg_id": "msg_1", "from_accid": "user_acc", "content": "Hello"} + + service.enqueue_delayed_message(mock_auth, message) + + mock_redis.setex.assert_called_once() + call_args = mock_redis.setex.call_args + value = call_args[0][1] + import json + payload = json.loads(call_args[0][2]) + assert payload["owner_huihui_id"] == "owner_huihui_123" + + +def test_enqueue_delayed_message_without_redis_logs_warning(mock_db, mock_boxim, mock_auth, mock_avatar): + """When Redis is not configured, enqueue_delayed_message should log a warning and not crash.""" + avatar_filter = MagicMock() + avatar_filter.first.return_value = mock_avatar + avatar_query = MagicMock() + avatar_query.filter.return_value = avatar_filter + mock_db.query.return_value = avatar_query + + service = TakeoverService(mock_db, mock_boxim) + message = {"msg_id": "msg_1", "from_accid": "user_acc", "content": "Hello"} + + service.enqueue_delayed_message(mock_auth, message) + + +# --- process_delayed_queue --- + + +@pytest.mark.asyncio +async def test_process_delayed_queue_no_redis(mock_db, mock_boxim): + """process_delayed_queue should return immediately without Redis.""" + service = TakeoverService(mock_db, mock_boxim) + await service.process_delayed_queue() + mock_db.query.assert_not_called() + + +@pytest.mark.asyncio +async def test_process_delayed_queue_processes_messages(mock_db, mock_boxim, mock_auth, mock_avatar): + """process_delayed_queue should read from Redis, resolve auth, and execute takeover.""" + mock_redis = MagicMock() + mock_redis.keys.return_value = ["takeover:delayed:target_user_123:msg_1"] + mock_redis.get.return_value = '{"from_accid": "user_acc", "content": "Hello"}' + + avatar_filter = MagicMock() + avatar_filter.first.return_value = mock_avatar + avatar_query = MagicMock() + avatar_query.filter.return_value = avatar_filter + + auth_filter = MagicMock() + auth_filter.filter.return_value = auth_filter + auth_filter.first.return_value = mock_auth + + def query_side_effect(model): + if model == Avatar: + return avatar_query + return auth_filter + + mock_db.query.side_effect = query_side_effect + + with patch("services.takeover_service.httpx.AsyncClient") as mock_client_class: + mock_response = MagicMock() + mock_response.json.return_value = {"code": 200, "data": {"answer": "Hello back"}} + mock_client_class.return_value.__aenter__.return_value.post.return_value = mock_response + + service = TakeoverService(mock_db, mock_boxim, mock_redis) + await service.process_delayed_queue() + + mock_boxim.send_p2p_message.assert_called_once() + mock_redis.delete.assert_called_once()