diff --git a/.gitignore b/.gitignore index 693038a..42763c5 100644 --- a/.gitignore +++ b/.gitignore @@ -66,9 +66,12 @@ backend/uploads/ /backend/services/document/processors/markdown_output /backend/output /backend/core/hcsx-scigpt2-innocentrhino-acm-f87f8026be3d.json -/backend/.deepeval/.deepeval_telemetry.txt -/backend/tests/ -/files/ +/backend/.deepeval/.deepeval_telemetry.txt +/backend/tests/ +!/backend/tests/ +/backend/tests/* +!/backend/tests/test_chat_memory_service.py +/files/ /backend/files/ /backend/services/document/processors/docling/docling_service_sequential.py diff --git a/README.md b/README.md index 56d2df4..6aadb7b 100644 --- a/README.md +++ b/README.md @@ -26,6 +26,7 @@ It supports: - **Improved authentication** through Better Auth with GitHub OAuth and Microsoft Entra support - **In-app evaluation workflows** powered by **DeepEval**, including custom evaluation steps and LLM-as-a-judge patterns - **Batch and interactive workflows** for extraction, review, and evaluation +- **Short-term chatbot memory** for follow-up questions within an independent chat conversation - **Production-oriented deployment paths** for Azure infrastructure and other containerized environments --- @@ -256,6 +257,7 @@ SummarizationTool-dev/ - [Frontend technical design docs](docs/frontend/README.md) — frontend architecture, page-by-page docs, component index, hooks, and TypeScript interfaces - [Glossary](docs/glossary.md) — definitions for all Azure services, tools, and project-specific terms - [Backend README](backend/README.md) — backend setup and processing details +- [Chat memory](docs/chat-memory.md) - chatbot memory behavior, API contract, and operational notes - [Backend technical design docs](docs/backend/README.md) — backend architecture, workflows, diagrams, data models, schemas, and appendices - [Backend class reference](docs/backend/appendices/class-reference.md) — field-level reference for backend ORM models, schemas, dataclasses, service attributes, and provider classes - [Migration guide](docs/superpowers/migration-guide.md) — architecture migration and platform transition notes diff --git a/backend/README.md b/backend/README.md index ee54f58..cf9bc8a 100644 --- a/backend/README.md +++ b/backend/README.md @@ -140,10 +140,18 @@ cd Summarization_tool/backend && uvicorn main:app --reload --port 8000 --host 0. **Resources:** - 🔗 [DeepEval G-Eval Metrics](https://deepeval.com/docs/metrics-llm-evals) - Official DeepEval documentation +### Chat Memory +- Chat requests require `chat_session_id` so each browser chat has independent memory. +- Backend memory is persisted with LangGraph checkpoints in PostgreSQL through the existing `DATABASE_URL`. +- Older turns that fall outside the recent-message prompt window are retained through a rolling conversation summary. +- Uploaded document markdown is sent as request-time context and is not stored as permanent chat memory. +- See [Chat Memory](../docs/chat-memory.md) for the API contract and operational notes. + --- ## 📚 Documentation - **[DeepEval Metrics](https://deepeval.com/docs/metrics-llm-evals)** - Official DeepEval LLM evaluation metrics documentation +- **[Chat Memory](../docs/chat-memory.md)** - Chat memory behavior, API contract, and operational notes - **[Examples](examples/)** - Python examples and usage patterns - **[API Docs](http://localhost:8000/docs)** - Interactive API documentation (when server running) diff --git a/backend/api/chat/router.py b/backend/api/chat/router.py index 2652e34..a6eb87b 100644 --- a/backend/api/chat/router.py +++ b/backend/api/chat/router.py @@ -1,73 +1,144 @@ """Chat API endpoint for support staff chatbot""" -from fastapi import APIRouter, Depends +from typing import Any, Optional + +from fastapi import APIRouter, Depends, HTTPException from fastapi.responses import JSONResponse -from pydantic import BaseModel -from typing import Optional +from pydantic import BaseModel, Field from core.auth import get_current_user -from services.llm.llm_service import LLMService +from services.chat_memory import ChatMemoryRequest, ChatMemoryService +from services.chat_memory.chat_memory_service import GENERIC_MODEL_ERROR_MESSAGE router = APIRouter(prefix="/api/chat", tags=["chat"]) -llm_service = LLMService() +chat_memory_service: Optional[ChatMemoryService] = None class ChatQueryRequest(BaseModel): - query: str + chat_session_id: str = Field(min_length=1) + query: str = Field(min_length=1) document_markdown: Optional[str] = None - model_type: str # "azure", "gemini", "anthropic", "llama", "azure-llama", "macbook" + model_type: str model_id: Optional[str] = None deployment: Optional[str] = None api_version: Optional[str] = None -@router.post("/query", dependencies=[Depends(get_current_user)]) -async def chat_query(request: ChatQueryRequest): - """ - Send a chat message with optional document context. - - When document_markdown is provided, it is injected into the prompt so the - model can answer questions about the uploaded document. - """ - if request.document_markdown: - user_prompt = ( - "The following document has been uploaded by the user:\n\n" - f"\n{request.document_markdown}\n\n\n" - f"User question: {request.query}" - ) - system_message = ( - "You are a helpful document assistant for Health Canada support staff. " - "Answer the user's question based on the provided document. " - "If the answer is not found in the document, say so clearly and offer " - "general guidance if possible." +class ChatHistoryMessage(BaseModel): + id: str + role: str + content: str + + +class ChatHistorySummary(BaseModel): + chat_session_id: str + title: str + message_count: int + latest_message: str + latest_checkpoint_id: Optional[str] = None + + +class ChatHistoryListResponse(BaseModel): + chats: list[ChatHistorySummary] + total: int + + +class ChatHistoryDetailResponse(BaseModel): + chat_session_id: str + messages: list[ChatHistoryMessage] + conversation_summary: str = "" + summarized_message_count: int = 0 + context_usage: Optional[dict[str, Any]] = None + + +def get_chat_memory_service() -> ChatMemoryService: + global chat_memory_service + if chat_memory_service is None: + chat_memory_service = ChatMemoryService() + return chat_memory_service + + +@router.get("/history", response_model=ChatHistoryListResponse) +async def list_chat_history( + current_user: dict = Depends(get_current_user), +): + try: + return await get_chat_memory_service().list_chat_sessions(current_user["id"]) + except Exception as exc: + raise HTTPException( + status_code=500, + detail=f"Error listing chat history: {str(exc)}", + ) from exc + + +@router.get("/history/{chat_session_id}", response_model=ChatHistoryDetailResponse) +async def get_chat_history( + chat_session_id: str, + current_user: dict = Depends(get_current_user), +): + try: + chat = await get_chat_memory_service().get_chat_session( + user_id=current_user["id"], + chat_session_id=chat_session_id, ) - else: - user_prompt = request.query - system_message = ( - "You are a helpful assistant for Health Canada support staff. " - "Answer questions clearly and concisely." + except Exception as exc: + raise HTTPException( + status_code=500, + detail=f"Error loading chat history: {str(exc)}", + ) from exc + + if chat is None: + raise HTTPException(status_code=404, detail="Chat history not found") + return chat + + +@router.delete("/history/{chat_session_id}") +async def delete_chat_history( + chat_session_id: str, + current_user: dict = Depends(get_current_user), +): + try: + deleted = await get_chat_memory_service().delete_chat_session( + user_id=current_user["id"], + chat_session_id=chat_session_id, ) + except Exception as exc: + raise HTTPException( + status_code=500, + detail=f"Error deleting chat history: {str(exc)}", + ) from exc - result = await llm_service.generate_paragraph( - user_prompt=user_prompt, - model_type=request.model_type, - model_id=request.model_id, - deployment=request.deployment, - api_version=request.api_version, - max_tokens=4096, - temperature=0.3, - system_message=system_message, - ) - - if not result.get("success"): + if not deleted: + raise HTTPException(status_code=404, detail="Chat history not found") + return {"message": f"Chat history {chat_session_id} deleted successfully"} + + +@router.post("/query", dependencies=[Depends(get_current_user)]) +async def chat_query( + request: ChatQueryRequest, + current_user: dict = Depends(get_current_user), +): + document_context = request.document_markdown + + try: + result = await get_chat_memory_service().invoke( + ChatMemoryRequest( + user_id=current_user["id"], + chat_session_id=request.chat_session_id, + query=request.query, + model_type=request.model_type, + model_id=request.model_id, + deployment=request.deployment, + api_version=request.api_version, + document_context=document_context, + ) + ) + return result + except Exception: return JSONResponse( status_code=500, content={ "success": False, - "error": result.get( - "error", "The model call failed. Please try again." - ), + "error": GENERIC_MODEL_ERROR_MESSAGE, }, ) - - return {"success": True, "response": result.get("content", "")} diff --git a/backend/requirements.txt b/backend/requirements.txt index ac2de57..930e095 100644 --- a/backend/requirements.txt +++ b/backend/requirements.txt @@ -18,6 +18,8 @@ deepeval>=3.6.9 google-cloud-aiplatform>=1.122.0 langchain-openai>=1.0.1 langchain-google-vertexai>=3.0.1 +langgraph>=1.2.4 +langgraph-checkpoint-postgres>=3.1.0 pytest>=8.4.2 pytest-asyncio>=1.2.0 anthropic[vertex]>=0.39.0 diff --git a/backend/services/chat_memory/__init__.py b/backend/services/chat_memory/__init__.py new file mode 100644 index 0000000..202685b --- /dev/null +++ b/backend/services/chat_memory/__init__.py @@ -0,0 +1,3 @@ +from .chat_memory_service import ChatMemoryRequest, ChatMemoryService + +__all__ = ["ChatMemoryRequest", "ChatMemoryService"] diff --git a/backend/services/chat_memory/chat_memory_service.py b/backend/services/chat_memory/chat_memory_service.py new file mode 100644 index 0000000..37940ce --- /dev/null +++ b/backend/services/chat_memory/chat_memory_service.py @@ -0,0 +1,768 @@ +from __future__ import annotations + +import asyncio +import atexit +import inspect +import weakref +from dataclasses import dataclass +from typing import Annotated, Any, Dict, List, Optional, TypedDict + +from services.llm.llm_service import LLMService + +LANGGRAPH_INSTALL_MESSAGE = ( + "LangGraph chat memory dependencies are unavailable. " + "Install backend requirements and retry." +) +GENERIC_MODEL_ERROR_MESSAGE = "The model call failed. Please try again." +MAX_HISTORY_MESSAGES_IN_PROMPT = 12 +SUMMARY_MAX_TOKENS = 1_024 +DEFAULT_MAX_CONTEXT_TOKENS = 128_000 +MAX_RESPONSE_TOKENS = 4_096 +CONTEXT_USAGE_HARD_LIMIT_PERCENTAGE = 95.0 +CONTEXT_WINDOW_ERROR_MESSAGE = ( + "Context window is too large. Remove a document, start a new chat, " + "or use a model with a larger context window." +) + + +class ModelProviderError(Exception): + pass + + +class ContextWindowExceededError(Exception): + def __init__(self, context_usage: Dict[str, Any]): + super().__init__(CONTEXT_WINDOW_ERROR_MESSAGE) + self.context_usage = context_usage + + +try: + from langchain_core.messages import AIMessage, BaseMessage, HumanMessage + from langchain_core.runnables import RunnableConfig + + LANGCHAIN_MESSAGES_IMPORT_ERROR: Optional[ImportError] = None +except ImportError as exc: + AIMessage = BaseMessage = HumanMessage = RunnableConfig = Any # type: ignore[assignment] + LANGCHAIN_MESSAGES_IMPORT_ERROR = exc + +try: + from langgraph.graph import END, START, StateGraph + from langgraph.graph.message import add_messages + + LANGGRAPH_GRAPH_IMPORT_ERROR: Optional[ImportError] = None +except ImportError as exc: + END = "__end__" + START = "__start__" + StateGraph = None + LANGGRAPH_GRAPH_IMPORT_ERROR = exc + + def add_messages(existing: List[Any], new: List[Any]) -> List[Any]: + return [*(existing or []), *(new or [])] + + +try: + from langgraph.checkpoint.memory import InMemorySaver as LangGraphInMemorySaver + + LANGGRAPH_MEMORY_IMPORT_ERROR: Optional[ImportError] = None +except ImportError: + try: + from langgraph.checkpoint.memory import MemorySaver as LangGraphInMemorySaver + + LANGGRAPH_MEMORY_IMPORT_ERROR = None + except ImportError as exc: + LangGraphInMemorySaver = None + LANGGRAPH_MEMORY_IMPORT_ERROR = exc + + +@dataclass +class ChatMemoryRequest: + user_id: str + chat_session_id: str + query: str + model_type: str + model_id: Optional[str] = None + deployment: Optional[str] = None + api_version: Optional[str] = None + document_context: Optional[str] = None + + +class ChatState(TypedDict, total=False): + messages: Annotated[List[BaseMessage], add_messages] + conversation_summary: str + summarized_message_count: int + context_usage: Dict[str, Any] + + +class ChatMemoryService: + """Persist per-user chat turns with LangGraph and build bounded prompts. + + `chat_session_id` is generated by the client, so the persisted LangGraph + thread also includes `user_id`. Document content is passed as request-time + context only and is deliberately excluded from stored chat messages. Older + turns are folded into `conversation_summary` so long chats retain useful + context after raw messages fall outside the recent-message prompt window. + """ + + def __init__( + self, + llm_service: Optional[LLMService] = None, + checkpointer: Optional[Any] = None, + use_memory_checkpointer: bool = False, + ): + self.llm_service = llm_service or LLMService() + self._use_memory_checkpointer = use_memory_checkpointer + self._async_checkpointer_context = None + self._did_register_atexit_close = False + self._initialization_lock = asyncio.Lock() + self._session_locks: weakref.WeakValueDictionary[str, asyncio.Lock] = ( + weakref.WeakValueDictionary() + ) + self._ensure_langgraph_dependencies(require_in_memory=use_memory_checkpointer) + self.checkpointer = checkpointer + self.graph = ( + self._build_graph(checkpointer) if checkpointer is not None else None + ) + + def _ensure_langgraph_dependencies(self, require_in_memory: bool = False) -> None: + import_errors: List[str] = [] + + if LANGCHAIN_MESSAGES_IMPORT_ERROR is not None: + import_errors.append( + f"langchain_core.messages: {LANGCHAIN_MESSAGES_IMPORT_ERROR}" + ) + if LANGGRAPH_GRAPH_IMPORT_ERROR is not None: + import_errors.append(f"langgraph.graph: {LANGGRAPH_GRAPH_IMPORT_ERROR}") + if require_in_memory and LANGGRAPH_MEMORY_IMPORT_ERROR is not None: + import_errors.append( + f"langgraph.checkpoint.memory: {LANGGRAPH_MEMORY_IMPORT_ERROR}" + ) + + if import_errors: + details = "; ".join(import_errors) + raise RuntimeError(f"{LANGGRAPH_INSTALL_MESSAGE} Import errors: {details}") + + def _get_async_postgres_saver_class(self) -> Any: + from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver + + return AsyncPostgresSaver + + async def _build_checkpointer(self) -> Any: + if self._use_memory_checkpointer: + return self._build_in_memory_checkpointer() + + try: + AsyncPostgresSaver = self._get_async_postgres_saver_class() + except Exception as exc: + raise RuntimeError( + f"Failed to import LangGraph Postgres checkpointer: {exc}" + ) from exc + + try: + from models.base import DATABASE_URL + + saver_context = AsyncPostgresSaver.from_conn_string(DATABASE_URL) + self._async_checkpointer_context = saver_context + saver = await saver_context.__aenter__() + self._register_checkpointer_cleanup() + if hasattr(saver, "setup"): + setup_result = saver.setup() + if inspect.isawaitable(setup_result): + await setup_result + return saver + except Exception as exc: + await self._close_async_checkpointer_context() + raise RuntimeError( + f"Failed to initialize LangGraph Postgres checkpointer: {exc}" + ) from exc + + def _register_checkpointer_cleanup(self) -> None: + if self._did_register_atexit_close: + return + atexit.register(self._close_checkpointer_context_at_exit) + self._did_register_atexit_close = True + + def _close_checkpointer_context_at_exit(self) -> None: + if self._async_checkpointer_context is None: + return + try: + asyncio.run(self._close_async_checkpointer_context()) + except Exception: + pass + + async def _close_async_checkpointer_context(self) -> None: + if self._async_checkpointer_context is None: + return + try: + await self._async_checkpointer_context.__aexit__(None, None, None) + except Exception: + pass + finally: + self._async_checkpointer_context = None + + def _build_in_memory_checkpointer(self) -> Any: + if LangGraphInMemorySaver is None: + detail = ( + str(LANGGRAPH_MEMORY_IMPORT_ERROR) + if LANGGRAPH_MEMORY_IMPORT_ERROR is not None + else "LangGraph in-memory saver is unavailable." + ) + raise RuntimeError( + f"{LANGGRAPH_INSTALL_MESSAGE} Unable to import LangGraph in-memory saver: {detail}" + ) from LANGGRAPH_MEMORY_IMPORT_ERROR + return LangGraphInMemorySaver() + + def _build_graph(self, checkpointer: Any) -> Any: + self._ensure_langgraph_dependencies() + + graph = StateGraph(ChatState) + graph.add_node("generate", self._generate_response) + graph.add_edge(START, "generate") + graph.add_edge("generate", END) + return graph.compile(checkpointer=checkpointer) + + async def _ensure_graph(self) -> Any: + if self.graph is not None: + return self.graph + + async with self._initialization_lock: + if self.graph is not None: + return self.graph + + checkpointer = self.checkpointer or await self._build_checkpointer() + self.checkpointer = checkpointer + self.graph = self._build_graph(checkpointer) + return self.graph + + def _get_session_lock(self, chat_session_id: str) -> asyncio.Lock: + lock = self._session_locks.get(chat_session_id) + if lock is None: + lock = asyncio.Lock() + self._session_locks[chat_session_id] = lock + return lock + + def _build_thread_id(self, user_id: str, chat_session_id: str) -> str: + return f"user:{user_id}:chat:{chat_session_id}" + + def _thread_prefix(self, user_id: str) -> str: + return f"user:{user_id}:chat:" + + def _chat_session_id_from_thread_id(self, user_id: str, thread_id: str) -> str: + prefix = self._thread_prefix(user_id) + if thread_id.startswith(prefix): + return thread_id[len(prefix) :] + return thread_id + + async def invoke(self, request: ChatMemoryRequest) -> Dict[str, Any]: + thread_id = self._build_thread_id(request.user_id, request.chat_session_id) + config = { + "configurable": { + "thread_id": thread_id, + "model_type": request.model_type, + "model_id": request.model_id, + "deployment": request.deployment, + "api_version": request.api_version, + "document_context": request.document_context, + "query": request.query, + } + } + lock = self._get_session_lock(thread_id) + graph = await self._ensure_graph() + + try: + async with lock: + result = await graph.ainvoke( + {}, + config=config, + ) + response = self._extract_ai_response(result.get("messages", [])) + return { + "success": True, + "response": response, + "chat_session_id": request.chat_session_id, + "context_usage": result.get("context_usage"), + } + except ContextWindowExceededError as exc: + return { + "success": False, + "response": "", + "chat_session_id": request.chat_session_id, + "error": CONTEXT_WINDOW_ERROR_MESSAGE, + "error_code": "context_window_exceeded", + "context_usage": exc.context_usage, + } + except ModelProviderError: + return { + "success": False, + "response": "", + "chat_session_id": request.chat_session_id, + "error": GENERIC_MODEL_ERROR_MESSAGE, + } + + async def list_chat_sessions(self, user_id: str, limit: int = 50) -> Dict[str, Any]: + await self._ensure_graph() + thread_rows = await self._list_thread_rows_for_user(user_id, limit=limit) + chats = [] + for row in thread_rows: + chat_session_id = self._chat_session_id_from_thread_id( + user_id, row["thread_id"] + ) + detail = await self.get_chat_session( + user_id=user_id, + chat_session_id=chat_session_id, + ) + if detail is None: + continue + messages = detail["messages"] + chats.append( + { + "chat_session_id": chat_session_id, + "title": self._build_chat_title(messages), + "message_count": len(messages), + "latest_message": self._build_chat_history_metadata(messages), + "latest_checkpoint_id": row.get("latest_checkpoint_id"), + } + ) + return {"chats": chats, "total": len(chats)} + + async def get_chat_session( + self, user_id: str, chat_session_id: str + ) -> Optional[Dict[str, Any]]: + graph = await self._ensure_graph() + thread_id = self._build_thread_id(user_id, chat_session_id) + state = await graph.aget_state({"configurable": {"thread_id": thread_id}}) + values = state.values or {} + messages = list(values.get("messages", [])) + if not messages and not values.get("conversation_summary"): + return None + return { + "chat_session_id": chat_session_id, + "messages": [ + self._message_to_response(message, index) + for index, message in enumerate(messages) + ], + "conversation_summary": str(values.get("conversation_summary") or ""), + "summarized_message_count": self._coerce_summarized_message_count( + values.get("summarized_message_count"), + total_messages=len(messages), + ), + "context_usage": values.get("context_usage"), + } + + async def delete_chat_session(self, user_id: str, chat_session_id: str) -> bool: + graph = await self._ensure_graph() + thread_id = self._build_thread_id(user_id, chat_session_id) + state = await graph.aget_state({"configurable": {"thread_id": thread_id}}) + values = state.values or {} + if not values.get("messages") and not values.get("conversation_summary"): + return False + + delete_thread = getattr(self.checkpointer, "adelete_thread", None) + if delete_thread is None: + raise RuntimeError( + "Chat history deletion is unavailable for this checkpointer." + ) + await delete_thread(thread_id) + self._session_locks.pop(thread_id, None) + return True + + async def _list_thread_rows_for_user( + self, user_id: str, limit: int + ) -> List[Dict[str, Any]]: + checkpointer = self.checkpointer + cursor_factory = getattr(checkpointer, "_cursor", None) + if cursor_factory is None: + return [] + + async with cursor_factory() as cur: + await cur.execute( + """ + SELECT thread_id, MAX(checkpoint_id) AS latest_checkpoint_id + FROM checkpoints + WHERE thread_id LIKE %s + GROUP BY thread_id + ORDER BY latest_checkpoint_id DESC + LIMIT %s + """, + (f"{self._thread_prefix(user_id)}%", int(limit)), + ) + rows = await cur.fetchall() + + return [ + { + "thread_id": row["thread_id"], + "latest_checkpoint_id": row["latest_checkpoint_id"], + } + for row in rows + ] + + def _message_to_response(self, message: BaseMessage, index: int) -> Dict[str, str]: + role = "assistant" if isinstance(message, AIMessage) else "user" + return { + "id": f"stored-{index}", + "role": role, + "content": str(message.content), + } + + def _build_chat_title(self, messages: List[Dict[str, str]]) -> str: + for message in messages: + if message["role"] == "user" and message["content"].strip(): + return self._truncate_preview(message["content"], max_length=64) + return "Untitled chat" + + def _build_chat_history_metadata(self, messages: List[Dict[str, str]]) -> str: + user_message_count = sum(1 for message in messages if message["role"] == "user") + if user_message_count == 1: + return "1 question" + return f"{user_message_count} questions" + + def _truncate_preview(self, value: str, max_length: int) -> str: + normalized = " ".join(value.split()) + if len(normalized) <= max_length: + return normalized + return normalized[: max_length - 1].rstrip() + "..." + + async def _generate_response( + self, + state: ChatState, + config: Optional[RunnableConfig] = None, + ) -> Dict[str, Any]: + configurable = (config or {}).get("configurable", {}) + prior_messages = list(state.get("messages", [])) + query = str(configurable.get("query") or "").strip() + if not query: + raise ModelProviderError("Missing chat query.") + current_user_message = HumanMessage(content=query) + messages = [*prior_messages, current_user_message] + conversation_summary = str(state.get("conversation_summary", "") or "") + summarized_message_count = self._coerce_summarized_message_count( + state.get("summarized_message_count"), + total_messages=len(messages), + ) + conversation_summary, summarized_message_count = ( + await self._ensure_summary_for_messages( + messages=messages, + conversation_summary=conversation_summary, + summarized_message_count=summarized_message_count, + target_message_count=self._messages_omitted_from_prompt(messages), + config=config, + ) + ) + prompt = self._build_user_prompt( + messages=messages, + conversation_summary=conversation_summary, + document_context=configurable.get("document_context"), + ) + system_message = self._build_system_message( + has_document_context=bool(configurable.get("document_context")) + ) + context_usage = self._build_context_usage( + user_prompt=prompt, + system_message=system_message, + messages=messages, + conversation_summary=conversation_summary, + document_context=configurable.get("document_context"), + model_type=str(configurable.get("model_type", "")), + model_id=configurable.get("model_id"), + deployment=configurable.get("deployment"), + ) + if self._is_context_window_too_large(context_usage): + raise ContextWindowExceededError(context_usage) + + try: + result = await self.llm_service.generate_paragraph( + user_prompt=prompt, + model_type=configurable["model_type"], + model_id=configurable.get("model_id"), + deployment=configurable.get("deployment"), + api_version=configurable.get("api_version"), + max_tokens=MAX_RESPONSE_TOKENS, + temperature=0.3, + system_message=system_message, + ) + except Exception as exc: + raise ModelProviderError(GENERIC_MODEL_ERROR_MESSAGE) from exc + if not result.get("success"): + raise ModelProviderError( + str(result.get("error", GENERIC_MODEL_ERROR_MESSAGE)) + ) + + ai_message = AIMessage(content=str(result.get("content", ""))) + updated_messages = [*messages, ai_message] + conversation_summary, summarized_message_count = ( + await self._ensure_summary_for_messages( + messages=updated_messages, + conversation_summary=conversation_summary, + summarized_message_count=summarized_message_count, + target_message_count=max( + 0, len(updated_messages) - MAX_HISTORY_MESSAGES_IN_PROMPT + ), + config=config, + ) + ) + + return { + "messages": [current_user_message, ai_message], + "conversation_summary": conversation_summary, + "summarized_message_count": summarized_message_count, + "context_usage": context_usage, + } + + def _build_user_prompt( + self, + messages: List[BaseMessage], + conversation_summary: str, + document_context: Optional[str], + ) -> str: + history_lines: List[str] = [] + current_query = "" + + included_messages = self._select_messages_for_prompt(messages) + + for index, message in enumerate(included_messages): + is_last = index == len(included_messages) - 1 + if isinstance(message, HumanMessage): + if is_last: + current_query = str(message.content) + else: + history_lines.append(f"User: {message.content}") + elif isinstance(message, AIMessage): + history_lines.append(f"Assistant: {message.content}") + + sections: List[str] = [] + omitted_count = max(0, len(messages) - len(included_messages)) + if conversation_summary.strip(): + sections.append( + "Summary of earlier conversation:\n" f"{conversation_summary.strip()}" + ) + elif omitted_count: + sections.append( + f"Earlier conversation omitted from this request: {omitted_count} message(s)." + ) + if history_lines: + sections.append("Conversation so far:\n" + "\n".join(history_lines)) + if document_context: + sections.append( + "Current document context for this request:\n" + f"{document_context}\n\n" + "Use this context when it is relevant. Do not assume it remains available in future turns unless it is provided again." + ) + sections.append(f"User question: {current_query}") + return "\n\n".join(sections) + + def _select_messages_for_prompt( + self, messages: List[BaseMessage] + ) -> List[BaseMessage]: + if len(messages) <= MAX_HISTORY_MESSAGES_IN_PROMPT + 1: + return messages + + current_message = messages[-1:] + recent_history = messages[-(MAX_HISTORY_MESSAGES_IN_PROMPT + 1) : -1] + return [*recent_history, *current_message] + + def _messages_omitted_from_prompt(self, messages: List[BaseMessage]) -> int: + included_messages = self._select_messages_for_prompt(messages) + return max(0, len(messages) - len(included_messages)) + + def _coerce_summarized_message_count(self, value: Any, total_messages: int) -> int: + try: + count = int(value or 0) + except (TypeError, ValueError): + return 0 + return min(max(0, count), total_messages) + + async def _ensure_summary_for_messages( + self, + messages: List[BaseMessage], + conversation_summary: str, + summarized_message_count: int, + target_message_count: int, + config: Optional[RunnableConfig], + ) -> tuple[str, int]: + target_message_count = min(max(0, target_message_count), len(messages)) + summarized_message_count = self._coerce_summarized_message_count( + summarized_message_count, + total_messages=len(messages), + ) + if target_message_count <= summarized_message_count: + return conversation_summary, summarized_message_count + + new_messages = messages[summarized_message_count:target_message_count] + if not new_messages: + return conversation_summary, summarized_message_count + + next_summary = await self._summarize_messages( + existing_summary=conversation_summary, + messages_to_summarize=new_messages, + config=config, + ) + if next_summary is None: + return conversation_summary, summarized_message_count + return next_summary, target_message_count + + async def _summarize_messages( + self, + existing_summary: str, + messages_to_summarize: List[BaseMessage], + config: Optional[RunnableConfig], + ) -> Optional[str]: + configurable = (config or {}).get("configurable", {}) + transcript = self._format_messages_for_summary(messages_to_summarize) + prompt = ( + "Update the rolling summary of this chat conversation.\n\n" + "Existing summary:\n" + f"{existing_summary.strip() or '(none)'}\n\n" + "New messages to fold into the summary:\n" + f"{transcript}\n\n" + "Return only the updated summary. Keep it concise but preserve durable " + "facts that may matter later, including user goals, document names or " + "topics, decisions, constraints, definitions, and unresolved questions. " + "Do not include full document excerpts." + ) + system_message = ( + "You summarize chat history for future context. Use only the provided " + "conversation. Do not add external facts." + ) + + try: + result = await self.llm_service.generate_paragraph( + user_prompt=prompt, + model_type=configurable["model_type"], + model_id=configurable.get("model_id"), + deployment=configurable.get("deployment"), + api_version=configurable.get("api_version"), + max_tokens=SUMMARY_MAX_TOKENS, + temperature=0.1, + system_message=system_message, + ) + except Exception: + return None + + if not result.get("success"): + return None + + summary = str(result.get("content", "")).strip() + return summary or existing_summary + + def _format_messages_for_summary(self, messages: List[BaseMessage]) -> str: + lines: List[str] = [] + for message in messages: + if isinstance(message, HumanMessage): + role = "User" + elif isinstance(message, AIMessage): + role = "Assistant" + else: + role = message.__class__.__name__ + lines.append(f"{role}: {message.content}") + return "\n".join(lines) + + def _build_context_usage( + self, + user_prompt: str, + system_message: str, + messages: List[BaseMessage], + conversation_summary: str, + document_context: Optional[str], + model_type: str, + model_id: Optional[str], + deployment: Optional[str], + ) -> Dict[str, Any]: + included_messages = self._select_messages_for_prompt(messages) + prompt_tokens = self._estimate_tokens(user_prompt) + system_tokens = self._estimate_tokens(system_message) + summary_tokens = self._estimate_tokens(conversation_summary) + document_tokens = self._estimate_tokens(document_context or "") + max_context_tokens = self._estimate_max_context_tokens( + model_type=model_type, + model_id=model_id, + deployment=deployment, + ) + estimated_tokens = prompt_tokens + system_tokens + percentage = ( + round((estimated_tokens / max_context_tokens) * 100, 3) + if max_context_tokens + else None + ) + + return { + "estimated_tokens": estimated_tokens, + "max_context_tokens": max_context_tokens, + "percentage": percentage, + "history_message_count": len(messages), + "included_history_message_count": len(included_messages), + "omitted_history_message_count": max( + 0, len(messages) - len(included_messages) + ), + "summary_tokens": summary_tokens, + "document_context_tokens": document_tokens, + "reserved_response_tokens": MAX_RESPONSE_TOKENS, + "hard_limit_percentage": CONTEXT_USAGE_HARD_LIMIT_PERCENTAGE, + "method": "estimated_chars_div_4", + } + + def _is_context_window_too_large(self, context_usage: Dict[str, Any]) -> bool: + max_context_tokens = context_usage.get("max_context_tokens") + estimated_tokens = context_usage.get("estimated_tokens") + percentage = context_usage.get("percentage") + + if not max_context_tokens or not estimated_tokens: + return False + if estimated_tokens + MAX_RESPONSE_TOKENS >= max_context_tokens: + return True + return ( + isinstance(percentage, (int, float)) + and percentage >= CONTEXT_USAGE_HARD_LIMIT_PERCENTAGE + ) + + def _estimate_tokens(self, text: str) -> int: + if not text: + return 0 + return max(1, (len(text) + 3) // 4) + + def _estimate_max_context_tokens( + self, + model_type: str, + model_id: Optional[str], + deployment: Optional[str], + ) -> int: + model_key = f"{model_id or ''} {deployment or ''}".lower() + compact_model_key = ( + model_key.replace("-", "") + .replace("_", "") + .replace(".", "") + .replace(" ", "") + ) + + if model_type == "gemini": + return 1_000_000 + if model_type == "anthropic": + return 200_000 + if model_type in ("llama", "azure-llama"): + return 128_000 + if model_type in ("macbook", "vllm"): + return 32_000 + if "gpt54mini" in compact_model_key or "gpt54nano" in compact_model_key: + return 400_000 + if "gpt54" in compact_model_key: + return 1_050_000 + if "gpt5" in compact_model_key: + return 400_000 + if "gpt-4o" in model_key or "gpt-5" in model_key or "o3" in model_key: + return 128_000 + return DEFAULT_MAX_CONTEXT_TOKENS + + def _build_system_message(self, has_document_context: bool) -> str: + if has_document_context: + return ( + "You are a helpful document assistant for Health Canada support staff. " + "Answer the user's question using the current document context and remembered conversation. " + "If the answer is not found in the document context, say so clearly and offer general guidance if possible." + ) + return ( + "You are a helpful assistant for Health Canada support staff. " + "Answer questions clearly and concisely using remembered conversation when relevant." + ) + + def _extract_ai_response(self, messages: List[BaseMessage]) -> str: + for message in reversed(messages): + if isinstance(message, AIMessage): + return str(message.content) + return "" diff --git a/backend/services/llm/anthropic.py b/backend/services/llm/anthropic.py index 79f34c0..e872d0b 100644 --- a/backend/services/llm/anthropic.py +++ b/backend/services/llm/anthropic.py @@ -394,6 +394,7 @@ async def generate_paragraph_with_anthropic( model_id: Optional[str] = None, max_tokens: int = 2048, temperature: float = 0.0, + system_message: Optional[str] = None, project_id_override: Optional[str] = None, location_override: Optional[str] = None, service_account_path_override: Optional[Path] = None, @@ -402,7 +403,9 @@ async def generate_paragraph_with_anthropic( used_model_id = model_id or "claude-sonnet-4-5@20250929" # System message for paragraph generation - system = "You are a scientific writing assistant. Your task is to synthesize extracted information into a cohesive, well-structured paragraph while maintaining complete accuracy. Follow the instructions exactly and preserve all factual details from the provided entities." + system = system_message or ( + "You are a scientific writing assistant. Your task is to synthesize extracted information into a cohesive, well-structured paragraph while maintaining complete accuracy. Follow the instructions exactly and preserve all factual details from the provided entities." + ) messages = [{"role": "user", "content": user_prompt}] diff --git a/backend/services/llm/llm_service.py b/backend/services/llm/llm_service.py index aa62465..64bdb4e 100644 --- a/backend/services/llm/llm_service.py +++ b/backend/services/llm/llm_service.py @@ -419,6 +419,7 @@ async def generate_paragraph( model_id, max_tokens, temperature, + system_message=system_message, ) self._record_session_metrics(session_id, "gcp", result) return result diff --git a/backend/tests/test_chat_memory_service.py b/backend/tests/test_chat_memory_service.py new file mode 100644 index 0000000..cf26a91 --- /dev/null +++ b/backend/tests/test_chat_memory_service.py @@ -0,0 +1,188 @@ +import sys +from pathlib import Path + +import pytest + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from services.chat_memory.chat_memory_service import ( # noqa: E402 + MAX_HISTORY_MESSAGES_IN_PROMPT, + AIMessage, + ChatMemoryRequest, + ChatMemoryService, + HumanMessage, +) + + +class FakeLLMService: + def __init__(self): + self.prompts = [] + self.summary_prompts = [] + + async def generate_paragraph(self, **kwargs): + user_prompt = kwargs["user_prompt"] + if user_prompt.startswith("Update the rolling summary"): + self.summary_prompts.append(kwargs) + return { + "success": True, + "content": "The user previously established that the project codename is Maple.", + } + + self.prompts.append(kwargs) + return {"success": True, "content": "Assistant response."} + + +class FakeGraph: + def __init__(self): + self.ainvoke_calls = [] + + async def ainvoke(self, graph_input, config): + self.ainvoke_calls.append((graph_input, config)) + query = config["configurable"]["query"] + return { + "messages": [ + HumanMessage(content=query), + AIMessage(content="Assistant response."), + ], + "context_usage": {"estimated_tokens": 12}, + } + + +@pytest.mark.asyncio +async def test_invoke_runs_chat_turn_through_langgraph(): + service = ChatMemoryService( + llm_service=FakeLLMService(), + use_memory_checkpointer=True, + ) + fake_graph = FakeGraph() + service.graph = fake_graph + + result = await service.invoke( + ChatMemoryRequest( + user_id="user-1", + chat_session_id="chat-1", + query="What is my name?", + model_type="azure", + ) + ) + + assert result["success"] is True + assert result["response"] == "Assistant response." + assert len(fake_graph.ainvoke_calls) == 1 + + graph_input, config = fake_graph.ainvoke_calls[0] + assert graph_input == {} + assert config["configurable"]["thread_id"] == "user:user-1:chat:chat-1" + assert config["configurable"]["model_type"] == "azure" + assert config["configurable"]["query"] == "What is my name?" + + +@pytest.mark.asyncio +async def test_old_messages_are_summarized_and_included_in_prompt(): + fake_llm = FakeLLMService() + service = ChatMemoryService( + llm_service=fake_llm, + use_memory_checkpointer=True, + ) + + total_turns = (MAX_HISTORY_MESSAGES_IN_PROMPT // 2) + 3 + for index in range(total_turns): + await service.invoke( + ChatMemoryRequest( + user_id="user-1", + chat_session_id="chat-1", + query=f"Question {index}", + model_type="azure", + ) + ) + + assert fake_llm.summary_prompts + latest_prompt = fake_llm.prompts[-1]["user_prompt"] + assert "Summary of earlier conversation:" in latest_prompt + assert "project codename is Maple" in latest_prompt + assert "Question 0" not in latest_prompt + + +@pytest.mark.asyncio +async def test_chat_history_detail_includes_context_usage_and_question_metadata(): + fake_llm = FakeLLMService() + service = ChatMemoryService( + llm_service=fake_llm, + use_memory_checkpointer=True, + ) + + await service.invoke( + ChatMemoryRequest( + user_id="user-1", + chat_session_id="chat-1", + query="What is PMRA?", + model_type="azure", + ) + ) + + detail = await service.get_chat_session( + user_id="user-1", + chat_session_id="chat-1", + ) + assert detail is not None + assert detail["context_usage"]["estimated_tokens"] > 0 + + assert service._build_chat_title(detail["messages"]) == "What is PMRA?" + assert service._build_chat_history_metadata(detail["messages"]) == "1 question" + assert "Assistant response" not in service._build_chat_history_metadata( + detail["messages"] + ) + + +@pytest.mark.asyncio +async def test_summary_failure_does_not_mark_messages_as_summarized(): + class FailingSummaryLLM(FakeLLMService): + async def generate_paragraph(self, **kwargs): + user_prompt = kwargs["user_prompt"] + if user_prompt.startswith("Update the rolling summary"): + self.summary_prompts.append(kwargs) + return {"success": False, "error": "summary unavailable"} + return await super().generate_paragraph(**kwargs) + + fake_llm = FailingSummaryLLM() + service = ChatMemoryService( + llm_service=fake_llm, + use_memory_checkpointer=True, + ) + + for index in range((MAX_HISTORY_MESSAGES_IN_PROMPT // 2) + 2): + await service.invoke( + ChatMemoryRequest( + user_id="user-1", + chat_session_id="chat-1", + query=f"Question {index}", + model_type="azure", + ) + ) + + graph = await service._ensure_graph() + state = await graph.aget_state( + {"configurable": {"thread_id": service._build_thread_id("user-1", "chat-1")}} + ) + + assert fake_llm.summary_prompts + assert state.values.get("conversation_summary") in (None, "") + assert state.values.get("summarized_message_count") in (None, 0) + + +def test_context_usage_reports_small_nonzero_percentages(): + service = ChatMemoryService(use_memory_checkpointer=True) + + context_usage = service._build_context_usage( + user_prompt="What is PMRA?", + system_message="You are a helpful assistant.", + messages=[], + conversation_summary="", + document_context=None, + model_type="gemini", + model_id=None, + deployment=None, + ) + + assert context_usage["estimated_tokens"] > 0 + assert 0 < context_usage["percentage"] < 0.1 diff --git a/frontend/components/ChatPage.tsx b/frontend/components/ChatPage.tsx index 1245211..d5f95df 100644 --- a/frontend/components/ChatPage.tsx +++ b/frontend/components/ChatPage.tsx @@ -1,4 +1,4 @@ -import { useState, useRef, useEffect, useCallback } from "react"; +import { useState, useRef, useEffect, useCallback, ReactNode } from "react"; import ReactMarkdown from "react-markdown"; import remarkGfm from "remark-gfm"; import { @@ -8,6 +8,8 @@ import { FileText, Loader2, AlertCircle, + History, + Trash2, ThumbsUp, ThumbsDown, Copy, @@ -18,6 +20,23 @@ import { } from "lucide-react"; import { Button } from "./ui/button"; import { Textarea } from "./ui/textarea"; +import { + Sheet, + SheetContent, + SheetDescription, + SheetHeader, + SheetTitle, +} from "./ui/sheet"; +import { + AlertDialog, + AlertDialogAction, + AlertDialogCancel, + AlertDialogContent, + AlertDialogDescription, + AlertDialogFooter, + AlertDialogHeader, + AlertDialogTitle, +} from "./ui/alert-dialog"; import { Select, SelectContent, @@ -31,6 +50,12 @@ import { ModelConfig, } from "../utils/modelSelection"; import { getValidToken } from "../utils/authUtils"; +import { + createChatSessionId, + getOrCreateChatSessionId, + resetChatSessionId, + setChatSessionId as persistChatSessionId, +} from "../utils/chatSession"; import { toast } from "./ui/sonner"; // ─── Types ──────────────────────────────────────────────────────────────────── @@ -41,7 +66,24 @@ interface Message { content: string; } +interface ChatHistorySummary { + chat_session_id: string; + title: string; + message_count: number; + latest_message: string; +} + +interface ChatHistoryDetail { + chat_session_id: string; + messages: Message[]; + conversation_summary?: string; + summarized_message_count?: number; + context_usage?: ContextUsage | null; +} + const MAX_DOCS = 5; +const REGENERATE_DISABLED_WITH_MEMORY_TITLE = + "Regenerate is unavailable for chat memory sessions until turn replacement is supported."; type DocEntry = | { status: "loading"; file: File; tempId: string } @@ -61,6 +103,27 @@ interface ChatPageProps { onSignOut?: () => void; } +interface MarkdownCodeProps { + inline?: boolean; + children?: ReactNode; +} + +interface ChatModelConfig extends ModelConfig { + model_type?: string; +} + +interface ContextUsage { + estimated_tokens: number; + max_context_tokens: number; + percentage: number | null; + omitted_history_message_count: number; + summary_tokens?: number; + document_context_tokens: number; + reserved_response_tokens?: number; + hard_limit_percentage?: number; + method: string; +} + // ─── Small sub-components ───────────────────────────────────────────────────── function AvatarAI() { @@ -78,6 +141,8 @@ interface MessageRowProps { rating: "up" | "down" | null; onRate: (id: string, rating: "up" | "down") => void; onCopy: (content: string) => void; + canRegenerate: boolean; + regenerateTitle: string; onRegenerate: (id: string) => void; isRegenerating: boolean; } @@ -87,6 +152,8 @@ function MessageRow({ rating, onRate, onCopy, + canRegenerate, + regenerateTitle, onRegenerate, isRegenerating, }: MessageRowProps) { @@ -158,7 +225,7 @@ function MessageRow({ li: ({ children }) => (
  • {children}
  • ), - code: ({ inline, children }: any) => + code: ({ inline, children }: MarkdownCodeProps) => inline ? ( {children} @@ -240,9 +307,9 @@ function MessageRow({ /> onRegenerate(message.id)} - disabled={isRegenerating} + disabled={!canRegenerate || isRegenerating} > >({}); + const [chatSessionId, setChatSessionId] = useState(() => + getOrCreateChatSessionId() + ); const [availableModels, setAvailableModels] = useState([]); const [selectedModelId, setSelectedModelId] = useState(""); @@ -312,6 +382,12 @@ export function ChatPage({ onSwitchToWorkflow, onSignOut }: ChatPageProps) { const [docs, setDocs] = useState>(new Map()); const [contextError, setContextError] = useState(false); + const [contextUsage, setContextUsage] = useState(null); + const [historyOpen, setHistoryOpen] = useState(false); + const [chatHistory, setChatHistory] = useState([]); + const [historyLoading, setHistoryLoading] = useState(false); + const [historyDeleteTarget, setHistoryDeleteTarget] = + useState(null); const removeDoc = useCallback((tempId: string) => { setDocs((prev) => { @@ -332,6 +408,7 @@ export function ChatPage({ onSwitchToWorkflow, onSignOut }: ChatPageProps) { const messagesEndRef = useRef(null); const fileInputRef = useRef(null); const textareaRef = useRef(null); + const activeRequestSessionIdRef = useRef(null); // ── Models ────────────────────────────────────────────────────────────────── useEffect(() => { @@ -350,9 +427,35 @@ export function ChatPage({ onSwitchToWorkflow, onSignOut }: ChatPageProps) { messagesEndRef.current?.scrollIntoView({ behavior: "smooth" }); }, [messages, isLoading]); + const fetchChatHistory = useCallback(async () => { + try { + setHistoryLoading(true); + const token = await getValidToken(); + if (!token) throw new Error("Not authenticated"); + const response = await fetch("/api/chat/history", { + headers: { Authorization: `Bearer ${token}` }, + }); + if (!response.ok) throw new Error("Failed to load chat history"); + const data = await response.json(); + setChatHistory(data.chats ?? []); + } catch (err: unknown) { + const msg = + err instanceof Error ? err.message : "Failed to load chat history"; + toast.error(msg); + setChatHistory([]); + } finally { + setHistoryLoading(false); + } + }, []); + + useEffect(() => { + if (historyOpen) { + fetchChatHistory(); + } + }, [fetchChatHistory, historyOpen]); // ── Document processing ───────────────────────────────────────────────────── const processFile = useCallback(async (file: File) => { - const tempId = crypto.randomUUID(); + const tempId = createChatSessionId(); setDocs((prev) => { const active = Array.from(prev.values()).filter( @@ -437,6 +540,7 @@ export function ChatPage({ onSwitchToWorkflow, onSignOut }: ChatPageProps) { (e: React.ChangeEvent) => { const files = Array.from(e.target.files ?? []); if (fileInputRef.current) fileInputRef.current.value = ""; + const available = MAX_DOCS - Array.from(docs.values()).filter((d) => d.status !== "error").length; @@ -473,6 +577,7 @@ export function ChatPage({ onSwitchToWorkflow, onSignOut }: ChatPageProps) { e.preventDefault(); dragCounter.current = 0; setIsDragOver(false); + const files = Array.from(e.dataTransfer.files); const available = MAX_DOCS - @@ -497,9 +602,9 @@ export function ChatPage({ onSwitchToWorkflow, onSignOut }: ChatPageProps) { if (!model) return null; const isGemini = model.provider === "Google Gemini"; const isAnthropic = model.provider === "Anthropic"; + const chatModel = model as ChatModelConfig; const isLlama = - model.provider === "Meta Llama" || - (model as any).model_type === "azure-llama"; + model.provider === "Meta Llama" || chatModel.model_type === "azure-llama"; return { modelType: isGemini ? "gemini" @@ -533,6 +638,8 @@ export function ChatPage({ onSwitchToWorkflow, onSignOut }: ChatPageProps) { return; } + const requestChatSessionId = chatSessionId; + activeRequestSessionIdRef.current = requestChatSessionId; setContextError(false); setIsLoading(true); @@ -561,6 +668,7 @@ export function ChatPage({ onSwitchToWorkflow, onSignOut }: ChatPageProps) { Authorization: `Bearer ${token}`, }, body: JSON.stringify({ + chat_session_id: requestChatSessionId, query, document_markdown: documentMarkdown, model_type: modelConfig.modelType, @@ -571,18 +679,40 @@ export function ChatPage({ onSwitchToWorkflow, onSignOut }: ChatPageProps) { }); const data = await res.json(); - if (!res.ok || !data.success) - throw new Error(data.error || "Request failed"); + + if (activeRequestSessionIdRef.current !== requestChatSessionId) { + return; + } + + if (data.context_usage) { + setContextUsage(data.context_usage); + } + + if (!res.ok || !data.success) { + const msg = data.error || "Request failed"; + if ( + data.error_code === "context_window_exceeded" || + isContextWindowError(msg) + ) { + setContextError(true); + return; + } + throw new Error(msg); + } setMessages((prev) => [ ...prev, { - id: crypto.randomUUID(), + id: createChatSessionId(), role: "assistant", content: data.response, }, ]); } catch (err: unknown) { + if (activeRequestSessionIdRef.current !== requestChatSessionId) { + return; + } + const msg = err instanceof Error ? err.message : "Something went wrong"; if (isContextWindowError(msg)) { setContextError(true); @@ -590,17 +720,20 @@ export function ChatPage({ onSwitchToWorkflow, onSignOut }: ChatPageProps) { setMessages((prev) => [ ...prev, { - id: crypto.randomUUID(), + id: createChatSessionId(), role: "assistant", content: `Error: ${msg}`, }, ]); } } finally { - setIsLoading(false); + if (activeRequestSessionIdRef.current === requestChatSessionId) { + activeRequestSessionIdRef.current = null; + setIsLoading(false); + } } }, - [docs, getModelConfig] + [chatSessionId, docs, getModelConfig] ); const handleSend = useCallback(async () => { @@ -608,7 +741,7 @@ export function ChatPage({ onSwitchToWorkflow, onSignOut }: ChatPageProps) { if (!trimmed || isLoading) return; setMessages((prev) => [ ...prev, - { id: crypto.randomUUID(), role: "user", content: trimmed }, + { id: createChatSessionId(), role: "user", content: trimmed }, ]); setInput(""); // Reset textarea height @@ -618,20 +751,83 @@ export function ChatPage({ onSwitchToWorkflow, onSignOut }: ChatPageProps) { await sendQuery(trimmed); }, [input, isLoading, sendQuery]); - const handleRegenerate = useCallback( - async (assistantMsgId: string) => { - // Find the user message that preceded this assistant message - const idx = messages.findIndex((m) => m.id === assistantMsgId); - if (idx < 1) return; - const preceding = messages[idx - 1]; - if (preceding.role !== "user") return; - - // Remove the assistant message and resend - setMessages((prev) => prev.filter((m) => m.id !== assistantMsgId)); - await sendQuery(preceding.content); - }, - [messages, sendQuery] - ); + const handleRegenerate = useCallback(async (_assistantMsgId: string) => { + toast.info(REGENERATE_DISABLED_WITH_MEMORY_TITLE); + }, []); + + const handleNewChat = useCallback(() => { + const nextSessionId = resetChatSessionId(); + activeRequestSessionIdRef.current = nextSessionId; + setIsLoading(false); + setChatSessionId(nextSessionId); + setMessages([]); + setRatings({}); + setContextError(false); + setContextUsage(null); + setDocs(new Map()); + }, []); + + const handleOpenHistory = useCallback(() => { + setHistoryOpen(true); + }, []); + + const handleRestoreChat = useCallback(async (history: ChatHistorySummary) => { + try { + const token = await getValidToken(); + if (!token) throw new Error("Not authenticated"); + const response = await fetch( + `/api/chat/history/${encodeURIComponent(history.chat_session_id)}`, + { + headers: { Authorization: `Bearer ${token}` }, + } + ); + if (!response.ok) throw new Error("Failed to restore chat"); + const data: ChatHistoryDetail = await response.json(); + activeRequestSessionIdRef.current = null; + persistChatSessionId(data.chat_session_id); + setChatSessionId(data.chat_session_id); + setMessages(data.messages ?? []); + setRatings({}); + setContextError(false); + setContextUsage(data.context_usage ?? null); + setDocs(new Map()); + setHistoryOpen(false); + toast.success("Chat restored"); + } catch (err: unknown) { + const msg = err instanceof Error ? err.message : "Failed to restore chat"; + toast.error(msg); + } + }, []); + + const handleDeleteChat = useCallback(async () => { + if (!historyDeleteTarget) return; + try { + const token = await getValidToken(); + if (!token) throw new Error("Not authenticated"); + const response = await fetch( + `/api/chat/history/${encodeURIComponent(historyDeleteTarget.chat_session_id)}`, + { + method: "DELETE", + headers: { Authorization: `Bearer ${token}` }, + } + ); + if (!response.ok) throw new Error("Failed to delete chat"); + setChatHistory((prev) => + prev.filter( + (chat) => chat.chat_session_id !== historyDeleteTarget.chat_session_id + ) + ); + if (historyDeleteTarget.chat_session_id === chatSessionId) { + handleNewChat(); + } + toast.success("Chat history deleted"); + } catch (err: unknown) { + const msg = err instanceof Error ? err.message : "Failed to delete chat"; + toast.error(msg); + } finally { + setHistoryDeleteTarget(null); + } + }, [chatSessionId, handleNewChat, historyDeleteTarget]); const handleKeyDown = useCallback( (e: React.KeyboardEvent) => { @@ -655,6 +851,48 @@ export function ChatPage({ onSwitchToWorkflow, onSignOut }: ChatPageProps) { }, []); const canSend = !!input.trim() && !isLoading && !!selectedModelId; + const contextPercentage = + typeof contextUsage?.percentage === "number" + ? contextUsage.percentage + : null; + const hasEstimatedContextTokens = (contextUsage?.estimated_tokens ?? 0) > 0; + const contextPercentageLabel = + contextPercentage === null + ? "estimated" + : hasEstimatedContextTokens && contextPercentage < 0.1 + ? "<0.1%" + : `${contextPercentage.toFixed(contextPercentage < 10 ? 1 : 0)}%`; + const contextTokenLabel = + contextUsage && contextUsage.max_context_tokens + ? `${contextUsage.estimated_tokens.toLocaleString()} / ${contextUsage.max_context_tokens.toLocaleString()} tokens` + : null; + const contextVisibleTokenLabel = contextUsage + ? `${contextUsage.estimated_tokens.toLocaleString()} tokens` + : null; + const contextLevel = + contextPercentage === null + ? "normal" + : contextPercentage >= 90 + ? "critical" + : contextPercentage >= 75 + ? "warning" + : "normal"; + const contextCircleValue = Math.min( + Math.max( + contextPercentage === null + ? 0 + : hasEstimatedContextTokens + ? Math.max(contextPercentage, 0.1) + : contextPercentage, + 0 + ), + 100 + ); + const contextCircleRadius = 9; + const contextCircleCircumference = 2 * Math.PI * contextCircleRadius; + const contextCircleOffset = + contextCircleCircumference * + (1 - (contextPercentage === null ? 0 : contextCircleValue / 100)); return (
    + + {onSwitchToWorkflow && ( + +
    + ))} + + )} + + + + + { + if (!open) setHistoryDeleteTarget(null); + }} + > + + + Delete Chat History + + This removes the selected chat history from the database. + + + + Cancel + + Delete + + + + {/* ── Message list ────────────────────────────────────────────────────── */}
    @@ -759,6 +1105,8 @@ export function ChatPage({ onSwitchToWorkflow, onSignOut }: ChatPageProps) { rating={ratings[msg.id] ?? null} onRate={handleRate} onCopy={handleCopy} + canRegenerate={false} + regenerateTitle={REGENERATE_DISABLED_WITH_MEMORY_TITLE} onRegenerate={handleRegenerate} isRegenerating={isLoading} /> @@ -938,10 +1286,64 @@ export function ChatPage({ onSwitchToWorkflow, onSignOut }: ChatPageProps) {
    {/* Disclaimer */} -

    - Science-GPT can make mistakes. Check important information - carefully. -

    +
    + + Science-GPT can make mistakes. Check important information + carefully. + + {contextUsage && ( + 0 + ? "Older chat turns are represented by a rolling summary." + : "" + }`} + > + + + Context {contextPercentageLabel} + {contextVisibleTokenLabel + ? ` · ${contextVisibleTokenLabel}` + : ""} + + + )} +
    diff --git a/frontend/utils/chatSession.ts b/frontend/utils/chatSession.ts new file mode 100644 index 0000000..fcf83ad --- /dev/null +++ b/frontend/utils/chatSession.ts @@ -0,0 +1,82 @@ +const CHAT_SESSION_STORAGE_KEY = "science-gpt.chat_session_id"; + +let inMemoryChatSessionId: string | null = null; + +function getCryptoApi(): Crypto | undefined { + if (typeof globalThis === "undefined") return undefined; + return globalThis.crypto; +} + +function getSessionStorage(): Storage | null { + if (typeof globalThis === "undefined" || !("sessionStorage" in globalThis)) { + return null; + } + + try { + return globalThis.sessionStorage; + } catch { + return null; + } +} + +function readStoredChatSessionId(): string | null { + const storage = getSessionStorage(); + if (!storage) return null; + + try { + return storage.getItem(CHAT_SESSION_STORAGE_KEY); + } catch { + return null; + } +} + +function writeStoredChatSessionId(sessionId: string): void { + const storage = getSessionStorage(); + if (!storage) return; + + try { + storage.setItem(CHAT_SESSION_STORAGE_KEY, sessionId); + } catch { + // Ignore blocked sessionStorage and fall back to in-memory IDs. + } +} + +export function createChatSessionId(): string { + const cryptoApi = getCryptoApi(); + if (cryptoApi?.randomUUID) { + return cryptoApi.randomUUID(); + } + + if (cryptoApi?.getRandomValues) { + const bytes = new Uint8Array(16); + cryptoApi.getRandomValues(bytes); + const randomHex = Array.from(bytes, (byte) => + byte.toString(16).padStart(2, "0") + ).join(""); + return `chat-${randomHex}`; + } + + throw new Error("Secure random ID generation is unavailable."); +} + +export function getOrCreateChatSessionId(): string { + const existing = readStoredChatSessionId() ?? inMemoryChatSessionId; + if (existing) return existing; + + const next = createChatSessionId(); + inMemoryChatSessionId = next; + writeStoredChatSessionId(next); + return next; +} + +export function resetChatSessionId(): string { + const next = createChatSessionId(); + inMemoryChatSessionId = next; + writeStoredChatSessionId(next); + return next; +} + +export function setChatSessionId(sessionId: string): void { + inMemoryChatSessionId = sessionId; + writeStoredChatSessionId(sessionId); +}