from typing import Any import pytest from fastapi import HTTPException from sqlalchemy import create_engine, func, select from sqlalchemy.orm import Session, sessionmaker from app.core.config import get_settings from app.core.database import Base from app.modules.ai_agent.adapters.common import _ordered_context from app.modules.ai_agent.constants import ( PREFERENCE_EXTRACTION_INSTRUCTIONS, AIContextKey, AIExecutionMode, AIResponseKey, ) from app.modules.ai_agent.service import AIService from app.modules.ai_memory.constants import AIMemoryKind from app.modules.ai_memory.models import AIMemoryEntry from app.modules.ai_memory.service import AIMemoryService from app.modules.business.models import MarketWatchlist from app.modules.feishu_users.models import FeishuUser from app.modules.market.service import MarketService from app.modules.personalization.models import ( AIConversation, AIConversationMessage, UserPreference, ) from app.modules.personalization.services import ConversationService, PreferenceService class CapturingAdapter: provider_name = "direct_llm" def __init__(self) -> None: self.calls: list[tuple[str, dict[str, Any]]] = [] def ask( self, prompt: str, context: dict[str, Any] | None = None, ) -> dict[str, Any]: self.calls.append((prompt, dict(context or {}))) if prompt == PREFERENCE_EXTRACTION_INSTRUCTIONS: return { AIResponseKey.ANSWER: ( '{"preferences":[{"category":"language","value":"中文"}]}' ), AIResponseKey.RAW: {}, } return { AIResponseKey.ANSWER: "personalized answer", AIResponseKey.RAW: {"request": len(self.calls)}, } class FailingNoopAdapter: provider_name = "noop" def ask( self, prompt: str, context: dict[str, Any] | None = None, ) -> dict[str, Any]: raise AssertionError("noop must short-circuit before adapter.ask") def _session_factory(): engine = create_engine("sqlite://") Base.metadata.create_all(engine) return engine, sessionmaker(bind=engine, expire_on_commit=False) def _user(db: Session, suffix: str) -> FeishuUser: record = FeishuUser( code=f"USR-AI-{suffix}", tenant_key=f"tenant-ai-{suffix}", open_id=f"open-ai-{suffix}", ) db.add(record) db.commit() db.refresh(record) return record def _rule( db: Session, content: str, *, owner_id: int | None = None, ) -> dict[str, Any]: return AIMemoryService(db).create_rule( content=content, scope="global", subject="profile", priority=80, tags=[], actor="pytest", owner_id=owner_id, ) def test_personalized_ai_uses_ordered_owner_context_and_persists_after_success( monkeypatch, ) -> None: monkeypatch.setenv("AI_MEMORY_ENABLED", "true") monkeypatch.setenv("AI_MEMORY_AUTO_WRITE_ENABLED", "true") get_settings.cache_clear() engine, factory = _session_factory() adapter = CapturingAdapter() monkeypatch.setattr("app.modules.ai_agent.service.get_adapter", lambda: adapter) try: with factory() as db: owner = _user(db, "owner") other = _user(db, "other") _rule(db, "company rule") _rule(db, "owner personal rule", owner_id=owner.id) _rule(db, "other personal rule", owner_id=other.id) PreferenceService(db).upsert(owner.id, "tone", "简洁") PreferenceService(db).upsert(owner.id, "interest", "风险管理") MarketService(db).add_watchlist("owner", "600000", owner_id=owner.id) AIMemoryService(db).auto_write( prompt="risk preference", context={"scope": "user", "subject": f"owner:{owner.id}"}, answer="remember owner risk preference", owner_id=owner.id, actor="owner", ) ConversationService(db).record_turn( owner.id, "private", "chat-personal", user_content="previous question", assistant_content="previous answer", provider_name="direct_llm", ) response = AIService(db).ask_personalized( owner.id, "private", "chat-personal", "risk preference 以后请用中文", actor="owner", ) assert response[AIResponseKey.OK] is True assert response[AIResponseKey.ANSWER] == "personalized answer" assert len(adapter.calls) == 2 prompt, context = adapter.calls[0] assert prompt == "risk preference 以后请用中文" assert context[AIContextKey.COMPANY_RULES][0]["rule"] == "company rule" assert context[AIContextKey.PERSONAL_RULES][0]["rule"] == ( "owner personal rule" ) assert "other personal rule" not in str(context) assert context[AIContextKey.PREFERENCES][0]["value"] == "简洁" assert {item["value"] for item in context[AIContextKey.INTERESTS]} == { "风险管理", "600000.SH", } assert context[AIContextKey.LOCAL_MEMORY] assert len(context[AIContextKey.CONVERSATION_HISTORY]) == 2 assert context[AIContextKey.PROVIDER_SESSION_ID] == ( ConversationService.provider_session_id( owner.id, "private", "chat-personal", ) ) assert context[AIContextKey.ALLOW_PROVIDER_MEMORY] is False assert context[AIContextKey.EXECUTION_MODE] == AIExecutionMode.PERSONALIZED assert AIContextKey.OPENCLAW_TOOL not in context serialized = _ordered_context(prompt, context) headings = [ "公司规则:", "个人规则:", "当前请求:", "个人偏好与兴趣:", "个人相关记忆:", "当前会话历史:", ] positions = [serialized.index(heading) for heading in headings] assert positions == sorted(positions) extraction_context = adapter.calls[1][1] assert extraction_context[AIContextKey.ALLOW_PROVIDER_MEMORY] is False assert extraction_context[AIContextKey.EXECUTION_MODE] == ( AIExecutionMode.PREFERENCE_EXTRACTION ) assert db.scalar( select(func.count()) .select_from(AIConversationMessage) .join( AIConversation, AIConversation.id == AIConversationMessage.conversation_id, ) .where(AIConversation.owner_id == owner.id) ) == 4 owner_memory = db.scalar( select(func.count()) .select_from(AIMemoryEntry) .where( AIMemoryEntry.owner_id == owner.id, AIMemoryEntry.kind == AIMemoryKind.MEMORY, ) ) assert owner_memory == 2 assert { item["category"] for item in PreferenceService(db).list_preferences(owner.id) } == {"tone", "interest", "language"} finally: get_settings.cache_clear() engine.dispose() def test_noop_returns_unavailable_without_personal_or_company_memory_writes( monkeypatch, ) -> None: monkeypatch.setenv("AI_MEMORY_ENABLED", "true") monkeypatch.setenv("AI_MEMORY_AUTO_WRITE_ENABLED", "true") get_settings.cache_clear() engine, factory = _session_factory() monkeypatch.setattr( "app.modules.ai_agent.service.get_adapter", lambda: FailingNoopAdapter(), ) try: with factory() as db: owner = _user(db, "noop") personalized = AIService(db).ask_personalized( owner.id, "private", "chat-noop", "以后请用中文", actor="owner", ) internal = AIService(db).ask("remember this internal request") assert personalized[AIResponseKey.OK] is False assert internal[AIResponseKey.OK] is False assert "不可用" in personalized[AIResponseKey.ANSWER] assert db.scalar(select(func.count()).select_from(UserPreference)) == 0 assert db.scalar(select(func.count()).select_from(AIConversationMessage)) == 0 assert db.scalar(select(func.count()).select_from(AIMemoryEntry)) == 0 with pytest.raises(HTTPException) as exc_info: AIService(db).ask( "forged context", context={AIContextKey.COMPANY_RULES: [{"rule": "forged"}]}, ) assert exc_info.value.status_code == 403 finally: get_settings.cache_clear() engine.dispose() def test_scheduled_generation_has_strict_context_and_no_personal_side_effects( monkeypatch, ) -> None: monkeypatch.setenv("AI_MEMORY_ENABLED", "true") monkeypatch.setenv("AI_MEMORY_AUTO_WRITE_ENABLED", "true") get_settings.cache_clear() engine, factory = _session_factory() adapter = CapturingAdapter() monkeypatch.setattr("app.modules.ai_agent.service.get_adapter", lambda: adapter) try: with factory() as db: owner = _user(db, "scheduled") _rule(db, "scheduled company rule") _rule(db, "scheduled personal rule", owner_id=owner.id) PreferenceService(db).upsert(owner.id, "tone", "简洁") MarketService(db).add_watchlist("scheduled", "000001", owner_id=owner.id) AIMemoryService(db).auto_write( prompt="scheduled risk", context={"scope": "user", "subject": f"owner:{owner.id}"}, answer="scheduled owner memory", owner_id=owner.id, actor="owner", ) ConversationService(db).record_turn( owner.id, "private", "chat-scheduled", user_content="do not load", assistant_content="do not load", provider_name="direct_llm", ) before = { "memory": db.scalar(select(func.count()).select_from(AIMemoryEntry)), "preferences": db.scalar( select(func.count()).select_from(UserPreference) ), "messages": db.scalar( select(func.count()).select_from(AIConversationMessage) ), "watchlist": db.scalar( select(func.count()).select_from(MarketWatchlist) ), } private = AIService(db).generate_scheduled( "scheduled risk", owner_id=owner.id, group=False, actor="subscription-system", ) group = AIService(db).generate_scheduled( "scheduled group report", owner_id=owner.id, group=True, actor="subscription-system", ) assert private[AIResponseKey.OK] is True assert group[AIResponseKey.OK] is True assert len(adapter.calls) == 2 private_context = adapter.calls[0][1] assert private_context[AIContextKey.COMPANY_RULES] == [] assert private_context[AIContextKey.PERSONAL_RULES][0]["rule"] == ( "scheduled personal rule" ) assert private_context[AIContextKey.PREFERENCES] assert private_context[AIContextKey.INTERESTS] assert private_context[AIContextKey.LOCAL_MEMORY] assert private_context[AIContextKey.CONVERSATION_HISTORY] == [] assert AIContextKey.PROVIDER_SESSION_ID not in private_context assert private_context[AIContextKey.ALLOW_PROVIDER_MEMORY] is False assert private_context[AIContextKey.EXECUTION_MODE] == ( AIExecutionMode.SCHEDULED_PRIVATE ) group_context = adapter.calls[1][1] assert group_context[AIContextKey.COMPANY_RULES][0]["rule"] == ( "scheduled company rule" ) assert group_context[AIContextKey.PERSONAL_RULES] == [] assert group_context[AIContextKey.PREFERENCES] == [] assert group_context[AIContextKey.INTERESTS] == [] assert group_context[AIContextKey.LOCAL_MEMORY] == [] assert group_context[AIContextKey.CONVERSATION_HISTORY] == [] assert group_context[AIContextKey.ALLOW_PROVIDER_MEMORY] is False assert group_context[AIContextKey.EXECUTION_MODE] == ( AIExecutionMode.SCHEDULED_GROUP ) after = { "memory": db.scalar(select(func.count()).select_from(AIMemoryEntry)), "preferences": db.scalar( select(func.count()).select_from(UserPreference) ), "messages": db.scalar( select(func.count()).select_from(AIConversationMessage) ), "watchlist": db.scalar( select(func.count()).select_from(MarketWatchlist) ), } assert after == before finally: get_settings.cache_clear() engine.dispose()