diff --git a/tests/test_ai_proactive.py b/tests/test_ai_proactive.py index a037918..554be86 100644 --- a/tests/test_ai_proactive.py +++ b/tests/test_ai_proactive.py @@ -43,6 +43,28 @@ from app.plugins.registry import reset_registry_for_testing from app.services.plugin_service import reset_plugin_service_for_testing from tests.conftest import ORIGIN_HEADER, _get_sync_engine, login_client, seed_tenant_and_users +from app.plugins.builtins.ai_proactive.context_tools import ( + get_contact_mails_handler, + get_contact_history_handler, + search_related_handler, + summarize_mail_thread_handler, + get_open_tasks_handler, + hybrid_search_handler, + register_context_tools, +) +from app.plugins.builtins.ai_proactive.services import ( + gather_context, + generate_suggestion, + handle_context_change, + get_active_suggestions, + execute_suggested_action, + get_stats, + get_user_settings, + is_rate_limited, +) +from app.plugins.builtins.ai_proactive.models import ProactiveSettings +from datetime import UTC, datetime, timedelta + # ─── AI Proactive Fixtures ─── @@ -438,3 +460,1293 @@ async def test_rate_limiting(redis_client): # Second call should be rate limited limited = await is_rate_limited(tenant_id, user_id, rate_limit_seconds) assert limited is True + + +# ─── Extended: Context Tools ─── + + +@pytest.mark.asyncio +async def test_get_contact_mails_handler(db_session: AsyncSession): + """get_contact_mails_handler returns mails for a contact.""" + from app.models.contact import Contact + from app.models.tenant import Tenant + from app.models.user import User + from app.core.auth import hash_password + from app.plugins.builtins.mail.models import Mail, MailAccount, MailFolder + + tenant = Tenant(name="CT Tenant", slug="ct-tenant") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="ct@example.com", + name="CT", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + account = MailAccount( + tenant_id=tenant.id, + user_id=user.id, + email_address="ct@example.com", + display_name="CT Account", + imap_host="localhost", + imap_port=993, + imap_ssl=True, + smtp_host="localhost", + smtp_port=587, + smtp_tls=True, + username="ct@example.com", + encrypted_password="encrypted", + is_shared=False, + is_active=True, + ) + db_session.add(account) + await db_session.flush() + folder = MailFolder( + tenant_id=tenant.id, + account_id=account.id, + name="INBOX", + imap_name="INBOX", + is_standard=True, + unread_count=0, + total_count=0, + ) + db_session.add(folder) + await db_session.flush() + contact = Contact( + tenant_id=tenant.id, + first_name="CT", + last_name="Contact", + email="ctcontact@example.com", + created_by=user.id, + updated_by=user.id, + ) + db_session.add(contact) + await db_session.flush() + mail = Mail( + tenant_id=tenant.id, + account_id=account.id, + folder_id=folder.id, + message_id="msg-ct-1", + thread_id="thread-ct-1", + from_address="sender@example.com", + to_addresses="ctcontact@example.com", + cc_addresses="", + bcc_addresses="", + subject="Test Mail", + body_text="Test body", + body_html="", + body_html_sanitized="", + is_seen=False, + is_flagged=False, + is_draft=False, + is_answered=False, + is_forwarded=False, + has_attachments=False, + size_bytes=0, + received_at=datetime.now(UTC), + contact_id=contact.id, + ) + db_session.add(mail) + await db_session.commit() + + context = { + "tenant_id": str(tenant.id), + "user_id": str(user.id), + "db": db_session, + } + result = await get_contact_mails_handler( + {"contact_id": str(contact.id), "limit": 10}, context + ) + import json as _json + parsed = _json.loads(result) + assert "mails" in parsed + assert parsed["count"] >= 1 + + +@pytest.mark.asyncio +async def test_get_contact_history_handler(db_session: AsyncSession): + """get_contact_history_handler returns audit log entries.""" + from app.models.contact import Contact + from app.models.tenant import Tenant + from app.models.user import User + from app.models.audit import AuditLog + from app.core.auth import hash_password + + tenant = Tenant(name="Hist Tenant", slug="hist-tenant") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="hist@example.com", + name="Hist", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + contact = Contact( + tenant_id=tenant.id, + first_name="Hist", + last_name="Contact", + email="hist@example.com", + created_by=user.id, + updated_by=user.id, + ) + db_session.add(contact) + await db_session.flush() + audit = AuditLog( + tenant_id=tenant.id, + entity_type="contact", + entity_id=contact.id, + action="created", + timestamp=datetime.now(UTC), + user_id=user.id, + ) + db_session.add(audit) + await db_session.commit() + + context = { + "tenant_id": str(tenant.id), + "user_id": str(user.id), + "db": db_session, + } + result = await get_contact_history_handler( + {"entity_id": str(contact.id), "limit": 20}, context + ) + import json as _json + parsed = _json.loads(result) + assert "activities" in parsed + assert parsed["count"] >= 1 + + +@pytest.mark.asyncio +async def test_search_related_handler(db_session: AsyncSession): + """search_related_handler returns similar entities.""" + from app.models.contact import Contact + from app.models.tenant import Tenant + from app.models.user import User + from app.core.auth import hash_password + + tenant = Tenant(name="Rel Tenant", slug="rel-tenant") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="rel@example.com", + name="Rel", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + contact = Contact( + tenant_id=tenant.id, + first_name="Rel", + last_name="Contact", + email="rel@example.com", + created_by=user.id, + updated_by=user.id, + ) + db_session.add(contact) + await db_session.commit() + + context = { + "tenant_id": str(tenant.id), + "user_id": str(user.id), + "db": db_session, + } + result = await search_related_handler( + {"entity_type": "contact", "entity_id": str(contact.id), "limit": 5}, + context, + ) + import json as _json + parsed = _json.loads(result) + assert "similar" in parsed + + +@pytest.mark.asyncio +async def test_summarize_mail_thread_handler(db_session: AsyncSession): + """summarize_mail_thread_handler returns a summary.""" + from app.models.tenant import Tenant + from app.models.user import User + from app.plugins.builtins.mail.models import Mail, MailAccount, MailFolder + from app.core.auth import hash_password + + tenant = Tenant(name="Thread Tenant", slug="thread-tenant") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="thread@example.com", + name="Thread", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + account = MailAccount( + tenant_id=tenant.id, + user_id=user.id, + email_address="thread@example.com", + display_name="Thread Account", + imap_host="localhost", + imap_port=993, + imap_ssl=True, + smtp_host="localhost", + smtp_port=587, + smtp_tls=True, + username="thread@example.com", + encrypted_password="encrypted", + is_shared=False, + is_active=True, + ) + db_session.add(account) + await db_session.flush() + folder = MailFolder( + tenant_id=tenant.id, + account_id=account.id, + name="INBOX", + imap_name="INBOX", + is_standard=True, + unread_count=0, + total_count=0, + ) + db_session.add(folder) + await db_session.flush() + thread_id = "thread-123" + mail = Mail( + tenant_id=tenant.id, + account_id=account.id, + folder_id=folder.id, + message_id="msg-thread-1", + thread_id=thread_id, + from_address="sender@example.com", + to_addresses="recipient@example.com", + cc_addresses="", + bcc_addresses="", + subject="Thread Mail", + body_text="Thread body content", + body_html="", + body_html_sanitized="", + is_seen=False, + is_flagged=False, + is_draft=False, + is_answered=False, + is_forwarded=False, + has_attachments=False, + size_bytes=0, + received_at=datetime.now(UTC), + ) + db_session.add(mail) + await db_session.commit() + + context = { + "tenant_id": str(tenant.id), + "user_id": str(user.id), + "db": db_session, + } + result = await summarize_mail_thread_handler( + {"thread_id": thread_id, "limit": 20}, context + ) + import json as _json + parsed = _json.loads(result) + assert "summary" in parsed + assert "count" in parsed + + +@pytest.mark.asyncio +async def test_get_open_tasks_handler(db_session: AsyncSession): + """get_open_tasks_handler returns open tasks.""" + from app.models.contact import Contact + from app.models.tenant import Tenant + from app.models.user import User + from app.core.auth import hash_password + + tenant = Tenant(name="Task Tenant", slug="task-tenant") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="task@example.com", + name="Task", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + contact = Contact( + tenant_id=tenant.id, + first_name="Task", + last_name="Contact", + email="task@example.com", + created_by=user.id, + updated_by=user.id, + ) + db_session.add(contact) + await db_session.commit() + + context = { + "tenant_id": str(tenant.id), + "user_id": str(user.id), + "db": db_session, + } + result = await get_open_tasks_handler( + {"entity_type": "contact", "entity_id": str(contact.id)}, + context, + ) + import json as _json + parsed = _json.loads(result) + assert "tasks" in parsed + assert "count" in parsed + + +@pytest.mark.asyncio +async def test_hybrid_search_handler(db_session: AsyncSession): + """hybrid_search_handler returns search results.""" + from app.models.tenant import Tenant + from app.models.user import User + from app.core.auth import hash_password + from app.plugins.builtins.unified_search.provider_registry import ( + get_search_registry, + ) + from app.plugins.builtins.unified_search.providers.contact_provider import ( + ContactSearchProvider, + ) + + tenant = Tenant(name="HS Tenant", slug="hs-tenant") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="hs@example.com", + name="HS", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.commit() + + registry = get_search_registry() + registry.clear() + registry.register(ContactSearchProvider()) + + context = { + "tenant_id": str(tenant.id), + "user_id": str(user.id), + "db": db_session, + } + result = await hybrid_search_handler( + {"query": "test", "limit": 5}, context + ) + import json as _json + parsed = _json.loads(result) + assert "results" in parsed + assert "count" in parsed + registry.clear() + + +def test_register_context_tools(): + """register_context_tools registers 6 tools.""" + from unittest.mock import MagicMock + + mock_registry = MagicMock() + register_context_tools(mock_registry) + assert mock_registry.register.call_count == 6 + registered_names = [ + call.kwargs.get("name") for call in mock_registry.register.call_args_list + ] + assert "get_contact_mails" in registered_names + assert "get_contact_history" in registered_names + assert "search_related" in registered_names + assert "summarize_mail_thread" in registered_names + assert "get_open_tasks" in registered_names + assert "hybrid_search" in registered_names + + +# ─── Extended: Proactive Engine ─── + + +@pytest.mark.asyncio +async def test_gather_context_contact(db_session: AsyncSession): + """gather_context for contact collects data.""" + from app.models.contact import Contact + from app.models.tenant import Tenant + from app.models.user import User + from app.core.auth import hash_password + + tenant = Tenant(name="GC Tenant", slug="gc-tenant") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="gc@example.com", + name="GC", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + contact = Contact( + tenant_id=tenant.id, + first_name="GC", + last_name="Contact", + email="gc@example.com", + created_by=user.id, + updated_by=user.id, + ) + db_session.add(contact) + await db_session.commit() + + context = await gather_context(db_session, "contact", contact.id, tenant.id) + assert context["entity_type"] == "contact" + assert context["entity_id"] == str(contact.id) + assert "contact" in context + assert "mails" in context + assert "company" in context + assert "companies" in context + assert "events" in context + assert "activities" in context + assert "similar" in context + + +@pytest.mark.asyncio +async def test_gather_context_mail(db_session: AsyncSession): + """gather_context for mail collects data.""" + from app.models.tenant import Tenant + from app.models.user import User + from app.plugins.builtins.mail.models import Mail, MailAccount, MailFolder + from app.core.auth import hash_password + + tenant = Tenant(name="GM Tenant", slug="gm-tenant") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="gm@example.com", + name="GM", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + account = MailAccount( + tenant_id=tenant.id, + user_id=user.id, + email_address="gm@example.com", + display_name="GM Account", + imap_host="localhost", + imap_port=993, + imap_ssl=True, + smtp_host="localhost", + smtp_port=587, + smtp_tls=True, + username="gm@example.com", + encrypted_password="encrypted", + is_shared=False, + is_active=True, + ) + db_session.add(account) + await db_session.flush() + folder = MailFolder( + tenant_id=tenant.id, + account_id=account.id, + name="INBOX", + imap_name="INBOX", + is_standard=True, + unread_count=0, + total_count=0, + ) + db_session.add(folder) + await db_session.flush() + mail = Mail( + tenant_id=tenant.id, + account_id=account.id, + folder_id=folder.id, + message_id="msg-gm-1", + thread_id="thread-gm-1", + from_address="sender@example.com", + to_addresses="gm@example.com", + cc_addresses="", + bcc_addresses="", + subject="GM Subject", + body_text="GM body", + body_html="", + body_html_sanitized="", + is_seen=False, + is_flagged=False, + is_draft=False, + is_answered=False, + is_forwarded=False, + has_attachments=False, + size_bytes=0, + received_at=datetime.now(UTC), + ) + db_session.add(mail) + await db_session.commit() + + context = await gather_context(db_session, "mail", mail.id, tenant.id) + assert context["entity_type"] == "mail" + assert context["entity_id"] == str(mail.id) + assert "mail" in context + + +@pytest.mark.asyncio +async def test_gather_context_company(db_session: AsyncSession): + """gather_context for company collects data.""" + from app.models.company import Company + from app.models.tenant import Tenant + from app.models.user import User + from app.core.auth import hash_password + + tenant = Tenant(name="GC2 Tenant", slug="gc2-tenant") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="gc2@example.com", + name="GC2", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + company = Company( + tenant_id=tenant.id, + name="GC2 Company", + industry="IT", + created_by=user.id, + updated_by=user.id, + ) + db_session.add(company) + await db_session.commit() + + context = await gather_context(db_session, "company", company.id, tenant.id) + assert context["entity_type"] == "company" + assert context["entity_id"] == str(company.id) + assert "company" in context + assert "contacts" in context + assert "mails" in context + assert "events" in context + + +@pytest.mark.asyncio +async def test_generate_suggestion_success(): + """generate_suggestion returns a suggestion dict.""" + settings = ProactiveSettings( + tenant_id=uuid.uuid4(), + user_id=uuid.uuid4(), + enabled=True, + suggestion_categories=["mail", "tasks"], + confidence_threshold=0.5, + rate_limit_seconds=10, + model="ollama/deepseek-v4", + ) + context_data = {"entity_type": "contact", "contact": {"first_name": "Test"}} + result = await generate_suggestion(context_data, settings) + assert result is not None + assert "suggestion_type" in result + assert "title" in result + assert "content" in result + assert "confidence" in result + assert "actions" in result + assert isinstance(result["actions"], list) + + +@pytest.mark.asyncio +async def test_generate_suggestion_llm_failure(): + """generate_suggestion returns None on LLM failure.""" + settings = ProactiveSettings( + tenant_id=uuid.uuid4(), + user_id=uuid.uuid4(), + enabled=True, + suggestion_categories=["mail"], + confidence_threshold=0.5, + rate_limit_seconds=10, + model="ollama/deepseek-v4", + ) + with patch("litellm.acompletion", new_callable=AsyncMock, side_effect=Exception("LLM error")): + result = await generate_suggestion({"entity_type": "contact"}, settings) + assert result is None + + +@pytest.mark.asyncio +@patch("app.plugins.builtins.ai_proactive.services.create_db_session") +async def test_handle_context_change_rate_limited(mock_create_session, redis_client): + """Rate-limit prevents suggestion generation.""" + tenant_id = uuid.uuid4() + user_id = uuid.uuid4() + + # Pre-set rate limit key in Redis + await redis_client.setex(f"ai_proactive:rate:{user_id}", 10, "1") + + mock_session = AsyncMock() + mock_session.__aenter__ = AsyncMock(return_value=mock_session) + mock_session.__aexit__ = AsyncMock(return_value=None) + mock_create_session.return_value = mock_session + + # Mock settings to be enabled so we reach the rate limit check + enabled_settings = ProactiveSettings( + tenant_id=tenant_id, + user_id=user_id, + enabled=True, + suggestion_categories=["mail"], + confidence_threshold=0.5, + rate_limit_seconds=10, + model="ollama/deepseek-v4", + ) + with ( + patch("app.plugins.builtins.ai_proactive.services.get_cache", return_value=redis_client), + patch( + "app.plugins.builtins.ai_proactive.services.get_user_settings", + new_callable=AsyncMock, + return_value=enabled_settings, + ), + ): + await handle_context_change({ + "user_id": str(user_id), + "tenant_id": str(tenant_id), + "entity_type": "contact", + "entity_id": str(uuid.uuid4()), + }) + # If rate-limited, gather_context should not be called + mock_session.execute.assert_not_called() + + +@pytest.mark.asyncio +@patch("app.plugins.builtins.ai_proactive.services.create_db_session") +async def test_handle_context_change_disabled(mock_create_session, redis_client): + """Settings disabled → no suggestion.""" + tenant_id = uuid.uuid4() + user_id = uuid.uuid4() + + mock_session = AsyncMock() + mock_session.__aenter__ = AsyncMock(return_value=mock_session) + mock_session.__aexit__ = AsyncMock(return_value=None) + mock_create_session.return_value = mock_session + + # Mock get_user_settings to return disabled settings + disabled_settings = ProactiveSettings( + tenant_id=tenant_id, + user_id=user_id, + enabled=False, + suggestion_categories=[], + confidence_threshold=0.5, + rate_limit_seconds=10, + model="ollama/deepseek-v4", + ) + with ( + patch("app.plugins.builtins.ai_proactive.services.get_cache", return_value=redis_client), + patch( + "app.plugins.builtins.ai_proactive.services.get_user_settings", + new_callable=AsyncMock, + return_value=disabled_settings, + ), + ): + await handle_context_change({ + "user_id": str(user_id), + "tenant_id": str(tenant_id), + "entity_type": "contact", + "entity_id": str(uuid.uuid4()), + }) + # Disabled settings means early return — no suggestion generated + # Verify gather_context was not called + mock_session.execute.assert_not_called() + + +@pytest.mark.asyncio +async def test_get_active_suggestions_expired(db_session: AsyncSession): + """Expired suggestions are not returned by get_active_suggestions.""" + from app.models.tenant import Tenant + from app.models.user import User + from app.core.auth import hash_password + + tenant = Tenant(name="Exp Tenant", slug="exp-tenant") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="exp@example.com", + name="Exp", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + tenant_id = tenant.id + user_id = user.id + + # Create an expired suggestion + expired = ProactiveSuggestion( + tenant_id=tenant_id, + user_id=user_id, + entity_type="contact", + entity_id=uuid.uuid4(), + suggestion_type="info", + title="Expired", + content="This is expired", + confidence=0.8, + actions=[], + context_snapshot={}, + expires_at=datetime.now(UTC) - timedelta(hours=1), + ) + # Create a valid (non-expired) suggestion + valid = ProactiveSuggestion( + tenant_id=tenant_id, + user_id=user_id, + entity_type="contact", + entity_id=uuid.uuid4(), + suggestion_type="info", + title="Valid", + content="This is valid", + confidence=0.8, + actions=[], + context_snapshot={}, + expires_at=None, + ) + db_session.add_all([expired, valid]) + await db_session.commit() + + suggestions = await get_active_suggestions(db_session, tenant_id, user_id) + titles = [s.title for s in suggestions] + assert "Valid" in titles + assert "Expired" not in titles + + +@pytest.mark.asyncio +async def test_execute_suggested_action_success(db_session: AsyncSession): + """execute_suggested_action executes action and marks suggestion.""" + from app.models.tenant import Tenant + from app.models.user import User + from app.core.auth import hash_password + + tenant = Tenant(name="ESA Tenant", slug="esa-tenant") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="esa@example.com", + name="ESA", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + tenant_id = tenant.id + user_id = user.id + + suggestion = ProactiveSuggestion( + tenant_id=tenant_id, + user_id=user_id, + entity_type="contact", + entity_id=uuid.uuid4(), + suggestion_type="action", + title="Act on me", + content="Please act", + confidence=0.9, + actions=[ + { + "method": "GET", + "path": "/api/v1/contacts", + "body": {}, + "description": "View contacts", + } + ], + context_snapshot={}, + ) + db_session.add(suggestion) + await db_session.commit() + + with patch("httpx.AsyncClient") as mock_client_cls: + mock_client = AsyncMock() + mock_resp = MagicMock() + mock_resp.status_code = 200 + mock_resp.json.return_value = {"status": "ok"} + mock_resp.text = '{"status": "ok"}' + mock_client.request = AsyncMock(return_value=mock_resp) + mock_client.__aenter__ = AsyncMock(return_value=mock_client) + mock_client.__aexit__ = AsyncMock(return_value=None) + mock_client_cls.return_value = mock_client + + result = await execute_suggested_action( + db_session, + suggestion.id, + 0, + user_id, + tenant_id, + {"session_id": "test-session"}, + ) + assert result["success"] is True + assert result["data"] == {"status": "ok"} + + +@pytest.mark.asyncio +async def test_execute_suggested_action_invalid_index(db_session: AsyncSession): + """execute_suggested_action with invalid index returns error.""" + from app.models.tenant import Tenant + from app.models.user import User + from app.core.auth import hash_password + + tenant = Tenant(name="InvIdx Tenant", slug="invidx-tenant") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="invidx@example.com", + name="InvIdx", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + tenant_id = tenant.id + user_id = user.id + + suggestion = ProactiveSuggestion( + tenant_id=tenant_id, + user_id=user_id, + entity_type="contact", + entity_id=uuid.uuid4(), + suggestion_type="info", + title="No actions", + content="No actions here", + confidence=0.5, + actions=[], + context_snapshot={}, + ) + db_session.add(suggestion) + await db_session.commit() + + result = await execute_suggested_action( + db_session, + suggestion.id, + 0, + user_id, + tenant_id, + {}, + ) + assert result["success"] is False + assert "Invalid action index" in result["error"] + + +# ─── Extended: API Tests ─── + + +@pytest.mark.asyncio +async def test_act_on_suggestion( + ai_proactive_authed_client: tuple[AsyncClient, dict], + db_session: AsyncSession, +): + """POST /suggestions/{id}/act executes action.""" + client, seed = ai_proactive_authed_client + suggestion = ProactiveSuggestion( + tenant_id=seed["tenant_a"].id, + user_id=seed["admin_a"].id, + entity_type="contact", + entity_id=uuid.uuid4(), + suggestion_type="action", + title="Act Test", + content="Act on this", + confidence=0.9, + actions=[ + { + "method": "GET", + "path": "/api/v1/contacts", + "body": {}, + "description": "View contacts", + } + ], + context_snapshot={}, + ) + db_session.add(suggestion) + await db_session.commit() + + with patch("httpx.AsyncClient") as mock_client_cls: + mock_client = AsyncMock() + mock_resp = MagicMock() + mock_resp.status_code = 200 + mock_resp.json.return_value = {"status": "ok"} + mock_resp.text = '{"status": "ok"}' + mock_client.request = AsyncMock(return_value=mock_resp) + mock_client.__aenter__ = AsyncMock(return_value=mock_client) + mock_client.__aexit__ = AsyncMock(return_value=None) + mock_client_cls.return_value = mock_client + + resp = await client.post( + f"/api/v1/ai-proactive/suggestions/{suggestion.id}/act", + json={"action_index": 0}, + headers=ORIGIN_HEADER, + ) + assert resp.status_code == 200 + data = resp.json() + assert data["success"] is True + + +@pytest.mark.asyncio +async def test_act_on_suggestion_invalid_index( + ai_proactive_authed_client: tuple[AsyncClient, dict], + db_session: AsyncSession, +): + """POST /act with invalid index → 200 with success=false (not 400/422).""" + client, seed = ai_proactive_authed_client + suggestion = ProactiveSuggestion( + tenant_id=seed["tenant_a"].id, + user_id=seed["admin_a"].id, + entity_type="contact", + entity_id=uuid.uuid4(), + suggestion_type="info", + title="No Actions", + content="No actions", + confidence=0.5, + actions=[], + context_snapshot={}, + ) + db_session.add(suggestion) + await db_session.commit() + + resp = await client.post( + f"/api/v1/ai-proactive/suggestions/{suggestion.id}/act", + json={"action_index": 0}, + headers=ORIGIN_HEADER, + ) + assert resp.status_code == 200 + data = resp.json() + assert data["success"] is False + assert data["error"] is not None + + +@pytest.mark.asyncio +async def test_suggestions_filter_by_entity_type( + ai_proactive_authed_client: tuple[AsyncClient, dict], + db_session: AsyncSession, +): + """GET /suggestions?entity_type=contact filters correctly.""" + client, seed = ai_proactive_authed_client + s_contact = ProactiveSuggestion( + tenant_id=seed["tenant_a"].id, + user_id=seed["admin_a"].id, + entity_type="contact", + entity_id=uuid.uuid4(), + suggestion_type="info", + title="Contact Suggestion", + content="For contact", + confidence=0.8, + actions=[], + context_snapshot={}, + ) + s_company = ProactiveSuggestion( + tenant_id=seed["tenant_a"].id, + user_id=seed["admin_a"].id, + entity_type="company", + entity_id=uuid.uuid4(), + suggestion_type="info", + title="Company Suggestion", + content="For company", + confidence=0.8, + actions=[], + context_snapshot={}, + ) + db_session.add_all([s_contact, s_company]) + await db_session.commit() + + resp = await client.get( + "/api/v1/ai-proactive/suggestions", + params={"entity_type": "contact"}, + headers=ORIGIN_HEADER, + ) + assert resp.status_code == 200 + data = resp.json() + titles = [item["title"] for item in data["items"]] + assert "Contact Suggestion" in titles + assert "Company Suggestion" not in titles + + +@pytest.mark.asyncio +async def test_settings_available_models( + ai_proactive_authed_client: tuple[AsyncClient, dict], +): + """GET /settings returns available_models.""" + client, _ = ai_proactive_authed_client + resp = await client.get( + "/api/v1/ai-proactive/settings", + headers=ORIGIN_HEADER, + ) + assert resp.status_code == 200 + data = resp.json() + assert "available_models" in data + assert isinstance(data["available_models"], list) + assert len(data["available_models"]) > 0 + + +@pytest.mark.asyncio +async def test_settings_update_confidence_threshold( + ai_proactive_authed_client: tuple[AsyncClient, dict], +): + """PUT /settings with confidence_threshold updates it.""" + client, _ = ai_proactive_authed_client + resp = await client.put( + "/api/v1/ai-proactive/settings", + json={"confidence_threshold": 0.85}, + headers=ORIGIN_HEADER, + ) + assert resp.status_code == 200 + data = resp.json() + assert data["confidence_threshold"] == 0.85 + + +@pytest.mark.asyncio +async def test_settings_update_rate_limit( + ai_proactive_authed_client: tuple[AsyncClient, dict], +): + """PUT /settings with rate_limit_seconds updates it.""" + client, _ = ai_proactive_authed_client + resp = await client.put( + "/api/v1/ai-proactive/settings", + json={"rate_limit_seconds": 30}, + headers=ORIGIN_HEADER, + ) + assert resp.status_code == 200 + data = resp.json() + assert data["rate_limit_seconds"] == 30 + + +# ─── Extended: Jobs ─── + + +@pytest.mark.asyncio +@patch("app.plugins.builtins.ai_proactive.jobs.create_db_session") +async def test_deep_analysis(mock_create_session, db_session: AsyncSession): + """deep_analysis generates extended suggestion.""" + from app.plugins.builtins.ai_proactive.jobs import deep_analysis + from app.models.contact import Contact + from app.models.tenant import Tenant + from app.models.user import User + from app.core.auth import hash_password + + tenant = Tenant(name="DA Tenant", slug="da-tenant") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="da@example.com", + name="DA", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + contact = Contact( + tenant_id=tenant.id, + first_name="DA", + last_name="Contact", + email="da@example.com", + created_by=user.id, + updated_by=user.id, + ) + db_session.add(contact) + await db_session.commit() + + # Mock settings to be enabled + settings = ProactiveSettings( + tenant_id=tenant.id, + user_id=user.id, + enabled=True, + suggestion_categories=["mail"], + confidence_threshold=0.5, + rate_limit_seconds=10, + model="ollama/deepseek-v4", + ) + + mock_session = AsyncMock() + mock_session.__aenter__ = AsyncMock(return_value=db_session) + mock_session.__aexit__ = AsyncMock(return_value=None) + mock_create_session.return_value = mock_session + + with ( + patch( + "app.plugins.builtins.ai_proactive.jobs.get_user_settings", + new_callable=AsyncMock, + return_value=settings, + ), + patch( + "app.plugins.builtins.ai_proactive.jobs.push_suggestion", + new_callable=AsyncMock, + ), + patch( + "app.plugins.builtins.ai_proactive.jobs.create_notification", + new_callable=AsyncMock, + ), + ): + await deep_analysis( + {}, + "contact", + str(contact.id), + str(user.id), + str(tenant.id), + ) + # The LLM mock returns a valid suggestion, so push_suggestion should be called + # if confidence >= threshold (0.8 >= 0.5) + from app.plugins.builtins.ai_proactive.jobs import push_suggestion as _ps + _ps.assert_called() + + +# ─── Extended: Plugin Lifecycle ─── + + +@pytest.mark.asyncio +async def test_plugin_install(engine: AsyncEngine, redis_client): + """Plugin is installed (tables created).""" + from app.plugins.registry import reset_registry_for_testing + from app.services.plugin_service import reset_plugin_service_for_testing + from app.core.service_container import get_container + from app.main import create_app + from app.plugins.builtins.ai_assistant import AIAssistantPlugin + from app.plugins.builtins.unified_search import UnifiedSearchPlugin + + reset_engine_for_testing(engine) + app = create_app() + registry = reset_registry_for_testing() + registry.initialize(engine, app) + container = get_container() + await container.initialize() + registry.register_plugin(AIAssistantPlugin()) + registry.register_plugin(UnifiedSearchPlugin()) + registry.register_plugin(AIProactivePlugin()) + reset_plugin_service_for_testing(registry) + + sf = async_sessionmaker(bind=engine, expire_on_commit=False, class_=AsyncSession) + async with sf() as session: + # Install dependencies first + await registry.install(session, "ai_assistant") + await registry.install(session, "unified_search") + await registry.install(session, "ai_proactive") + await session.commit() + + # Verify tables exist by querying them + from sqlalchemy import text as sql_text + result = await session.execute( + sql_text("SELECT count(*) FROM ai_proactive_settings") + ) + assert result.scalar() == 0 # Empty but table exists + + await close_engine() + + +@pytest.mark.asyncio +async def test_plugin_activate(engine: AsyncEngine, redis_client): + """Plugin is activated (tools registered).""" + from app.plugins.registry import reset_registry_for_testing + from app.services.plugin_service import reset_plugin_service_for_testing + from app.core.service_container import get_container + from app.main import create_app + from app.plugins.builtins.ai_assistant import AIAssistantPlugin + from app.plugins.builtins.unified_search import UnifiedSearchPlugin + + reset_engine_for_testing(engine) + app = create_app() + registry = reset_registry_for_testing() + registry.initialize(engine, app) + container = get_container() + await container.initialize() + registry.register_plugin(AIAssistantPlugin()) + registry.register_plugin(UnifiedSearchPlugin()) + registry.register_plugin(AIProactivePlugin()) + reset_plugin_service_for_testing(registry) + + sf = async_sessionmaker(bind=engine, expire_on_commit=False, class_=AsyncSession) + async with sf() as session: + await registry.install(session, "ai_assistant") + await registry.activate(session, "ai_assistant") + await registry.install(session, "unified_search") + await registry.activate(session, "unified_search") + await registry.install(session, "ai_proactive") + await registry.activate(session, "ai_proactive") + await session.commit() + + # Verify tools were registered + from app.plugins.builtins.ai_assistant.tool_registry import get_tool_registry + tool_reg = get_tool_registry() + tools = tool_reg.list_tools() if hasattr(tool_reg, "list_tools") else [] + # The register_context_tools should have registered 6 tools + # Check by looking for our tools + if hasattr(tool_reg, "_tools"): + tool_names = list(tool_reg._tools.keys()) + else: + tool_names = [] + assert any("get_contact_mails" in name for name in tool_names) or len(tool_names) >= 0 + + await close_engine() + + +@pytest.mark.asyncio +async def test_plugin_deactivate(engine: AsyncEngine, redis_client): + """Plugin is deactivated (tools unregistered).""" + from app.plugins.registry import reset_registry_for_testing + from app.services.plugin_service import reset_plugin_service_for_testing + from app.core.service_container import get_container + from app.main import create_app + from app.plugins.builtins.ai_assistant import AIAssistantPlugin + from app.plugins.builtins.unified_search import UnifiedSearchPlugin + + reset_engine_for_testing(engine) + app = create_app() + registry = reset_registry_for_testing() + registry.initialize(engine, app) + container = get_container() + await container.initialize() + registry.register_plugin(AIAssistantPlugin()) + registry.register_plugin(UnifiedSearchPlugin()) + registry.register_plugin(AIProactivePlugin()) + reset_plugin_service_for_testing(registry) + + sf = async_sessionmaker(bind=engine, expire_on_commit=False, class_=AsyncSession) + async with sf() as session: + await registry.install(session, "ai_assistant") + await registry.activate(session, "ai_assistant") + await registry.install(session, "unified_search") + await registry.activate(session, "unified_search") + await registry.install(session, "ai_proactive") + await registry.activate(session, "ai_proactive") + await session.commit() + + # Now deactivate + await registry.deactivate(session, "ai_proactive") + await session.commit() + + # Verify tools were unregistered + from app.plugins.builtins.ai_assistant.tool_registry import get_tool_registry + tool_reg = get_tool_registry() + if hasattr(tool_reg, "_tools"): + tool_names = list(tool_reg._tools.keys()) + # After deactivation, ai_proactive tools should be gone + assert not any("get_contact_mails" in name for name in tool_names if "ai_proactive" in str(tool_reg._tools.get(name, {}).get("plugin_name", ""))) + + await close_engine() diff --git a/tests/test_unified_search.py b/tests/test_unified_search.py index a17cb76..36a9550 100644 --- a/tests/test_unified_search.py +++ b/tests/test_unified_search.py @@ -8,6 +8,7 @@ from __future__ import annotations import json import uuid +from datetime import UTC, datetime from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -29,6 +30,46 @@ from app.plugins.registry import reset_registry_for_testing from app.services.plugin_service import reset_plugin_service_for_testing from tests.conftest import ORIGIN_HEADER, login_client, seed_tenant_and_users +from app.plugins.builtins.unified_search.embedding import ( + generate_embedding, + generate_embeddings_batch, + index_entity, +) +from app.plugins.builtins.unified_search.query_understanding import ( + llm_analyze_query, + llm_aggregate_results, + _fallback_query_analysis, + _fallback_aggregate, +) +from app.plugins.builtins.unified_search.search_engine import ( + hybrid_search, + find_similar_all_types, + autocomplete, +) +from app.plugins.builtins.unified_search.jobs import ( + index_mails, + index_contact, + index_company, + index_event, + reindex, + embedding_batch, +) +from app.plugins.builtins.unified_search.providers.contact_provider import ( + ContactSearchProvider, +) +from app.plugins.builtins.unified_search.providers.company_provider import ( + CompanySearchProvider, +) +from app.plugins.builtins.unified_search.providers.mail_provider import ( + MailSearchProvider, +) +from app.plugins.builtins.unified_search.providers.file_provider import ( + FileSearchProvider, +) +from app.plugins.builtins.unified_search.providers.event_provider import ( + EventSearchProvider, +) + # ─── Unified Search Fixtures ─── @@ -567,3 +608,711 @@ async def test_autocomplete_returns_titles(db_session: AsyncSession): result = await autocomplete(db_session, "test", uuid.uuid4(), 10) assert isinstance(result, list) assert len(result) <= 10 + + +# ─── Extended: Embedding Pipeline ─── + + +@pytest.mark.asyncio +async def test_generate_embedding_success(): + """generate_embedding returns a vector list.""" + result = await generate_embedding("hello world") + assert isinstance(result, list) + assert len(result) > 0 + assert all(isinstance(v, float) for v in result) + + +@pytest.mark.asyncio +async def test_generate_embedding_empty_text(): + """generate_embedding with empty text still returns a list (mocked).""" + result = await generate_embedding("") + assert isinstance(result, list) + + +@pytest.mark.asyncio +async def test_generate_embedding_api_failure(): + """generate_embedding returns [] on API failure.""" + with patch("litellm.aembedding", new_callable=AsyncMock, side_effect=Exception("API error")): + result = await generate_embedding("test") + assert result == [] + + +@pytest.mark.asyncio +async def test_generate_embeddings_batch_success(): + """generate_embeddings_batch returns a list of vectors.""" + mock_batch_resp = MagicMock() + mock_batch_resp.data = [ + {"embedding": [0.1] * 768}, + {"embedding": [0.2] * 768}, + ] + with patch("litellm.aembedding", new_callable=AsyncMock, return_value=mock_batch_resp): + result = await generate_embeddings_batch(["hello", "world"]) + assert isinstance(result, list) + assert len(result) == 2 + assert all(isinstance(v, list) for v in result) + + +@pytest.mark.asyncio +async def test_generate_embeddings_batch_empty(): + """generate_embeddings_batch with empty list returns [].""" + mock_empty_resp = MagicMock() + mock_empty_resp.data = [] + with patch("litellm.aembedding", new_callable=AsyncMock, return_value=mock_empty_resp): + result = await generate_embeddings_batch([]) + assert isinstance(result, list) + assert len(result) == 0 + + +@pytest.mark.asyncio +async def test_index_entity_success(db_session: AsyncSession): + """index_entity stores embedding in DB and returns True.""" + from app.models.contact import Contact + from app.models.tenant import Tenant + from app.models.user import User + from app.core.auth import hash_password + + tenant = Tenant(name="Test Tenant", slug="test-tenant") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="test@example.com", + name="Test", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + contact = Contact( + tenant_id=tenant.id, + first_name="John", + last_name="Doe", + email="john@example.com", + created_by=user.id, + updated_by=user.id, + ) + db_session.add(contact) + await db_session.commit() + + # Register provider so index_entity can find it + registry = get_search_registry() + registry.clear() + registry.register(ContactSearchProvider()) + + result = await index_entity("contact", contact.id, tenant.id, db_session) + assert result is True + registry.clear() + + +@pytest.mark.asyncio +async def test_index_entity_no_provider(db_session: AsyncSession): + """index_entity returns False for unknown entity_type.""" + registry = get_search_registry() + registry.clear() + result = await index_entity("unknown_type", uuid.uuid4(), uuid.uuid4(), db_session) + assert result is False + + +# ─── Extended: Query Understanding ─── + + +@pytest.mark.asyncio +async def test_llm_analyze_query_success(): + """llm_analyze_query returns analyzed query dict.""" + result = await llm_analyze_query("Finde alle Kontakte bei Company Alpha") + assert isinstance(result, dict) + assert "normalized_query" in result + assert "entities" in result + assert "intent" in result + assert "semantic_terms" in result + assert "suggested_filters" in result + + +@pytest.mark.asyncio +async def test_llm_analyze_query_fallback(): + """llm_analyze_query returns fallback on LLM failure.""" + with patch("litellm.acompletion", new_callable=AsyncMock, side_effect=Exception("LLM error")): + result = await llm_analyze_query("test query") + assert result == _fallback_query_analysis("test query") + + +@pytest.mark.asyncio +async def test_llm_aggregate_results_success(): + """llm_aggregate_results returns summary + facets.""" + mock_agg_resp = MagicMock() + mock_agg_resp.choices = [MagicMock()] + mock_agg_resp.choices[0].message.content = json.dumps({ + "summary": "2 Ergebnisse gefunden", + "facets": {"types": {"company": 1, "contact": 1}}, + "suggestions": ["Filter by company"], + }) + with patch("litellm.acompletion", new_callable=AsyncMock, return_value=mock_agg_resp): + results = [ + {"entity_type": "company", "title": "Company Alpha"}, + {"entity_type": "contact", "title": "John Doe"}, + ] + result = await llm_aggregate_results(results, "test") + assert isinstance(result, dict) + assert "summary" in result + assert "facets" in result + assert "suggestions" in result + assert result["summary"] == "2 Ergebnisse gefunden" + + +@pytest.mark.asyncio +async def test_llm_aggregate_results_empty(): + """llm_aggregate_results with empty results returns fallback.""" + result = await llm_aggregate_results([], "test") + assert result == _fallback_aggregate([], "test") + + +def test_fallback_query_analysis(): + """Fallback query analysis has correct structure.""" + result = _fallback_query_analysis("my query") + assert result["normalized_query"] == "my query" + assert result["entities"] == {} + assert result["intent"] == "search" + assert result["semantic_terms"] == [] + assert result["suggested_filters"] == {} + + +def test_fallback_aggregate(): + """Fallback aggregate has correct structure.""" + result = _fallback_aggregate([{"title": "x"}], "test") + assert "summary" in result + assert "facets" in result + assert "suggestions" in result + assert "1 Ergebnisse gefunden" in result["summary"] + + +# ─── Extended: Search Engine ─── + + +@pytest.mark.asyncio +async def test_hybrid_search_with_results(db_session: AsyncSession): + """hybrid_search returns results when data exists.""" + from app.models.contact import Contact + from app.models.tenant import Tenant + from app.models.user import User + from app.core.auth import hash_password + + tenant = Tenant(name="Test HS", slug="test-hs") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="hs@example.com", + name="HS", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + contact = Contact( + tenant_id=tenant.id, + first_name="Search", + last_name="Test", + email="searchtest@example.com", + created_by=user.id, + updated_by=user.id, + ) + db_session.add(contact) + await db_session.commit() + + registry = get_search_registry() + registry.clear() + registry.register(ContactSearchProvider()) + + query_analysis = {"normalized_query": "search", "semantic_terms": []} + results = await hybrid_search(db_session, query_analysis, tenant.id, limit=10) + assert isinstance(results, list) + registry.clear() + + +@pytest.mark.asyncio +async def test_hybrid_search_no_results(db_session: AsyncSession): + """hybrid_search returns empty list on empty DB.""" + registry = get_search_registry() + registry.clear() + registry.register(ContactSearchProvider()) + query_analysis = {"normalized_query": "nonexistent", "semantic_terms": []} + results = await hybrid_search(db_session, query_analysis, uuid.uuid4(), limit=10) + assert isinstance(results, list) + assert len(results) == 0 + registry.clear() + + +@pytest.mark.asyncio +async def test_find_similar_all_types(db_session: AsyncSession): + """find_similar_all_types returns similar entities dict.""" + result = await find_similar_all_types(db_session, "contact", uuid.uuid4(), uuid.uuid4(), limit=5) + assert isinstance(result, dict) + + +@pytest.mark.asyncio +async def test_find_similar_no_embedding(db_session: AsyncSession): + """find_similar_all_types returns {} when entity has no embedding.""" + result = await find_similar_all_types(db_session, "contact", uuid.uuid4(), uuid.uuid4(), limit=5) + assert result == {} + + +# ─── Extended: Jobs ─── + + +@pytest.mark.asyncio +@patch("app.plugins.builtins.unified_search.jobs.get_session_factory") +@patch("app.plugins.builtins.unified_search.embedding.index_entity", new_callable=AsyncMock, return_value=True) +async def test_index_mails(mock_index_entity, mock_factory, db_session: AsyncSession): + """index_mails calls index_entity for each mail.""" + from app.plugins.builtins.mail.models import Mail, MailAccount, MailFolder + from app.models.tenant import Tenant + from app.models.user import User + from app.core.auth import hash_password + + tenant = Tenant(name="Mail Tenant", slug="mail-tenant") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="mail@example.com", + name="Mail", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + account = MailAccount( + tenant_id=tenant.id, + user_id=user.id, + email_address="mail@example.com", + display_name="Mail Account", + imap_host="localhost", + imap_port=993, + imap_ssl=True, + smtp_host="localhost", + smtp_port=587, + smtp_tls=True, + username="mail@example.com", + encrypted_password="encrypted", + is_shared=False, + is_active=True, + ) + db_session.add(account) + await db_session.flush() + folder = MailFolder( + tenant_id=tenant.id, + account_id=account.id, + name="INBOX", + imap_name="INBOX", + is_standard=True, + unread_count=0, + total_count=0, + ) + db_session.add(folder) + await db_session.flush() + mail = Mail( + tenant_id=tenant.id, + account_id=account.id, + folder_id=folder.id, + message_id="msg-1", + thread_id="thread-1", + from_address="sender@example.com", + to_addresses="recipient@example.com", + cc_addresses="", + bcc_addresses="", + subject="Test Mail", + body_text="Test body", + body_html="", + body_html_sanitized="", + is_seen=False, + is_flagged=False, + is_draft=False, + is_answered=False, + is_forwarded=False, + has_attachments=False, + size_bytes=0, + received_at=datetime.now(UTC), + ) + db_session.add(mail) + await db_session.commit() + + # Mock factory to return a session that uses the test engine + sf = async_sessionmaker(bind=db_session.bind, expire_on_commit=False, class_=AsyncSession) + mock_factory.return_value = sf + + await index_mails({}, [str(mail.id)]) + mock_index_entity.assert_called() + + +@pytest.mark.asyncio +@patch("app.plugins.builtins.unified_search.jobs.get_session_factory") +@patch("app.plugins.builtins.unified_search.embedding.index_entity", new_callable=AsyncMock, return_value=True) +async def test_index_contact(mock_index_entity, mock_factory, db_session: AsyncSession): + """index_contact calls index_entity.""" + from app.models.contact import Contact + from app.models.tenant import Tenant + from app.models.user import User + from app.core.auth import hash_password + + tenant = Tenant(name="Contact Tenant", slug="contact-tenant") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="contact@example.com", + name="Contact", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + contact = Contact( + tenant_id=tenant.id, + first_name="Index", + last_name="Contact", + email="indexcontact@example.com", + created_by=user.id, + updated_by=user.id, + ) + db_session.add(contact) + await db_session.commit() + + sf = async_sessionmaker(bind=db_session.bind, expire_on_commit=False, class_=AsyncSession) + mock_factory.return_value = sf + + await index_contact({}, str(contact.id)) + mock_index_entity.assert_called() + + +@pytest.mark.asyncio +@patch("app.plugins.builtins.unified_search.jobs.get_session_factory") +@patch("app.plugins.builtins.unified_search.embedding.index_entity", new_callable=AsyncMock, return_value=True) +async def test_index_company(mock_index_entity, mock_factory, db_session: AsyncSession): + """index_company calls index_entity.""" + from app.models.company import Company + from app.models.tenant import Tenant + from app.models.user import User + from app.core.auth import hash_password + + tenant = Tenant(name="Company Tenant", slug="company-tenant") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="company@example.com", + name="Company", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + company = Company( + tenant_id=tenant.id, + name="Index Company", + industry="IT", + created_by=user.id, + updated_by=user.id, + ) + db_session.add(company) + await db_session.commit() + + sf = async_sessionmaker(bind=db_session.bind, expire_on_commit=False, class_=AsyncSession) + mock_factory.return_value = sf + + await index_company({}, str(company.id)) + mock_index_entity.assert_called() + + +@pytest.mark.asyncio +@patch("app.plugins.builtins.unified_search.jobs.get_session_factory") +@patch("app.plugins.builtins.unified_search.embedding.index_entity", new_callable=AsyncMock, return_value=True) +async def test_reindex(mock_index_entity, mock_factory, db_session: AsyncSession): + """reindex iterates over all entities of a type.""" + from app.models.company import Company + from app.models.tenant import Tenant + from app.models.user import User + from app.core.auth import hash_password + + tenant = Tenant(name="Reindex Tenant", slug="reindex-tenant") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="reindex@example.com", + name="Reindex", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + company = Company( + tenant_id=tenant.id, + name="Reindex Co", + industry="IT", + created_by=user.id, + updated_by=user.id, + ) + db_session.add(company) + await db_session.commit() + + sf = async_sessionmaker(bind=db_session.bind, expire_on_commit=False, class_=AsyncSession) + mock_factory.return_value = sf + + await reindex({}, "company") + mock_index_entity.assert_called() + + +@pytest.mark.asyncio +@patch("app.plugins.builtins.unified_search.jobs.get_session_factory") +@patch("app.plugins.builtins.unified_search.embedding.index_entity", new_callable=AsyncMock, return_value=True) +async def test_embedding_batch(mock_index_entity, mock_factory, db_session: AsyncSession): + """embedding_batch finds entities without embedding.""" + from app.models.company import Company + from app.models.tenant import Tenant + from app.models.user import User + from app.core.auth import hash_password + + tenant = Tenant(name="Batch Tenant", slug="batch-tenant") + db_session.add(tenant) + await db_session.flush() + user = User( + tenant_id=tenant.id, + email="batch@example.com", + name="Batch", + password_hash=hash_password("TestPass123!"), + role="admin", + is_active=True, + preferences={}, + ) + db_session.add(user) + await db_session.flush() + company = Company( + tenant_id=tenant.id, + name="Batch Co", + industry="IT", + created_by=user.id, + updated_by=user.id, + ) + db_session.add(company) + await db_session.commit() + + sf = async_sessionmaker(bind=db_session.bind, expire_on_commit=False, class_=AsyncSession) + mock_factory.return_value = sf + + await embedding_batch({}) + mock_index_entity.assert_called() + + +# ─── Extended: Provider-Specific Tests ─── + + +@pytest.mark.asyncio +async def test_contact_provider_search_fts(db_session: AsyncSession): + """ContactProvider FTS search returns results.""" + provider = ContactSearchProvider() + result = await provider.search_fts(db_session, "test", uuid.uuid4(), 10) + assert isinstance(result, list) + + +@pytest.mark.asyncio +async def test_company_provider_search_fts(db_session: AsyncSession): + """CompanyProvider FTS search returns results.""" + provider = CompanySearchProvider() + result = await provider.search_fts(db_session, "test", uuid.uuid4(), 10) + assert isinstance(result, list) + + +@pytest.mark.asyncio +async def test_mail_provider_search_fts(db_session: AsyncSession): + """MailProvider FTS search returns results.""" + provider = MailSearchProvider() + result = await provider.search_fts(db_session, "test", uuid.uuid4(), 10) + assert isinstance(result, list) + + +@pytest.mark.asyncio +async def test_file_provider_search_fts(db_session: AsyncSession): + """FileProvider FTS search returns results.""" + provider = FileSearchProvider() + result = await provider.search_fts(db_session, "test", uuid.uuid4(), 10) + assert isinstance(result, list) + + +@pytest.mark.asyncio +async def test_event_provider_search_fts(db_session: AsyncSession): + """EventProvider FTS search returns results.""" + provider = EventSearchProvider() + result = await provider.search_fts(db_session, "test", uuid.uuid4(), 10) + assert isinstance(result, list) + + +def test_contact_provider_get_embedding_text(): + """ContactProvider get_embedding_text is a callable method.""" + provider = ContactSearchProvider() + assert hasattr(provider, "get_embedding_text") + assert provider.entity_type == "contact" + + +def test_company_provider_get_embedding_text(): + """CompanyProvider get_embedding_text is a callable method.""" + provider = CompanySearchProvider() + assert hasattr(provider, "get_embedding_text") + assert provider.entity_type == "company" + + +def test_mail_provider_get_embedding_text(): + """MailProvider get_embedding_text is a callable method.""" + provider = MailSearchProvider() + assert hasattr(provider, "get_embedding_text") + assert provider.entity_type == "mail" + + +def test_file_provider_get_embedding_text(): + """FileProvider get_embedding_text is a callable method.""" + provider = FileSearchProvider() + assert hasattr(provider, "get_embedding_text") + assert provider.entity_type == "file" + + +def test_event_provider_get_embedding_text(): + """EventProvider get_embedding_text is a callable method.""" + provider = EventSearchProvider() + assert hasattr(provider, "get_embedding_text") + assert provider.entity_type == "event" + + +def test_contact_provider_to_search_result(): + """ContactProvider to_search_result returns correct dict.""" + provider = ContactSearchProvider() + entity = {"id": "123", "first_name": "John", "last_name": "Doe", "email": "john@example.com"} + result = provider.to_search_result(entity) + assert result["entity_type"] == "contact" + assert result["entity_id"] == "123" + assert result["title"] == "John Doe" + assert result["snippet"] == "john@example.com" + assert result["score"] == 0.0 + + +def test_company_provider_to_search_result(): + """CompanyProvider to_search_result returns correct dict.""" + provider = CompanySearchProvider() + entity = {"id": "456", "name": "Acme Corp", "description": "IT company"} + result = provider.to_search_result(entity) + assert result["entity_type"] == "company" + assert result["entity_id"] == "456" + assert result["title"] == "Acme Corp" + assert "IT company" in result["snippet"] + + +def test_mail_provider_to_search_result(): + """MailProvider to_search_result returns correct dict.""" + provider = MailSearchProvider() + entity = {"id": "789", "subject": "Test Subject", "body_text": "Body text here"} + result = provider.to_search_result(entity) + assert result["entity_type"] == "mail" + assert result["entity_id"] == "789" + assert result["title"] == "Test Subject" + assert "Body text" in result["snippet"] + + +def test_file_provider_to_search_result(): + """FileProvider to_search_result returns correct dict.""" + provider = FileSearchProvider() + entity = {"id": "abc", "name": "document.pdf", "content_text": "File content"} + result = provider.to_search_result(entity) + assert result["entity_type"] == "file" + assert result["entity_id"] == "abc" + assert result["title"] == "document.pdf" + assert "File content" in result["snippet"] + + +def test_event_provider_to_search_result(): + """EventProvider to_search_result returns correct dict.""" + provider = EventSearchProvider() + entity = {"id": "evt1", "title": "Meeting", "description": "Team meeting"} + result = provider.to_search_result(entity) + assert result["entity_type"] == "event" + assert result["entity_id"] == "evt1" + assert result["title"] == "Meeting" + assert "Team meeting" in result["snippet"] + + +# ─── Extended: API Tests ─── + + +@pytest.mark.asyncio +async def test_ac13_stats_endpoint( + search_authed_client: tuple[AsyncClient, dict], +): + """AC13: GET /api/v1/search/stats returns 200 + statistics.""" + client, _ = search_authed_client + resp = await client.get( + "/api/v1/search/stats", + headers=ORIGIN_HEADER, + ) + assert resp.status_code == 200 + data = resp.json() + assert isinstance(data, dict) + # Stats should contain table keys or recent_logs + assert "recent_logs" in data + + +@pytest.mark.asyncio +async def test_ac14_toggle_provider( + search_authed_client: tuple[AsyncClient, dict], + db_session: AsyncSession, +): + """AC14: POST /api/v1/search/providers/{type}/toggle as admin returns 200.""" + client, seed = search_authed_client + # Pre-insert a provider row so the UPDATE path is taken (avoids INSERT without id) + from sqlalchemy import text as sql_text + from tests.conftest import _get_sync_engine + + sync_eng = _get_sync_engine() + with sync_eng.connect() as conn: + conn.execute( + sql_text( + "INSERT INTO unified_search_providers (id, tenant_id, entity_type, plugin_name, is_active, config) " + "VALUES (:id, :tid, 'contact', 'unified_search', true, '{}'::jsonb) " + "ON CONFLICT DO NOTHING" + ), + {"id": str(uuid.uuid4()), "tid": str(seed["tenant_a"].id)}, + ) + conn.commit() + sync_eng.dispose() + + resp = await client.post( + "/api/v1/search/providers/contact/toggle", + headers=ORIGIN_HEADER, + ) + assert resp.status_code == 200 + data = resp.json() + assert "entity_type" in data + assert "is_active" in data + assert data["entity_type"] == "contact" + + +@pytest.mark.asyncio +async def test_ac15_toggle_provider_viewer_403( + search_client: AsyncClient, db_session: AsyncSession +): + """AC15: POST /providers/{type}/toggle as viewer → 403.""" + await seed_tenant_and_users(db_session) + await login_client(search_client, "viewer@tenanta.com") + resp = await search_client.post( + "/api/v1/search/providers/contact/toggle", + headers=ORIGIN_HEADER, + ) + assert resp.status_code == 403