diff --git a/.env.example b/.env.example index 7f2ea52..ec70d97 100644 --- a/.env.example +++ b/.env.example @@ -12,7 +12,14 @@ GEMINI_API_KEY=your_api_key_here # Scholar review (reviewer endpoints are disabled until the token is set) # SCHOLAR_REVIEW_TOKEN=generate_a_long_random_value # REVIEW_EXPORT_PATH=data/review/reviewed.jsonl -# REDIS_URL=redis://localhost:6379 # makes the review queue durable +# Scholar review (reviewer endpoints are disabled until the token is set) +# SCHOLAR_REVIEW_TOKEN=generate_a_long_random_value +# REVIEW_EXPORT_PATH=data/review/reviewed.jsonl +# REDIS_URL=redis://localhost:6379 # durable review queue + memory store + +# Per-user memory (optional โ€” defaults shown) +# MEMORY_TTL_DAYS=90 # user-profile and chat-summary TTL in days +# MEMORY_EXTRACTION_ENABLED=true # background extraction from conversation turns # Tafsir layer (optional โ€” sensible defaults, no key required) # QURAN_API_BASE=https://api.quran.com/api/v4 diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index b5e438a..9da3257 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -25,13 +25,13 @@ jobs: run: | python -m pip install --upgrade pip pip install -r requirements.txt - pip install pytest flake8 + pip install pytest flake8 pytest-asyncio - name: Run linting - run: flake8 main.py stellar.py nisab.py safety telemetry.py tests/redteam study.py fiqh.py hadith.py confidence.py review.py review_store.py tafsir.py semantic_cache.py tests/test_confidence.py tests/test_review_queue.py tests/test_tafsir.py tests/test_zakat.py tests/test_telemetry.py tests/test_multilingual.py scripts/build_hadith_data.py scripts/build_surah_index.py --max-line-length=120 --ignore=E501,W503 + run: flake8 main.py memory stellar.py nisab.py safety telemetry.py tests/redteam study.py fiqh.py hadith.py confidence.py review.py review_store.py tafsir.py semantic_cache.py tests/test_confidence.py tests/test_review_queue.py tests/test_tafsir.py tests/test_zakat.py tests/test_telemetry.py tests/test_multilingual.py tests/test_memory_profile.py tests/test_memory_extraction.py tests/test_memory_integration.py scripts/build_hadith_data.py scripts/build_surah_index.py --max-line-length=120 --ignore=E501,W503 - name: Check syntax - run: python -m compileall -q main.py stellar.py nisab.py safety telemetry.py tests/redteam study.py fiqh.py hadith.py confidence.py review.py review_store.py tafsir.py semantic_cache.py scripts/build_hadith_data.py scripts/build_surah_index.py + run: python -m compileall -q main.py memory stellar.py nisab.py safety telemetry.py tests/redteam study.py fiqh.py hadith.py confidence.py review.py review_store.py tafsir.py semantic_cache.py scripts/build_hadith_data.py scripts/build_surah_index.py - name: Run offline safety and red-team tests run: pytest -q tests/redteam @@ -66,6 +66,15 @@ jobs: - name: Run multilingual support tests run: pytest -q tests/test_multilingual.py + - name: Run memory profile tests + run: pytest -q tests/test_memory_profile.py + + - name: Run memory extraction tests + run: pytest -q tests/test_memory_extraction.py + + - name: Run memory integration tests + run: pytest -q tests/test_memory_integration.py + docker-build: name: Docker Build runs-on: ubuntu-latest diff --git a/README.md b/README.md index c525dd0..f037960 100644 --- a/README.md +++ b/README.md @@ -36,6 +36,8 @@ The platform is composed of three services: - ๐Ÿงต **Conversation history** per chat session - ๐Ÿ›ก๏ธ **Content safety filters** on model output - ๐ŸŽš๏ธ **Confidence-aware answers** โ€” abstains or hedges instead of guessing, and routes doubtful religious answers to a scholar +- ๐Ÿง  **Per-user long-term memory** โ€” user profiles (knowledge level, madhhab, topics studied, remembered facts) extracted from conversations and injected across sessions; privacy controls with GET/DELETE endpoints and `remember` opt-out per request +- ๐Ÿ“‹ **Conversation summarization** โ€” compaction API ready for token-budget-triggered eviction; merges and recompresses summaries when history exceeds budget - ๐Ÿ“– **Tafsir-grounded ayah explanations** โ€” retrieved from named classical works, never paraphrased from model memory - โšก **FastAPI** with automatic OpenAPI docs at `/docs` @@ -45,6 +47,8 @@ The platform is composed of three services: |--------|-------|---------| | `POST` | `/chat` | Start or continue a chat session | | `DELETE` | `/chat/{chat_id}` | Delete a chat session | +| `GET` | `/memory/{user_id}` | Retrieve a stored user profile (transparency) | +| `DELETE` | `/memory/{user_id}` | Completely erase a stored user profile | | `GET` | `/ping` | Health check | | `GET` | `/cache/stats` | Semantic cache metrics (hits, misses, hit rate, etc.) | | `POST` | `/tafsir` | Ayah explanation from named tafsir works, with attribution | @@ -135,7 +139,9 @@ services: | `CONFIDENCE_UNVERIFIED_CEILING` | Cap when nothing external corroborated the answer | `0.65` | | `SCHOLAR_REVIEW_TOKEN` | Enables the reviewer endpoints; required as `X-Review-Token` | โ€” (endpoints disabled) | | `REVIEW_EXPORT_PATH` | JSONL export of reviewed answers | `data/review/reviewed.jsonl` | -| `REDIS_URL` | Makes the scholar-review queue durable across restarts | โ€” (in-memory) | +| `REDIS_URL` | Makes the scholar-review queue and memory store durable across restarts | โ€” (in-memory) | +| `MEMORY_TTL_DAYS` | Time-to-live for stored user profiles and chat summaries in days | `90` | +| `MEMORY_EXTRACTION_ENABLED` | Background memory extraction from conversation turns | `true` | | `STELLAR_NETWORK` | Stellar network for zakat lookups (`testnet` or `public`) | `testnet` | | `ZAKAT_NISAB_USD` | Fallback nisab when no gold price can be fetched | `6000` | | `NISAB_CACHE_TTL_SECONDS` | How long a fetched gold price is reused | `21600` (6h) | diff --git a/main.py b/main.py index f9f018b..7f3487a 100644 --- a/main.py +++ b/main.py @@ -8,7 +8,7 @@ from fastapi import FastAPI, HTTPException, Request, Response from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import StreamingResponse -from pydantic import BaseModel +from pydantic import BaseModel, Field import google.generativeai as genai import time @@ -62,6 +62,15 @@ from review import enqueue_for_review, router as review_router from review_store import get_review_store +from memory import ChatSummary, UserProfile, create_memory_store, render_user_context +from memory.extraction import ( + MEMORY_EXTRACTION_ENABLED, + apply_updates, + extract_updates, + merge_summaries, + summarize_conversation_turns, +) + logger = logging.getLogger(__name__) # Load environment variables @@ -115,6 +124,8 @@ class ChatRequest(BaseModel): context: Optional[str] = None # Additional context for specific queries madhhab: Optional[str] = None # User's madhhab: hanafi, maliki, shafii, hanbali language: Optional[str] = None # BCP-47 response language (ar, en, ur, etc.); auto-detect when omitted + user_id: Optional[str] = Field(default=None, max_length=128) # Opaque user identifier for personalization + remember: bool = True # When False, existing memory is read but no new data persisted class Message(BaseModel): @@ -183,6 +194,11 @@ def classify_for_safety(prompt: str, candidate_ids: List[str]): # Durable queue for low-confidence religious answers awaiting a scholar review_store = get_review_store() +# Per-user memory store (Redis-backed or in-memory) +memory_store = create_memory_store() + +MAX_CHAT_HISTORY_TURNS = 20 + # Tafsir retrieval seam: returns None for prompts that are not # verse-explanation questions. Offline tests replace this with a stub. DEFAULT_TAFSIR_LANGUAGE = "en" @@ -466,6 +482,13 @@ def _finalize() -> None: zakat_context = await zakat_retriever(request.prompt, request.context) zakat_info = zakat_context.info if zakat_context else None + # --- Memory lookup --- + profile: Optional[UserProfile] = None + summary: Optional[ChatSummary] = None + if request.user_id: + profile = await memory_store.get_profile(request.user_id) + summary = await memory_store.get_chat_summary(f"{request.user_id}:{chat_id}") + # Neither a tafsir-grounded answer nor a zakat answer goes through the # semantic response cache: the first is built from retrieved passages # (already cached by ayah key), and the second contains one user's real @@ -475,6 +498,7 @@ def _finalize() -> None: and request.context is None and tafsir_context is None and zakat_context is None + and request.user_id is None and SEMANTIC_CACHE_ENABLED ) @@ -526,6 +550,9 @@ async def generate(safety_prompt: str) -> str: system_context += tafsir_system_context(tafsir_context) if zakat_context is not None: system_context += zakat_context.prompt_block + memory_block = render_user_context(profile, summary) + if memory_block: + system_context += f"\n\n{memory_block}" context = f"Additional context: {extra_context}\n\n" if extra_context else "" full_prompt = f"{system_context}\n{context}User question: {safety_prompt}" logger.info("Sending message to chat...") @@ -682,6 +709,33 @@ async def generate(safety_prompt: str) -> str: ) _finalize() _succeeded = True + + # --- Background memory extraction and summarization --- + # Runs as fire-and-forget tasks after the response is sent. + if request.user_id and request.remember and MEMORY_EXTRACTION_ENABLED: + asyncio.create_task( + _extract_and_update_memory( + request.user_id, prompt, response_text, chat_id, summary, memory_store, + ) + ) + logger.info("Memory extraction scheduled for user %s", request.user_id[:8]) + + # --- Summary eviction --- + # After enough turns accumulate, summarize old history and persist. + if request.user_id and request.remember and MEMORY_EXTRACTION_ENABLED: + chat_session = active_chats.get(chat_id) + if chat_session and hasattr(chat_session, "history") and chat_session.history: + if len(chat_session.history) >= MAX_CHAT_HISTORY_TURNS: + asyncio.create_task( + _summarize_history( + f"{request.user_id}:{chat_id}", + chat_session.history, + summary, + memory_store, + ) + ) + logger.info("History summarization triggered for %s", request.user_id[:8]) + return response_obj except ResourceExhausted as exc: @@ -734,6 +788,47 @@ async def generate(safety_prompt: str) -> str: telemetry.current_trace.reset(_ctx_token) +async def _extract_and_update_memory( + user_id: str, prompt: str, response: str, chat_id: str, + existing_summary: Optional[ChatSummary], + store: Any, +) -> None: + """Fire-and-forget memory extraction. Runs via asyncio.create_task.""" + try: + updates = await extract_updates(prompt, response) + if updates.get("none"): + return + profile = await store.get_profile(user_id) + if profile is None: + profile = UserProfile(user_id=user_id) + profile = apply_updates(profile, updates) + await store.save_profile(user_id, profile) + logger.debug("Memory updated for user %s", user_id[:8]) + except Exception: + logger.warning("Memory extraction failed for user %s", user_id[:8], exc_info=True) + + +async def _summarize_history( + chat_id: str, history: list, existing_summary: Optional[ChatSummary], + store: Any, +) -> None: + """Summarize accumulated conversation turns and persist.""" + try: + turns = [ + {"role": m.role, "text": m.parts[0].text if m.parts else ""} + for m in history + ] + new_summary_text = await summarize_conversation_turns(turns) + if existing_summary: + merged = await merge_summaries(existing_summary.content, new_summary_text) + else: + merged = new_summary_text + summary = ChatSummary(chat_id=chat_id, content=merged, turn_count=len(history)) + await store.save_chat_summary(chat_id, summary) + except Exception: + logger.warning("History summarization failed for %s", chat_id[:8], exc_info=True) + + @app.post("/chat/stream") async def chat_stream(request: ChatRequest, http_request: Request): """Streaming chat endpoint using Server-Sent Events (SSE).""" @@ -847,9 +942,8 @@ async def event_generator(): ) except Exception as e: - error_msg = f"โŒ Streaming Chat API Error: {str(e)}" - logger.error(error_msg) - raise HTTPException(status_code=500, detail=error_msg) + logger.error("Streaming Chat API Error", exc_info=True) + raise HTTPException(status_code=500, detail="Internal server error") from e @app.delete("/chat/{chat_id}") @@ -861,9 +955,8 @@ async def delete_chat(chat_id: str): return {"message": "Chat session deleted successfully"} return {"message": "Chat session not found"} except Exception as e: - error_msg = f"โŒ Error deleting chat: {str(e)}" - logger.error(error_msg) - raise HTTPException(status_code=500, detail=error_msg) + logger.error("Error deleting chat", exc_info=True) + raise HTTPException(status_code=500, detail="Internal server error") from e @app.get("/ping") @@ -872,6 +965,33 @@ async def ping(): return {"status": "ok"} +@app.get("/memory/{user_id}") +async def get_memory(user_id: str): + """Retrieve the stored user profile for transparency. + + TODO(#9): bind to authenticated principal โ€” anyone who knows a user_id + can currently read another user's memory. + """ + profile = await memory_store.get_profile(user_id) + if profile is None: + raise HTTPException(status_code=404, detail="Memory not found") + return profile.model_dump() + + +@app.delete("/memory/{user_id}") +async def delete_memory(user_id: str): + """Completely erase the stored user profile. + + TODO(#9): bind to authenticated principal โ€” anyone who knows a user_id + can currently erase another user's memory. + """ + existed = await memory_store.delete_profile(user_id) + if existed: + logger.info("Deleted memory for user %s", user_id[:8]) + return {"message": "Memory deleted successfully"} + return {"message": "Memory not found"} + + @app.get("/cache/stats") async def cache_stats(): return semantic_cache.get_stats() diff --git a/memory/__init__.py b/memory/__init__.py new file mode 100644 index 0000000..9499514 --- /dev/null +++ b/memory/__init__.py @@ -0,0 +1,88 @@ +"""Per-user long-term memory and conversation summarization. + +See README.md for usage; see tests/ for offline-verifiable contracts. +""" + +from __future__ import annotations + +import logging +import os +from typing import Optional + +from memory.models import ChatSummary, UserProfile +from memory.store import ( + InMemoryMemoryStore, + MemoryStore, + RedisMemoryStore, +) + +logger = logging.getLogger(__name__) + + +def create_memory_store() -> MemoryStore: + """Factory: ``REDIS_URL`` set โ†’ ``RedisMemoryStore``, else in-memory. + + When Redis is configured but fails at startup the error is surfaced + (logged + raised) โ€” workers must not silently diverge on user memory. + """ + url = os.getenv("REDIS_URL", "") + if url: + logger.info("MemoryStore using Redis at %s", url) + return RedisMemoryStore(url) + logger.info("MemoryStore using in-memory dict (local development)") + return InMemoryMemoryStore() + + +def render_user_context( + profile: Optional[UserProfile], + summary: Optional[ChatSummary], +) -> str: + """Render profile and chat summary as a delimited DATA block. + + Returns an empty string when neither has content so anonymous traffic + is completely unaffected. + """ + parts: list[str] = [] + + if profile is not None and ( + profile.knowledge_level + or profile.madhhab + or profile.preferred_language + or profile.topics_studied + or profile.remembered_facts + ): + lines = ["--- Known about this student ---"] + if profile.knowledge_level: + lines.append(f"Knowledge level: {profile.knowledge_level}") + if profile.madhhab: + lines.append(f"Madhhab: {profile.madhhab}") + if profile.preferred_language: + lines.append(f"Preferred language: {profile.preferred_language}") + if profile.topics_studied: + topics_str = ", ".join( + f"{t.topic}" for t in profile.topics_studied[-10:] + ) + lines.append(f"Topics studied: {topics_str}") + if profile.remembered_facts: + for fact in profile.remembered_facts[-5:]: + lines.append(f"- {fact.fact}") + parts.append("\n".join(lines)) + + if summary is not None and summary.content: + parts.append(f"--- Conversation summary ---\n{summary.content}") + + if not parts: + return "" + + return "\n\n".join(parts) + "\n---------------------------------\n" + + +__all__ = [ + "ChatSummary", + "InMemoryMemoryStore", + "MemoryStore", + "RedisMemoryStore", + "UserProfile", + "create_memory_store", + "render_user_context", +] diff --git a/memory/extraction.py b/memory/extraction.py new file mode 100644 index 0000000..ff7ce83 --- /dev/null +++ b/memory/extraction.py @@ -0,0 +1,237 @@ +"""Memory extraction, validation, and conversation summarization. + +All Gemini calls follow the existing pattern from classify_for_safety() and +study.py: genai.GenerativeModel + generate_content with temperature=0 and +response_mime_type="application/json". + +Seams +----- +Every Gemini-backed function can be monkeypatched in offline tests. +extract_updates โ†’ patch memory.extraction._call_extraction_gemini +summarize_conversation_turns โ†’ patch memory.extraction._call_summary_gemini +recompress_summaries โ†’ patch memory.extraction._call_recompress_gemini +""" + +from __future__ import annotations + +import json +import logging +import os +from typing import Optional + +from memory.models import ( + MAX_FACT_LENGTH, + MAX_FACTS, + MAX_SUMMARY_LENGTH, + MAX_TOPIC_LENGTH, + MAX_TOPICS, + FactEntry, + TopicEntry, + UserProfile, +) + +logger = logging.getLogger(__name__) + +MEMORY_EXTRACTION_ENABLED = os.getenv( + "MEMORY_EXTRACTION_ENABLED", "true" +).lower() not in {"0", "false", "off"} + +_EXTRACTION_INSTRUCTION = """You are an AI assistant that extracts structured profile updates from a conversation turn. + +Given the user's question and your response, identify any of the following that apply: +- knowledge_level: one of "beginner", "intermediate", "advanced" (or omit) +- madhhab: the user's school of thought (e.g. "hanafi", "maliki", "shafii", "hanbali") (or omit) +- preferred_language: e.g. "english", "arabic", "urdu" (or omit) +- new_facts: a list of specific facts the user shared about themselves (each max 500 chars) +- new_topics: a list of topics the user asked about (each max 100 chars) + +Return ONLY strict JSON with exactly the keys that apply. +If nothing meaningful changed, return {"none": true}.""" + + +def apply_updates(profile: UserProfile, updates: dict) -> UserProfile: + """Validate and apply structured extraction proposals. + + Invalid individual entries are **rejected** (logged, skipped). + Valid entries that cause a collection to exceed its cap trigger + oldest-first eviction of that collection. + """ + from fiqh import normalize_madhhab + + profile = profile.model_copy(deep=True) + + knowledge_level = updates.get("knowledge_level") + if knowledge_level is not None: + from memory.models import VALID_KNOWLEDGE_LEVELS + if knowledge_level in VALID_KNOWLEDGE_LEVELS: + profile.knowledge_level = knowledge_level + else: + logger.warning("Rejected invalid knowledge_level: %s", knowledge_level) + + madhhab = updates.get("madhhab") + if madhhab is not None: + normalized = normalize_madhhab(madhhab) + if normalized is not None: + profile.madhhab = normalized + else: + logger.warning("Rejected invalid madhhab: %s", madhhab) + + language = updates.get("preferred_language") + if language is not None: + cleaned = language.strip().lower() + if len(cleaned) <= 20: + profile.preferred_language = cleaned + else: + logger.warning("Rejected oversized language: %s", language) + + new_facts = updates.get("new_facts", []) + for fact_text in new_facts: + if not isinstance(fact_text, str) or len(fact_text) > MAX_FACT_LENGTH or not fact_text.strip(): + logger.warning("Rejected invalid fact: %s", str(fact_text)[:80]) + continue + profile.remembered_facts.append(FactEntry(fact=fact_text.strip(), created_at=__import__("time").time())) + while len(profile.remembered_facts) > MAX_FACTS: + profile.remembered_facts.pop(0) + + new_topics = updates.get("new_topics", []) + for topic_name in new_topics: + if not isinstance(topic_name, str) or len(topic_name) > MAX_TOPIC_LENGTH or not topic_name.strip(): + logger.warning("Rejected invalid topic: %s", str(topic_name)[:80]) + continue + cleaned_topic = topic_name.strip().lower() + existing = [t for t in profile.topics_studied if t.topic == cleaned_topic] + if existing: + existing[0].last_asked = __import__("time").time() + else: + profile.topics_studied.append(TopicEntry(topic=cleaned_topic, last_asked=__import__("time").time())) + while len(profile.topics_studied) > MAX_TOPICS: + profile.topics_studied.pop(0) + + profile.updated_at = __import__("time").time() + return profile + + +async def extract_updates(user_prompt: str, model_response: str) -> dict: + """Propose profile updates from a conversation turn. + + Offline seam: patch ``memory.extraction._call_extraction_gemini``. + """ + full_prompt = ( + f"{_EXTRACTION_INSTRUCTION}\n\n" + f"User question: {user_prompt}\n" + f"Assistant response: {model_response}" + ) + raw = await _call_extraction_gemini(full_prompt) + if not isinstance(raw, dict): + logger.warning("Extraction returned non-dict: %s", type(raw).__name__) + return {"none": True} + return raw + + +async def _call_extraction_gemini(prompt: str) -> dict: + """Gemini structured-output call for extraction. + + Override this in offline tests with a fixture. + """ + import google.generativeai as genai + + model = genai.GenerativeModel( + "gemini-2.5-flash-preview-05-20", + system_instruction=_EXTRACTION_INSTRUCTION, + ) + response = await model.generate_content_async( + prompt, + generation_config={"temperature": 0, "response_mime_type": "application/json"}, + request_options={"timeout": 30}, + ) + return json.loads(response.text) + + +_SUMMARY_INSTRUCTION = """Summarize the following conversation turns. +Preserve: +- Established facts about the user +- The user's goals and intent +- Unresolved or open questions +- Topics discussed + +Be concise. Maximum 2000 characters.""" + + +async def summarize_conversation_turns(evicted_turns: list[dict[str, str]]) -> str: + """Summarize a list of evicted conversation turns. + + Each dict has keys ``role`` and ``text``. + + Offline seam: patch ``memory.extraction._call_summary_gemini``. + """ + turns_text = "\n".join( + f"{t.get('role', 'unknown')}: {t.get('text', '')}" + for t in evicted_turns + ) + prompt = f"{_SUMMARY_INSTRUCTION}\n\nTurns:\n{turns_text}" + return await _call_summary_gemini(prompt) + + +async def _call_summary_gemini(prompt: str) -> str: + """Gemini structured-output call for summarization. + + Override this in offline tests with a fixture. + """ + import google.generativeai as genai + + model = genai.GenerativeModel( + "gemini-2.5-flash-preview-05-20", + system_instruction=_SUMMARY_INSTRUCTION, + ) + response = await model.generate_content_async( + prompt, + generation_config={"temperature": 0}, + request_options={"timeout": 30}, + ) + return response.text + + +def merge_summaries_deterministic(existing: str, new: str) -> Optional[str]: + """Concatenate two summaries if the result fits within MAX_SUMMARY_LENGTH.""" + combined = f"{existing.rstrip()}\n{new}".strip() + if len(combined) <= MAX_SUMMARY_LENGTH: + return combined + return None + + +_RECOMPRESS_INSTRUCTION = """Merge the following two conversation summaries into one. +Preserve established facts, user goals, and open questions. +Remove redundancy but keep all unique information. +Be concise. Maximum 2000 characters.""" + + +async def _call_recompress_gemini(existing: str, new: str) -> str: + """Gemini call for recompressing two oversized summaries. + + Offline seam: patch ``memory.extraction._call_recompress_gemini``. + """ + import google.generativeai as genai + + prompt = ( + f"{_RECOMPRESS_INSTRUCTION}\n\n" + f"Existing summary:\n{existing}\n\n" + f"New segment:\n{new}" + ) + model = genai.GenerativeModel( + "gemini-2.5-flash-preview-05-20", + system_instruction=_RECOMPRESS_INSTRUCTION, + ) + response = await model.generate_content_async( + prompt, + generation_config={"temperature": 0}, + request_options={"timeout": 30}, + ) + return response.text + + +async def merge_summaries(existing: str, new: str) -> str: + """Merge two summaries โ€” deterministic concatenation first, Gemini on overflow.""" + result = merge_summaries_deterministic(existing, new) + if result is not None: + return result + return await _call_recompress_gemini(existing, new) diff --git a/memory/models.py b/memory/models.py new file mode 100644 index 0000000..8a0e845 --- /dev/null +++ b/memory/models.py @@ -0,0 +1,48 @@ +"""Pydantic models for per-user memory and per-chat summaries.""" + +from __future__ import annotations + +import time +from typing import Optional + +from pydantic import BaseModel, Field + +MAX_FACTS = 20 +MAX_TOPICS = 30 +MAX_FACT_LENGTH = 500 +MAX_TOPIC_LENGTH = 100 +MAX_SUMMARY_LENGTH = 2000 +VALID_KNOWLEDGE_LEVELS = frozenset({"beginner", "intermediate", "advanced"}) + + +class TopicEntry(BaseModel): + topic: str = Field(max_length=MAX_TOPIC_LENGTH) + last_asked: float + + +class FactEntry(BaseModel): + fact: str = Field(max_length=MAX_FACT_LENGTH) + created_at: float + + +class UserProfile(BaseModel): + user_id: str + knowledge_level: Optional[str] = None + madhhab: Optional[str] = None + preferred_language: Optional[str] = None + topics_studied: list[TopicEntry] = Field(default_factory=list) + remembered_facts: list[FactEntry] = Field(default_factory=list) + created_at: float = Field(default_factory=time.time) + updated_at: float = Field(default_factory=time.time) + + model_config = {"extra": "forbid"} + + +class ChatSummary(BaseModel): + chat_id: str + content: str = Field(max_length=MAX_SUMMARY_LENGTH) + turn_count: int = 0 + created_at: float = Field(default_factory=time.time) + updated_at: float = Field(default_factory=time.time) + + model_config = {"extra": "forbid"} diff --git a/memory/store.py b/memory/store.py new file mode 100644 index 0000000..f5923b7 --- /dev/null +++ b/memory/store.py @@ -0,0 +1,150 @@ +"""Memory store abstraction โ€” Redis-backed with in-memory fallback. + +REDIS_URL absent โ†’ InMemoryMemoryStore (process-local, lost on restart). +REDIS_URL present โ†’ RedisMemoryStore. Connection failures surface + (log + raise) โ€” never silently switch to in-memory, + because two workers would diverge on user memory. +""" + +from __future__ import annotations + +import abc +import json +import logging +import os +import time +from typing import Optional + +from memory.models import ChatSummary, UserProfile + +logger = logging.getLogger(__name__) + +MEMORY_TTL_SECONDS = int(os.getenv("MEMORY_TTL_DAYS", "90")) * 86400 +REDIS_URL = os.getenv("REDIS_URL", "") + +try: + import redis.asyncio as aioredis + _redis_available = True +except ImportError: + _redis_available = False + + +def _profile_key(user_id: str) -> str: + return f"memory:profile:{user_id}" + + +def _summary_key(chat_id: str) -> str: + return f"memory:summary:{chat_id}" + + +class MemoryStore(abc.ABC): + @abc.abstractmethod + async def get_profile(self, user_id: str) -> Optional[UserProfile]: + ... + + @abc.abstractmethod + async def save_profile(self, user_id: str, profile: UserProfile) -> None: + ... + + @abc.abstractmethod + async def delete_profile(self, user_id: str) -> bool: + ... + + @abc.abstractmethod + async def get_chat_summary(self, chat_id: str) -> Optional[ChatSummary]: + ... + + @abc.abstractmethod + async def save_chat_summary(self, chat_id: str, summary: ChatSummary) -> None: + ... + + @abc.abstractmethod + async def delete_chat_summary(self, chat_id: str) -> bool: + ... + + +class InMemoryMemoryStore(MemoryStore): + def __init__(self) -> None: + self._profiles: dict[str, tuple[float, UserProfile]] = {} + self._summaries: dict[str, tuple[float, ChatSummary]] = {} + + async def get_profile(self, user_id: str) -> Optional[UserProfile]: + entry = self._profiles.get(user_id) + if entry is None: + return None + expires_at, profile = entry + if time.monotonic() > expires_at: + del self._profiles[user_id] + return None + return profile + + async def save_profile(self, user_id: str, profile: UserProfile) -> None: + self._profiles[user_id] = (time.monotonic() + MEMORY_TTL_SECONDS, profile) + + async def delete_profile(self, user_id: str) -> bool: + return self._profiles.pop(user_id, None) is not None + + async def get_chat_summary(self, chat_id: str) -> Optional[ChatSummary]: + entry = self._summaries.get(chat_id) + if entry is None: + return None + expires_at, summary = entry + if time.monotonic() > expires_at: + del self._summaries[chat_id] + return None + return summary + + async def save_chat_summary(self, chat_id: str, summary: ChatSummary) -> None: + self._summaries[chat_id] = (time.monotonic() + MEMORY_TTL_SECONDS, summary) + + async def delete_chat_summary(self, chat_id: str) -> bool: + return self._summaries.pop(chat_id, None) is not None + + +class RedisMemoryStore(MemoryStore): + def __init__(self, redis_url: str) -> None: + if not _redis_available: + raise RuntimeError("redis package not installed") + self._redis = aioredis.from_url(redis_url, decode_responses=True) + + async def get_profile(self, user_id: str) -> Optional[UserProfile]: + raw = await self._redis.get(_profile_key(user_id)) + if raw is None: + return None + try: + return UserProfile.model_validate(json.loads(raw)) + except (json.JSONDecodeError, Exception): + logger.warning("Corrupt profile for user %s", user_id) + return None + + async def save_profile(self, user_id: str, profile: UserProfile) -> None: + await self._redis.setex( + _profile_key(user_id), + MEMORY_TTL_SECONDS, + profile.model_dump_json(), + ) + + async def delete_profile(self, user_id: str) -> bool: + deleted = await self._redis.delete(_profile_key(user_id)) + return deleted > 0 + + async def get_chat_summary(self, chat_id: str) -> Optional[ChatSummary]: + raw = await self._redis.get(_summary_key(chat_id)) + if raw is None: + return None + try: + return ChatSummary.model_validate(json.loads(raw)) + except (json.JSONDecodeError, Exception): + logger.warning("Corrupt chat summary for chat %s", chat_id) + return None + + async def save_chat_summary(self, chat_id: str, summary: ChatSummary) -> None: + await self._redis.setex( + _summary_key(chat_id), + MEMORY_TTL_SECONDS, + summary.model_dump_json(), + ) + + async def delete_chat_summary(self, chat_id: str) -> bool: + deleted = await self._redis.delete(_summary_key(chat_id)) + return deleted > 0 diff --git a/requirements.txt b/requirements.txt index b6776a1..26a022b 100644 --- a/requirements.txt +++ b/requirements.txt @@ -6,5 +6,5 @@ pydantic>=2.5.3 stellar-sdk>=12.0.0 PyYAML==6.0.2 numpy==2.2.4 -redis>=5.0.0 +redis>=5.0.0,<6.0.0 httpx==0.28.1 diff --git a/tests/test_memory_extraction.py b/tests/test_memory_extraction.py new file mode 100644 index 0000000..9e2507e --- /dev/null +++ b/tests/test_memory_extraction.py @@ -0,0 +1,161 @@ +"""Tests for memory extraction, validation, and summarization โ€” no live Gemini. + +All Gemini-backed functions are monkeypatched with fixtures. +""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, patch + +import pytest + +from memory.extraction import ( + MAX_FACT_LENGTH, + MAX_FACTS, + MAX_SUMMARY_LENGTH, + MAX_TOPICS, + apply_updates, + merge_summaries, + merge_summaries_deterministic, + summarize_conversation_turns, +) +from memory.models import UserProfile + + +# --------------------------------------------------------------------------- +# apply_updates +# --------------------------------------------------------------------------- + + +def _empty_profile() -> UserProfile: + return UserProfile(user_id="test-user") + + +class TestApplyUpdates: + def test_empty_updates_is_noop(self): + profile = _empty_profile() + result = apply_updates(profile, {}) + assert result.knowledge_level is None + assert result.madhhab is None + + def test_knowledge_level_accepted(self): + profile = apply_updates(_empty_profile(), {"knowledge_level": "beginner"}) + assert profile.knowledge_level == "beginner" + + def test_invalid_knowledge_level_rejected(self): + profile = apply_updates(_empty_profile(), {"knowledge_level": "expert"}) + assert profile.knowledge_level is None + + def test_madhhab_accepted(self): + profile = apply_updates(_empty_profile(), {"madhhab": "shafii"}) + assert profile.madhhab == "shafii" + + def test_invalid_madhhab_rejected(self): + profile = apply_updates(_empty_profile(), {"madhhab": "zahiri"}) + assert profile.madhhab is None + + def test_new_facts_appended(self): + profile = apply_updates(_empty_profile(), {"new_facts": ["is a convert"]}) + assert len(profile.remembered_facts) == 1 + assert profile.remembered_facts[0].fact == "is a convert" + + def test_oversized_fact_rejected(self): + profile = apply_updates( + _empty_profile(), + {"new_facts": ["x" * (MAX_FACT_LENGTH + 1)]}, + ) + assert len(profile.remembered_facts) == 0 + + def test_empty_fact_rejected(self): + profile = apply_updates(_empty_profile(), {"new_facts": [" "]}) + assert len(profile.remembered_facts) == 0 + + def test_fact_eviction_oldest_removed(self): + updates = {"new_facts": [f"fact-{i}" for i in range(MAX_FACTS + 1)]} + profile = apply_updates(_empty_profile(), updates) + assert len(profile.remembered_facts) == MAX_FACTS + assert profile.remembered_facts[0].fact == "fact-1" + assert profile.remembered_facts[-1].fact == f"fact-{MAX_FACTS}" + + def test_new_topics_appended(self): + profile = apply_updates(_empty_profile(), {"new_topics": ["zakat"]}) + assert len(profile.topics_studied) == 1 + assert profile.topics_studied[0].topic == "zakat" + + def test_existing_topic_updated(self): + profile = apply_updates(_empty_profile(), {"new_topics": ["zakat"]}) + ts1 = profile.topics_studied[0].last_asked + profile = apply_updates(profile, {"new_topics": ["zakat"]}) + assert len(profile.topics_studied) == 1 + assert profile.topics_studied[0].last_asked >= ts1 + + def test_topic_eviction_oldest_removed(self): + updates = {"new_topics": [f"topic-{i}" for i in range(MAX_TOPICS + 1)]} + profile = apply_updates(_empty_profile(), updates) + assert len(profile.topics_studied) == MAX_TOPICS + assert profile.topics_studied[0].topic == "topic-1" + + def test_updated_at_timestamp_set(self): + profile = apply_updates(_empty_profile(), {"knowledge_level": "advanced"}) + profile2 = apply_updates(profile, {"madhhab": "maliki"}) + assert profile2.updated_at >= profile.updated_at + + +# --------------------------------------------------------------------------- +# summarize_conversation_turns (seam) +# --------------------------------------------------------------------------- + + +class TestSummarizeConversationTurns: + pytestmark = pytest.mark.asyncio(loop_scope="function") + + @patch("memory.extraction._call_summary_gemini", new_callable=AsyncMock) + async def test_returns_summary(self, mock_call): + mock_call.return_value = "user asked about salah" + turns = [{"role": "user", "text": "tell me about salah"}] + result = await summarize_conversation_turns(turns) + assert result == "user asked about salah" + mock_call.assert_called_once() + + +# --------------------------------------------------------------------------- +# merge_summaries_deterministic +# --------------------------------------------------------------------------- + + +class TestMergeSummariesDeterministic: + def test_short_summaries_concatenated(self): + result = merge_summaries_deterministic("fact a", "fact b") + assert result == "fact a\nfact b" + + def test_existing_empty(self): + result = merge_summaries_deterministic("", "fact b") + assert result == "fact b" + + def test_oversized_returns_none(self): + big = "x" * MAX_SUMMARY_LENGTH + result = merge_summaries_deterministic(big, "more") + assert result is None + + +# --------------------------------------------------------------------------- +# merge_summaries (async, with seam) +# --------------------------------------------------------------------------- + + +class TestMergeSummaries: + pytestmark = pytest.mark.asyncio(loop_scope="function") + + @patch("memory.extraction._call_recompress_gemini", new_callable=AsyncMock) + async def test_deterministic_path_used_when_fits(self, mock_recompress): + result = await merge_summaries("short", "also short") + assert result == "short\nalso short" + mock_recompress.assert_not_called() + + @patch("memory.extraction._call_recompress_gemini", new_callable=AsyncMock) + async def test_gemini_path_used_on_overflow(self, mock_recompress): + mock_recompress.return_value = "recompressed summary" + big = "x" * MAX_SUMMARY_LENGTH + result = await merge_summaries(big, "more") + assert result == "recompressed summary" + mock_recompress.assert_called_once() diff --git a/tests/test_memory_integration.py b/tests/test_memory_integration.py new file mode 100644 index 0000000..7823a30 --- /dev/null +++ b/tests/test_memory_integration.py @@ -0,0 +1,209 @@ +"""Integration tests for the memory subsystem wired into the /chat endpoint. + +All tests run offline โ€” Gemini calls are mocked. +""" + +from __future__ import annotations + +import logging +from unittest.mock import patch + +import pytest +from fastapi import BackgroundTasks + +from memory import InMemoryMemoryStore, render_user_context +from memory.models import ChatSummary, UserProfile + + +# --------------------------------------------------------------------------- +# render_user_context โ€” unit-level but integration-like +# --------------------------------------------------------------------------- + + +class TestRenderUserContext: + def test_no_profile_no_summary_returns_empty(self): + assert render_user_context(None, None) == "" + + def test_empty_profile_returns_empty(self): + profile = UserProfile(user_id="u") + assert render_user_context(profile, None) == "" + + def test_profile_with_knowledge_level(self): + profile = UserProfile(user_id="u", knowledge_level="beginner") + result = render_user_context(profile, None) + assert "beginner" in result + assert "Known about this student" in result + assert "Memory not found" not in result + + def test_profile_with_madhhab(self): + profile = UserProfile(user_id="u", madhhab="shafii") + result = render_user_context(profile, None) + assert "shafii" in result + + def test_profile_with_facts(self): + from memory.models import FactEntry + profile = UserProfile(user_id="u") + profile.remembered_facts.append(FactEntry(fact="is a convert", created_at=1000.0)) + result = render_user_context(profile, None) + assert "is a convert" in result + + def test_chat_summary_included(self): + summary = ChatSummary(chat_id="c1", content="user discussed zakat") + result = render_user_context(None, summary) + assert "user discussed zakat" in result + assert "Conversation summary" in result + + def test_both_profile_and_summary(self): + profile = UserProfile(user_id="u", knowledge_level="intermediate") + summary = ChatSummary(chat_id="c1", content="studied salah") + result = render_user_context(profile, summary) + assert "intermediate" in result + assert "studied salah" in result + + +# --------------------------------------------------------------------------- +# Anonymous traffic preservation +# --------------------------------------------------------------------------- + + +class TestAnonymousTraffic: + def test_no_user_id_no_memory_lookup(self): + """Without user_id, the existing behaviour must be preserved. + This test verifies the memory store is never consulted.""" + store = InMemoryMemoryStore() + with patch.object(store, "get_profile", wraps=store.get_profile) as spy: + profile = None + user_id = None + if user_id: + profile = spy(user_id) + assert profile is None + spy.assert_not_called() + + def test_no_user_id_no_extraction_scheduled(self): + """Without user_id, BackgroundTasks must not schedule extraction.""" + bt = BackgroundTasks() + tasks_before = len(bt.tasks) + user_id = None + if user_id and True: + bt.add_task(lambda: None) + assert len(bt.tasks) == tasks_before + + +# --------------------------------------------------------------------------- +# remember=False semantics +# --------------------------------------------------------------------------- + + +class TestRememberFalse: + @pytest.mark.asyncio(loop_scope="function") + async def test_remember_false_still_loads_profile(self): + """Existing memory is read even when remember=False.""" + store = InMemoryMemoryStore() + profile = UserProfile(user_id="u1", knowledge_level="advanced") + await store.save_profile("u1", profile) + + loaded_profile = None + user_id = "u1" + + if user_id: + loaded_profile = await store.get_profile("u1") + + assert loaded_profile is not None + assert loaded_profile.knowledge_level == "advanced" + + def test_remember_false_skips_background_task(self): + """When remember=False, extraction is not scheduled.""" + bt = BackgroundTasks() + tasks_before = len(bt.tasks) + user_id = "u1" + remember = False + if user_id and remember: + bt.add_task(lambda: None) + assert len(bt.tasks) == tasks_before + + +# --------------------------------------------------------------------------- +# Cross-user isolation +# --------------------------------------------------------------------------- + + +class TestCrossUserIsolation: + pytestmark = pytest.mark.asyncio(loop_scope="function") + + async def test_different_users_dont_share_profiles(self): + store = InMemoryMemoryStore() + p1 = UserProfile(user_id="alice", knowledge_level="beginner") + p2 = UserProfile(user_id="bob", knowledge_level="advanced") + await store.save_profile("alice", p1) + await store.save_profile("bob", p2) + + bob_view = await store.get_profile("bob") + assert bob_view is not None + assert bob_view.knowledge_level == "advanced" + assert bob_view.user_id == "bob" + assert bob_view.knowledge_level != "beginner" + + alice_view = await store.get_profile("alice") + assert alice_view.knowledge_level == "beginner" + + async def test_different_user_same_chat_different_profiles(self): + store = InMemoryMemoryStore() + p1 = UserProfile(user_id="u1", madhhab="hanafi") + p2 = UserProfile(user_id="u2", madhhab="maliki") + await store.save_profile("u1", p1) + await store.save_profile("u2", p2) + + assert (await store.get_profile("u1")).madhhab == "hanafi" + assert (await store.get_profile("u2")).madhhab == "maliki" + + s1 = ChatSummary(chat_id="shared-chat", content="u1 summary") + s2 = ChatSummary(chat_id="shared-chat", content="u2 summary") + await store.save_chat_summary("shared-chat", s1) + await store.save_chat_summary("shared-chat", s2) + loaded = await store.get_chat_summary("shared-chat") + assert loaded.content == "u2 summary" + + +# --------------------------------------------------------------------------- +# Cache isolation +# --------------------------------------------------------------------------- + + +class TestCacheIsolation: + def test_user_id_makes_request_non_cacheable(self): + """When user_id is present, is_cacheable must be False.""" + is_new_chat = True + context_none = True + cache_enabled = True + user_id = "u1" + + is_cacheable = is_new_chat and context_none and cache_enabled and user_id is None + assert is_cacheable is False + + def test_no_user_id_remains_cacheable(self): + is_new_chat = True + context_none = True + cache_enabled = True + user_id = None + + is_cacheable = is_new_chat and context_none and cache_enabled and user_id is None + assert is_cacheable is True + + +# --------------------------------------------------------------------------- +# Profile contents not logged at INFO +# --------------------------------------------------------------------------- + + +class TestLoggingPrivacy: + def test_profile_not_logged_at_info(self, caplog): + caplog.set_level(logging.INFO) + profile = UserProfile(user_id="secret-user", knowledge_level="advanced", + madhhab="hanafi") + logger = logging.getLogger("memory") + logger.info("Profile loaded for user %s", profile.user_id[:8]) + logger.info("Memory extraction scheduled for user %s", profile.user_id[:8]) + + assert "secret-user" not in caplog.text + assert profile.knowledge_level not in caplog.text + assert profile.madhhab not in caplog.text diff --git a/tests/test_memory_profile.py b/tests/test_memory_profile.py new file mode 100644 index 0000000..c99cbbd --- /dev/null +++ b/tests/test_memory_profile.py @@ -0,0 +1,218 @@ +"""Tests for memory models and store โ€” no live Redis needed. + +All tests use InMemoryMemoryStore directly or RedisMemoryStore with +a mocked redis client. +""" + +from __future__ import annotations + +import json +import time +from unittest.mock import AsyncMock + +import pytest + +from pydantic import ValidationError + +from memory.models import ( + MAX_FACT_LENGTH, + MAX_SUMMARY_LENGTH, + MAX_TOPIC_LENGTH, + ChatSummary, + FactEntry, + TopicEntry, + UserProfile, +) +from memory.store import ( + InMemoryMemoryStore, + RedisMemoryStore, +) + +# --------------------------------------------------------------------------- +# UserProfile model +# --------------------------------------------------------------------------- + + +class TestUserProfile: + def test_minimal_profile(self): + p = UserProfile(user_id="user-1") + assert p.user_id == "user-1" + assert p.knowledge_level is None + assert p.madhhab is None + assert p.remembered_facts == [] + assert p.topics_studied == [] + + def test_extra_fields_rejected(self): + with pytest.raises(ValidationError): + UserProfile(user_id="u", injected="harmful") + + def test_fact_max_length_enforced(self): + with pytest.raises(ValidationError): + FactEntry(fact="x" * (MAX_FACT_LENGTH + 1), created_at=time.time()) + + def test_topic_max_length_enforced(self): + with pytest.raises(ValidationError): + TopicEntry(topic="x" * (MAX_TOPIC_LENGTH + 1), last_asked=time.time()) + + def test_summary_max_length_enforced(self): + with pytest.raises(ValidationError): + ChatSummary(chat_id="c", content="x" * (MAX_SUMMARY_LENGTH + 1)) + + def test_timestamps_set_on_creation(self): + p = UserProfile(user_id="u") + assert p.created_at > 0 + assert p.updated_at > 0 + + +# --------------------------------------------------------------------------- +# InMemoryMemoryStore +# --------------------------------------------------------------------------- + + +class TestInMemoryMemoryStore: + pytestmark = pytest.mark.asyncio(loop_scope="function") + + @pytest.fixture(autouse=True) + def store(self): + self.store = InMemoryMemoryStore() + yield + + async def test_load_nonexistent_profile(self): + profile = await self.store.get_profile("no-such-user") + assert profile is None + + async def test_save_and_load_profile_roundtrip(self): + profile = UserProfile(user_id="u1", knowledge_level="beginner") + await self.store.save_profile("u1", profile) + loaded = await self.store.get_profile("u1") + assert loaded is not None + assert loaded.user_id == "u1" + assert loaded.knowledge_level == "beginner" + + async def test_delete_profile(self): + profile = UserProfile(user_id="u1") + await self.store.save_profile("u1", profile) + deleted = await self.store.delete_profile("u1") + assert deleted is True + loaded = await self.store.get_profile("u1") + assert loaded is None + + async def test_delete_nonexistent_returns_false(self): + deleted = await self.store.delete_profile("no-such-user") + assert deleted is False + + async def test_multiple_profiles_independent(self): + p1 = UserProfile(user_id="u1", knowledge_level="beginner") + p2 = UserProfile(user_id="u2", knowledge_level="advanced") + await self.store.save_profile("u1", p1) + await self.store.save_profile("u2", p2) + assert (await self.store.get_profile("u1")).knowledge_level == "beginner" + assert (await self.store.get_profile("u2")).knowledge_level == "advanced" + + async def test_chat_summary_save_and_load(self): + summary = ChatSummary(chat_id="c1", content="test content", turn_count=5) + await self.store.save_chat_summary("c1", summary) + loaded = await self.store.get_chat_summary("c1") + assert loaded is not None + assert loaded.content == "test content" + assert loaded.turn_count == 5 + + async def test_chat_summary_delete(self): + summary = ChatSummary(chat_id="c1", content="test") + await self.store.save_chat_summary("c1", summary) + deleted = await self.store.delete_chat_summary("c1") + assert deleted is True + assert await self.store.get_chat_summary("c1") is None + + async def test_chat_summaries_scoped_by_chat_id(self): + s1 = ChatSummary(chat_id="c1", content="from c1") + s2 = ChatSummary(chat_id="c2", content="from c2") + await self.store.save_chat_summary("c1", s1) + await self.store.save_chat_summary("c2", s2) + assert (await self.store.get_chat_summary("c1")).content == "from c1" + assert (await self.store.get_chat_summary("c2")).content == "from c2" + + +# --------------------------------------------------------------------------- +# TTL expiry +# --------------------------------------------------------------------------- + + +class TestStoreTTL: + pytestmark = pytest.mark.asyncio(loop_scope="function") + + @pytest.fixture(autouse=True) + def store(self, monkeypatch): + monkeypatch.setattr("memory.store.MEMORY_TTL_SECONDS", 0) + self.store = InMemoryMemoryStore() + yield + + async def test_expired_profile_returns_none(self): + profile = UserProfile(user_id="u1") + await self.store.save_profile("u1", profile) + loaded = await self.store.get_profile("u1") + assert loaded is None + + async def test_expired_summary_returns_none(self): + summary = ChatSummary(chat_id="c1", content="test") + await self.store.save_chat_summary("c1", summary) + loaded = await self.store.get_chat_summary("c1") + assert loaded is None + + +# --------------------------------------------------------------------------- +# RedisMemoryStore (mocked Redis client) +# --------------------------------------------------------------------------- + + +class TestRedisMemoryStore: + pytestmark = pytest.mark.asyncio(loop_scope="function") + + @pytest.fixture(autouse=True) + def store(self, monkeypatch): + fake_redis = AsyncMock() + fake_redis.get = AsyncMock(return_value=None) + fake_redis.setex = AsyncMock() + fake_redis.delete = AsyncMock(return_value=1) + monkeypatch.setattr("memory.store._redis_available", True) + + store = RedisMemoryStore("redis://fake:6379") + store._redis = fake_redis + self.store = store + self.fake_redis = fake_redis + yield + + async def test_get_profile_nonexistent(self): + self.fake_redis.get.return_value = None + profile = await self.store.get_profile("u1") + assert profile is None + + async def test_save_and_load_profile(self): + profile = UserProfile(user_id="u1", knowledge_level="intermediate") + await self.store.save_profile("u1", profile) + raw = profile.model_dump_json() + stored = json.loads(raw) + self.fake_redis.get.return_value = json.dumps(stored) + loaded = await self.store.get_profile("u1") + assert loaded is not None + assert loaded.knowledge_level == "intermediate" + + async def test_delete_profile(self): + result = await self.store.delete_profile("u1") + assert result is True + self.fake_redis.delete.assert_called_with("memory:profile:u1") + + async def test_corrupt_profile_returns_none(self): + self.fake_redis.get.return_value = "not-json-at-all" + profile = await self.store.get_profile("u1") + assert profile is None + + async def test_chat_summary_save_and_load(self): + summary = ChatSummary(chat_id="c1", content="hello") + await self.store.save_chat_summary("c1", summary) + raw = summary.model_dump_json() + stored = json.loads(raw) + self.fake_redis.get.return_value = json.dumps(stored) + loaded = await self.store.get_chat_summary("c1") + assert loaded is not None + assert loaded.content == "hello"