diff --git a/pkg-py/CHANGELOG.md b/pkg-py/CHANGELOG.md index bd65723be..5d2029c9e 100644 --- a/pkg-py/CHANGELOG.md +++ b/pkg-py/CHANGELOG.md @@ -5,6 +5,18 @@ All notable changes to this project will be documented in this file. The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.1.0/), and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). +## [0.6.1] - 2026-08-14 + +### Bug fixes + +* Conversation history now rejects records written with unsupported future schema versions instead of attempting an unsafe downgrade. + +* Fixed a race between the chat greeting and conversation history restore: reloading a page that restored a previous conversation could briefly flash the app's greeting, and starting a new chat after a session began with a restored conversation could fail to show any greeting at all. Greeting resolution now defers to history's own restore decision instead of racing the client's independent greeting request. + +* Restoring a bookmark that contains a malformed message (for instance one written by an incompatible shinychat version) now warns and skips just that message, instead of raising and dropping every message after it. + +* Fixed conversation history and bookmarks failing to serialize a chatlas `ContentToolResult` when its supported dictionary-form `extra["display"]` contains HTML. Dictionary displays are now normalized through `ToolResultDisplay` before JSON serialization. (Related to #295) + ## [0.6.0] - 2026-07-06 ### New features diff --git a/pkg-py/src/shinychat/_chat.py b/pkg-py/src/shinychat/_chat.py index 4d21ffd94..7511ff45d 100644 --- a/pkg-py/src/shinychat/_chat.py +++ b/pkg-py/src/shinychat/_chat.py @@ -3,6 +3,7 @@ import inspect import json import re +import warnings from contextlib import asynccontextmanager from dataclasses import dataclass from typing import ( @@ -472,9 +473,12 @@ async def _on_slash_command(): self._setup_client(client) if greeting is not None: - from ._chat_client import setup_greeting + if self.history._controller is not None: + self.history.setup_greeting(greeting) + else: + from ._chat_client import setup_greeting - setup_greeting(self, greeting, self._session) + setup_greeting(self, greeting, self._session) def _setup_client( self, @@ -1363,10 +1367,26 @@ async def _restore_bookmark_message(self, message_dict: Any) -> None: try: stored = StoredMessage.model_validate(message_dict) except ValidationError as e: - raise ValueError( - "Cannot restore bookmark message: invalid or missing fields " - "(bookmark likely written by an incompatible shinychat version)." - ) from e + # Skip rather than raise: raising here would abort the caller's + # restore loop, silently dropping every message after this one + # too (Shiny's on_restore error handling only shows a banner, it + # doesn't resume the loop). + # + # include_input=False: the default error string embeds the + # offending value, which for a chat message is arbitrary (and + # possibly sensitive) message content -- keep the warning to + # locations/reasons only. + details = "; ".join( + f"{'.'.join(str(p) for p in err['loc'])}: {err['msg']}" + for err in e.errors(include_input=False) + ) + warnings.warn( + "Skipping malformed bookmarked chat message: invalid or " + "missing fields (bookmark likely written by an incompatible " + f"shinychat version). {details}", + stacklevel=2, + ) + return self._store_message(stored) await self._send_append_message(stored) diff --git a/pkg-py/src/shinychat/_chat_bookmark.py b/pkg-py/src/shinychat/_chat_bookmark.py index 72da68dfa..20317f7bd 100644 --- a/pkg-py/src/shinychat/_chat_bookmark.py +++ b/pkg-py/src/shinychat/_chat_bookmark.py @@ -11,6 +11,8 @@ runtime_checkable, ) +from ._chatlas_serialization import serialize_chatlas_turn + if TYPE_CHECKING: from chatlas import Chat from htmltools import Tagified @@ -85,7 +87,7 @@ async def get_state() -> Jsonifiable: turns: list[Turn[Any]] = client.get_turns() return { "version": 1, - "turns": [turn.model_dump(mode="json") for turn in turns], + "turns": [serialize_chatlas_turn(turn) for turn in turns], } return get_state diff --git a/pkg-py/src/shinychat/_chat_client.py b/pkg-py/src/shinychat/_chat_client.py index 1bd55c69e..58b9e9902 100644 --- a/pkg-py/src/shinychat/_chat_client.py +++ b/pkg-py/src/shinychat/_chat_client.py @@ -187,6 +187,32 @@ def messages_to_turns( return turns +async def resolve_greeting( + chat: "Chat", + greeting: "str | HTML | Tag | TagList | ChatGreeting | Callable[..., Any]", +) -> None: + """Resolve `greeting` (static content or a callable) and set it on `chat`.""" + from htmltools import HTML, Tag, TagList + + from ._chat_types import ChatGreeting + + if isinstance(greeting, (str, HTML, Tag, TagList, ChatGreeting)): + return await chat.set_greeting(greeting) + + sig = inspect.signature(greeting) + if "client" in sig.parameters and chat.client is not None: + client_copy = copy.deepcopy(chat.client.value) + client_copy.set_turns([]) + result = greeting(client=client_copy) + else: + result = greeting() + + if inspect.isawaitable(result): + result = await result + + await chat.set_greeting(result) # type: ignore[arg-type] + + def setup_greeting( chat: "Chat", greeting: "str | HTML | Tag | TagList | ChatGreeting | Callable[..., Any] | None", @@ -195,33 +221,16 @@ def setup_greeting( if greeting is None: return - from htmltools import HTML, Tag, TagList from shiny import reactive from shiny.module import ResolvedId from shiny.session import session_context - from ._chat_types import ChatGreeting - with session_context(session): greeting_requested_id = ResolvedId(f"{chat.id}_greeting_requested") @reactive.effect @reactive.event(session.input[greeting_requested_id]) async def _on_greeting_requested() -> None: - if isinstance(greeting, (str, HTML, Tag, TagList, ChatGreeting)): - return await chat.set_greeting(greeting) - - sig = inspect.signature(greeting) - if "client" in sig.parameters and chat.client is not None: - client_copy = copy.deepcopy(chat.client.value) - client_copy.set_turns([]) - result = greeting(client=client_copy) - else: - result = greeting() - - if inspect.isawaitable(result): - result = await result - - await chat.set_greeting(result) # type: ignore[arg-type] + await resolve_greeting(chat, greeting) chat._effects.append(_on_greeting_requested) diff --git a/pkg-py/src/shinychat/_chatlas_serialization.py b/pkg-py/src/shinychat/_chatlas_serialization.py new file mode 100644 index 000000000..dc08498bc --- /dev/null +++ b/pkg-py/src/shinychat/_chatlas_serialization.py @@ -0,0 +1,36 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +if TYPE_CHECKING: + from chatlas import Turn + + +def serialize_chatlas_turn(turn: Turn[Any]) -> dict[str, Any]: + """Serialize a chatlas turn after normalizing supported rich displays.""" + from chatlas import ContentToolResult + + from ._chat_normalize_chatlas import ToolResultDisplay + + normalized = turn + for index, content in enumerate(turn.contents): + if not isinstance(content, ContentToolResult): + continue + extra = content.extra + if not isinstance(extra, dict) or not isinstance( + extra.get("display"), dict + ): + continue + + if normalized is turn: + normalized = turn.model_copy(deep=True) + normalized_content = normalized.contents[index] + assert isinstance(normalized_content, ContentToolResult) + normalized_content.extra = dict(normalized_content.extra) + display = normalized_content.extra["display"] + normalized_content.extra["display"] = { + **display, + **ToolResultDisplay(**display).model_dump(mode="json"), + } + + return normalized.model_dump(mode="json") diff --git a/pkg-py/src/shinychat/_history.py b/pkg-py/src/shinychat/_history.py index fc453aeca..5923ec754 100644 --- a/pkg-py/src/shinychat/_history.py +++ b/pkg-py/src/shinychat/_history.py @@ -23,12 +23,18 @@ fallback_title, generate_title, ) -from ._history_types import ConversationRecord, new_conversation_record +from ._history_types import ( + ConversationRecord, + check_schema_version, + new_conversation_record, +) if TYPE_CHECKING: + from htmltools import HTML, Tag, TagList from shiny.module import ResolvedId from ._chat import Chat + from ._chat_types import ChatGreeting @dataclasses.dataclass(frozen=True) @@ -232,6 +238,9 @@ def __init__( ) = None # Internal hook: fired before a conversation is removed from the store. self.on_evict: Callable[[str], Awaitable[None]] | None = None + # Internal hook: fired whenever it is known whether the active + # conversation is a restore (True) or a fresh/new one (False). + self.on_settled: Callable[[bool], Awaitable[None]] | None = None self.max_store_bytes: int | None = max_store_bytes self._title_task: asyncio.Task[None] | None = None # replay_ui awaits per message, so on_response can fire mid-replay; @@ -242,6 +251,14 @@ def __init__( self._suppress_next_save: bool = False self._over_budget_warned: bool = False + async def _get_record( + self, partition: ConversationPartition, conv_id: str + ) -> ConversationRecord | None: + record = await self.store.get(partition, conv_id) + if record is not None: + check_schema_version(record.schema_version) + return record + # -- save ----------------------------------------------------------- async def on_response(self) -> None: @@ -340,6 +357,11 @@ def cancel_pending(self) -> None: if self._title_task is not None and not self._title_task.done(): self._title_task.cancel() + async def notify_settled(self, restored: bool) -> None: + """Called whenever the active conversation's restore state is known.""" + if self.on_settled is not None: + await self.on_settled(restored) + async def _evict_one(self, conv_id: str) -> None: assert self.partition is not None if self.on_evict is not None: @@ -402,7 +424,7 @@ async def switch_to(self, conv_id: str) -> None: return # Load BEFORE mutating anything: a failed load must leave the # current conversation untouched. - target = await self.store.get(self.partition, conv_id) + target = await self._get_record(self.partition, conv_id) if target is None: raise RuntimeError(f"Conversation {conv_id!r} no longer exists.") @@ -427,6 +449,7 @@ async def new_chat(self) -> None: self.record = None if self.on_active_id_change is not None: await self.on_active_id_change(None) + await self.notify_settled(False) await self.send_history_update() async def replay_ui(self, record: ConversationRecord) -> None: @@ -434,6 +457,7 @@ async def replay_ui(self, record: ConversationRecord) -> None: self._suppress_next_save = True try: await self.chat.clear_messages() + await self.chat.set_greeting(None) for node_id in record.path_node_ids(): node = record.nodes[node_id] stored = node.ui or [ @@ -466,7 +490,7 @@ async def rename(self, conv_id: str, title: str) -> None: record = ( self.record if self.record is not None and self.record.id == conv_id - else await self.store.get(self.partition, conv_id) + else await self._get_record(self.partition, conv_id) ) if record is None: return @@ -547,6 +571,7 @@ def __init__( ) -> None: self._chat = chat self._started: bool = False + self._controller: HistoryController | None = None self._save_callbacks: "list[Callable[[dict[str, Any]], None]]" = [] self._restore_callbacks: "list[Callable[[dict[str, Any]], None]]" = [] cfg = config if config is not None else HistoryOptions() @@ -616,6 +641,23 @@ def _(values): self._restore_callbacks.append(fn) return fn + def setup_greeting( + self, + greeting: "str | HTML | Tag | TagList | ChatGreeting | Callable[..., Any]", + ) -> None: + """Resolve a greeting after history determines whether it restored.""" + from ._chat_client import resolve_greeting + + chat = self._chat + controller = self._controller + assert controller is not None + + async def _on_settled(restored: bool) -> None: + if not restored: + await resolve_greeting(chat, greeting) + + controller.on_settled = _on_settled + def _start(self) -> None: chat = self._chat chat_client = chat.client @@ -657,6 +699,7 @@ def _start(self) -> None: restore_callbacks=self._restore_callbacks, max_store_bytes=max_store_bytes, ) + self._controller = controller if restore_mode == "url": @@ -726,7 +769,7 @@ async def _on_evict(conv_id: str) -> None: if controller.partition is None: rec = None else: - rec = await controller.store.get( + rec = await controller._get_record( controller.partition, conv_id ) state_id = ( @@ -818,9 +861,13 @@ async def _init_history(): restored_conv_id = str(raw_id) if raw_id else None if restored_conv_id is not None: - target = await controller.store.get( - controller.partition, restored_conv_id - ) + try: + target = await controller._get_record( + controller.partition, restored_conv_id + ) + except Exception as e: + await notify_error("Could not load conversation", e) + target = None if target is not None: adapter.set_turns_json(target.path_turns()) await controller.replay_ui(target) @@ -829,6 +876,7 @@ async def _init_history(): controller.record = target await controller.send_history_update() initialized = True + await controller.notify_settled(True) return # Priority 2: restore from the mode-specific ID source. @@ -848,9 +896,13 @@ async def _init_history(): current_id = None if current_id: - pointed = await controller.store.get( - controller.partition, current_id - ) + try: + pointed = await controller._get_record( + controller.partition, current_id + ) + except Exception as e: + await notify_error("Could not load conversation", e) + pointed = None if pointed is not None: adapter.set_turns_json(pointed.path_turns()) await controller.replay_ui(pointed) @@ -858,6 +910,7 @@ async def _init_history(): controller.record = pointed await controller.send_history_update() initialized = True + await controller.notify_settled(controller.record is not None) @reactive.effect @reactive.event(chat.messages, ignore_init=True) diff --git a/pkg-py/src/shinychat/_history_client.py b/pkg-py/src/shinychat/_history_client.py index d3992e277..b9b99bfef 100644 --- a/pkg-py/src/shinychat/_history_client.py +++ b/pkg-py/src/shinychat/_history_client.py @@ -4,6 +4,7 @@ from ._chat_bookmark import is_chatlas_chat_client from ._chat_client import ChatClient +from ._chatlas_serialization import serialize_chatlas_turn @runtime_checkable @@ -36,7 +37,7 @@ def get_turns_json(self) -> list[dict[str, Any]]: raw = self._turns_client() turns = raw.get_turns() if is_chatlas_chat_client(raw): - return [t.model_dump(mode="json") for t in turns] + return [serialize_chatlas_turn(t) for t in turns] return list(turns) def get_turns_grouped(self) -> list[list[dict[str, Any]]]: diff --git a/pkg-py/src/shinychat/_history_store.py b/pkg-py/src/shinychat/_history_store.py index 30fc734f7..6d5c0d812 100644 --- a/pkg-py/src/shinychat/_history_store.py +++ b/pkg-py/src/shinychat/_history_store.py @@ -16,6 +16,8 @@ ConversationMeta, ConversationNode, ConversationRecord, + UnsupportedSchemaVersionError, + check_schema_version, ) logger = logging.getLogger(__name__) @@ -166,6 +168,9 @@ async def list( continue try: raw = json.loads(record_file.read_text(encoding="utf-8")) + schema_version = check_schema_version( + raw.get("schema_version") + ) nodes_raw = raw.get("nodes", {}) nodes = {} for nid, nd in nodes_raw.items(): @@ -175,7 +180,7 @@ async def list( turns=[], ) rec = ConversationRecord( - schema_version=raw.get("schema_version", 1), + schema_version=schema_version, id=raw["id"], title=raw["title"], title_source=raw.get("title_source"), @@ -192,6 +197,8 @@ async def list( f.stat().st_size for f in d.iterdir() if f.is_file() ) metas.append(rec.meta(size_bytes=size_bytes)) + except UnsupportedSchemaVersionError: + raise except Exception as e: logger.warning("Unreadable conversation %s: %s", d.name, e) continue @@ -211,6 +218,7 @@ async def get( return None raw = json.loads(record_file.read_text(encoding="utf-8")) + schema_version = check_schema_version(raw.get("schema_version")) turns_map: dict[int, dict[str, Any]] = {} turns_file = conv_dir / "turns.jsonl" @@ -248,7 +256,7 @@ async def get( ) return ConversationRecord( - schema_version=raw.get("schema_version", 1), + schema_version=schema_version, id=raw["id"], title=raw["title"], title_source=raw.get("title_source"), @@ -266,8 +274,18 @@ async def get( async def put( self, partition: ConversationPartition, record: ConversationRecord ) -> None: + check_schema_version(record.schema_version) + partition_dir = await self._partition_dir(partition) conv_dir = safe_conv_path(partition_dir, record.id) + # Validate the on-disk schema version before creating/modifying + # anything, so an unsupported existing record is rejected fail-closed + # rather than partially overwritten. + record_file = conv_dir / "record.json" + if record_file.is_file(): + raw = json.loads(record_file.read_text(encoding="utf-8")) + check_schema_version(raw.get("schema_version")) + conv_dir.mkdir(parents=True, exist_ok=True) ws = self._get_or_init_write_state(partition, record.id, conv_dir) @@ -411,6 +429,8 @@ async def get( async def put( self, partition: ConversationPartition, record: ConversationRecord ) -> None: + check_schema_version(record.schema_version) + if partition not in self._data: self._data[partition] = {} self._data[partition][record.id] = record diff --git a/pkg-py/src/shinychat/_history_types.py b/pkg-py/src/shinychat/_history_types.py index c4f2daf0b..a18241377 100644 --- a/pkg-py/src/shinychat/_history_types.py +++ b/pkg-py/src/shinychat/_history_types.py @@ -43,6 +43,29 @@ class ConversationNode(BaseModel): ui: list[dict[str, Any]] | None = None +MIN_SCHEMA_VERSION = 1 +MAX_SCHEMA_VERSION = 1 + + +class UnsupportedSchemaVersionError(ValueError): + def __init__(self, version: object) -> None: + super().__init__( + f"Unsupported conversation record schema version: {version!r} " + f"(supported: {MIN_SCHEMA_VERSION}-{MAX_SCHEMA_VERSION})" + ) + + +def check_schema_version(version: object) -> int: + # None means the record predates schema_version entirely; treat as 1. + version = 1 if version is None else version + if ( + type(version) is int + and MIN_SCHEMA_VERSION <= version <= MAX_SCHEMA_VERSION + ): + return version + raise UnsupportedSchemaVersionError(version) + + class ConversationRecord(BaseModel): schema_version: int = 1 id: str diff --git a/pkg-py/tests/playwright/chat/history_greeting_no_flash/app.py b/pkg-py/tests/playwright/chat/history_greeting_no_flash/app.py new file mode 100644 index 000000000..c36c21ce4 --- /dev/null +++ b/pkg-py/tests/playwright/chat/history_greeting_no_flash/app.py @@ -0,0 +1,71 @@ +from __future__ import annotations + +import os +import tempfile +from typing import Any, AsyncGenerator +from unittest.mock import MagicMock + +import chatlas +from chatlas import Turn +from chatlas._turn import AssistantTurn +from shiny import App, Inputs, Outputs, Session, ui +from shinychat import Chat, chat_greeting, chat_ui +from shinychat.types import FileConversationStore, HistoryOptions + + +class EchoChatClient(chatlas.Chat): + def __init__(self) -> None: + provider = MagicMock() + provider.name = "echo" + provider.model = "echo" + super().__init__(provider) + + async def stream_async( + self, *args: Any, **kwargs: Any + ) -> AsyncGenerator[str, None]: # type: ignore[override] + user_input = str(args[0]) if args else "" + self._turns.extend( + [ + Turn(role="user", contents=user_input), + AssistantTurn(contents=f"echo: {user_input}"), + ] + ) + + async def _gen() -> AsyncGenerator[str, None]: + yield f"echo: {user_input}" + + return _gen() + + +_store_dir_cache: dict[int, str] = {} + + +def _get_store_dir() -> str: + pid = os.getpid() + if pid not in _store_dir_cache: + _store_dir_cache[pid] = tempfile.mkdtemp( + prefix="shinychat-greeting-history-" + ) + return _store_dir_cache[pid] + + +def app_ui(request: object) -> ui.Tag: + return chat_ui("chat", greeting=chat_greeting("## Welcome!", persistent=True)) + + +def server(input: Inputs, output: Outputs, session: Session) -> None: + store_dir = _get_store_dir() + + Chat( + id="chat", + client=EchoChatClient(), + greeting=chat_greeting("## Welcome!", persistent=True), + history=HistoryOptions( + store=FileConversationStore(dir=store_dir), + scope="test-user", + title=None, + ), + ) + + +app = App(app_ui, server, bookmark_store="server") diff --git a/pkg-py/tests/playwright/chat/history_greeting_no_flash/test_history_greeting_no_flash.py b/pkg-py/tests/playwright/chat/history_greeting_no_flash/test_history_greeting_no_flash.py new file mode 100644 index 000000000..bd56da684 --- /dev/null +++ b/pkg-py/tests/playwright/chat/history_greeting_no_flash/test_history_greeting_no_flash.py @@ -0,0 +1,57 @@ +from __future__ import annotations + +from playwright.sync_api import Page, expect +from shiny.run import ShinyAppProc +from shinychat.playwright import ChatController + + +def open_drawer(page: Page) -> None: + page.locator(".shiny-chat-history-trigger").click() + expect(page.locator(".shiny-chat-history-drawer")).to_be_visible() + + +def test_greeting_does_not_flash_on_restored_conversation( + page: Page, local_app: ShinyAppProc +) -> None: + """ + Reloading a page that restores a previous conversation must not show + the app's greeting again — it belongs to brand-new conversations only. + """ + page.goto(local_app.url) + chat = ChatController(page, "chat") + expect(chat.loc).to_be_visible(timeout=30_000) + chat.expect_greeting("Welcome", timeout=30_000) + + chat.set_user_input("hello") + chat.send_user_input(method="enter") + chat.expect_latest_message("echo: hello", timeout=30_000) + + # Ensure the active conversation ID has been written to localStorage before + # reloading. The ID lands once the history_update action reaches the client, + # which happens after send_history_update(). The drawer shows the item at + # that point, so opening it is a reliable sync point. + open_drawer(page) + expect(page.locator(".shiny-chat-history-item")).to_have_count( + 1, timeout=10_000 + ) + page.keyboard.press("Escape") + page.locator(".shiny-chat-history-drawer").wait_for(state="hidden") + + page.reload() + expect(chat.loc).to_be_visible(timeout=30_000) + + # Transcript must be restored ... + chat.expect_latest_message("echo: hello", timeout=30_000) + # ... and the greeting must never reappear. + expect(chat.loc_greeting).to_have_count(0, timeout=5_000) + + # Starting a new chat from a session that began by restoring a + # conversation must still resolve the app's greeting — it was never + # requested/resolved for *this* session, so there's nothing cached + # client-side to fall back on. + page.locator(".shiny-chat-history-trigger").click() + expect(page.locator(".shiny-chat-history-drawer")).to_be_visible() + page.locator(".shiny-chat-history-new").click() + expect(page.locator(".shiny-chat-history-drawer")).not_to_be_visible() + + chat.expect_greeting("Welcome", timeout=30_000) diff --git a/pkg-py/tests/playwright/chat/history_switch_clears_greeting/app.py b/pkg-py/tests/playwright/chat/history_switch_clears_greeting/app.py new file mode 100644 index 000000000..b69eacc88 --- /dev/null +++ b/pkg-py/tests/playwright/chat/history_switch_clears_greeting/app.py @@ -0,0 +1,71 @@ +from __future__ import annotations + +import os +import tempfile +from typing import Any, AsyncGenerator +from unittest.mock import MagicMock + +import chatlas +from chatlas import Turn +from chatlas._turn import AssistantTurn +from shiny import App, Inputs, Outputs, Session, ui +from shinychat import Chat, chat_greeting, chat_ui +from shinychat.types import FileConversationStore, HistoryOptions + + +class EchoChatClient(chatlas.Chat): + def __init__(self) -> None: + provider = MagicMock() + provider.name = "echo" + provider.model = "echo" + super().__init__(provider) + + async def stream_async( + self, *args: Any, **kwargs: Any + ) -> AsyncGenerator[str, None]: # type: ignore[override] + user_input = str(args[0]) if args else "" + self._turns.extend( + [ + Turn(role="user", contents=user_input), + AssistantTurn(contents=f"echo: {user_input}"), + ] + ) + + async def _gen() -> AsyncGenerator[str, None]: + yield f"echo: {user_input}" + + return _gen() + + +_store_dir_cache: dict[int, str] = {} + + +def _get_store_dir() -> str: + pid = os.getpid() + if pid not in _store_dir_cache: + _store_dir_cache[pid] = tempfile.mkdtemp( + prefix="shinychat-greeting-switch-" + ) + return _store_dir_cache[pid] + + +def app_ui(request: object) -> ui.Tag: + return chat_ui("chat", greeting=chat_greeting("## Welcome!", persistent=True)) + + +def server(input: Inputs, output: Outputs, session: Session) -> None: + store_dir = _get_store_dir() + + Chat( + id="chat", + client=EchoChatClient(), + greeting=chat_greeting("## Welcome!", persistent=True), + history=HistoryOptions( + store=FileConversationStore(dir=store_dir), + scope="test-user", + title=None, + ), + ) + + +app = App(app_ui, server, bookmark_store="server") diff --git a/pkg-py/tests/playwright/chat/history_switch_clears_greeting/test_history_switch_clears_greeting.py b/pkg-py/tests/playwright/chat/history_switch_clears_greeting/test_history_switch_clears_greeting.py new file mode 100644 index 000000000..441178de9 --- /dev/null +++ b/pkg-py/tests/playwright/chat/history_switch_clears_greeting/test_history_switch_clears_greeting.py @@ -0,0 +1,56 @@ +from __future__ import annotations + +from playwright.sync_api import Page, expect +from shiny.run import ShinyAppProc +from shinychat.playwright import ChatController + + +def open_drawer(page: Page) -> None: + page.locator(".shiny-chat-history-trigger").click() + expect(page.locator(".shiny-chat-history-drawer")).to_be_visible() + + +def test_greeting_clears_when_switching_to_old_conversation( + page: Page, local_app: ShinyAppProc +) -> None: + """ + Starting a fresh conversation shows the greeting; switching to a + previously-saved conversation via the history drawer must hide it. + + Selectors and flow mirror + `pkg-py/tests/playwright/chat/history/test_history.py::test_history_full_flow`. + """ + page.goto(local_app.url) + chat = ChatController(page, "chat") + expect(chat.loc).to_be_visible(timeout=30_000) + chat.expect_greeting("Welcome", timeout=30_000) + + chat.set_user_input("first conversation") + chat.send_user_input(method="enter") + chat.expect_latest_message("echo: first conversation", timeout=30_000) + + open_drawer(page) + items = page.locator(".shiny-chat-history-item") + expect(items).to_have_count(1, timeout=10_000) + + # New chat: clears the transcript (and, per this fix, the greeting + # should reappear here as a normal fresh-conversation greeting). + page.locator(".shiny-chat-history-new").click() + expect(page.locator(".shiny-chat-history-drawer")).not_to_be_visible() + + chat.set_user_input("second conversation") + chat.send_user_input(method="enter") + chat.expect_latest_message("echo: second conversation", timeout=30_000) + + open_drawer(page) + expect(page.locator(".shiny-chat-history-item")).to_have_count( + 2, timeout=10_000 + ) + + # Switch back to the first conversation via the drawer. + page.locator( + ".shiny-chat-history-item", has_text="first conversation" + ).click() + + chat.expect_latest_message("echo: first conversation", timeout=30_000) + expect(chat.loc_greeting).to_have_count(0, timeout=5_000) diff --git a/pkg-py/tests/pytest/test_bookmark_serialization.py b/pkg-py/tests/pytest/test_bookmark_serialization.py index b06c558d7..b207c1fd7 100644 --- a/pkg-py/tests/pytest/test_bookmark_serialization.py +++ b/pkg-py/tests/pytest/test_bookmark_serialization.py @@ -9,32 +9,46 @@ from __future__ import annotations import json +from typing import Any +import pytest from chatlas import ContentToolResult, Turn from htmltools import HTMLDependency, tags +from shinychat._chatlas_serialization import serialize_chatlas_turn from shinychat.types import ToolResultDisplay -def test_turn_serialization_with_htmldep_in_tool_result(): +@pytest.mark.parametrize("as_dict", [False, True]) +def test_turn_serialization_with_htmldep_in_tool_result(as_dict: bool): """Turn containing ToolResultDisplay with HTMLDependency round-trips through JSON.""" - display = ToolResultDisplay( + typed_display = ToolResultDisplay( html=tags.div( "Widget output", HTMLDependency("my-dep", "1.0", source={"subdir": "."}), ), title="My Widget", ) + display: Any = typed_display + if as_dict: + display = { + "html": typed_display.html, + "title": typed_display.title, + "application_metadata": {"widget_id": "my-widget"}, + } result = ContentToolResult(value="done", extra={"display": display}) turn = Turn(role="user", contents=[result]) - # This is what _chat_bookmark.py's get_chatlas_state does - dumped = turn.model_dump(mode="json") + dumped = serialize_chatlas_turn(turn) # Must be JSON-serializable json_str = json.dumps(dumped) # Verify the serialized dependencies are JSON dicts (not live HTMLDependency objects) display_data = dumped["contents"][0]["extra"]["display"] + if as_dict: + assert display_data["application_metadata"] == { + "widget_id": "my-widget" + } deps = display_data["html"]["dependencies"] assert len(deps) == 1 assert deps[0]["name"] == "my-dep" diff --git a/pkg-py/tests/pytest/test_chat.py b/pkg-py/tests/pytest/test_chat.py index 09b10931f..6033f05b9 100644 --- a/pkg-py/tests/pytest/test_chat.py +++ b/pkg-py/tests/pytest/test_chat.py @@ -700,6 +700,63 @@ def test_bookmark_omits_side_effect_only_slash_command(): ] +def test_restore_bookmark_message_warns_and_skips_malformed(): + # A malformed stored message (e.g. a bookmark written by an incompatible + # shinychat version) must not abort the whole restore loop -- letting it + # raise would hit Shiny's generic on_restore error handling, which shows + # a banner and silently drops every message after the bad one. + with session_context(test_session): + chat = Chat(id="chat_restore_malformed") + sent: list[dict[str, Any]] = [] + + async def _capture(action: Any, deps: Any = None) -> None: + sent.append(action) + + chat._send_action = _capture # type: ignore[method-assign] + + saved: list[Any] = [ + {"role": "user", "segments": [{"content": "before", "content_type": "markdown"}]}, + {"role": "user"}, # missing required `segments` + {"role": "assistant", "segments": [{"content": "after", "content_type": "markdown"}]}, + ] + + async def _exercise() -> None: + with pytest.warns(UserWarning, match="incompatible shinychat version"): + for message_dict in saved: + await chat._restore_bookmark_message(message_dict) + + run_async(_exercise) + + contents = [ + a["message"]["segments"][0]["content"] for a in sent if a["type"] == "message" + ] + assert contents == ["before", "after"] + + +def test_restore_bookmark_message_warning_omits_the_offending_value(): + # pydantic's default ValidationError string embeds the invalid input + # value, which for a chat message is arbitrary (and possibly sensitive) + # content -- the warning must not repeat it. + with session_context(test_session): + chat = Chat(id="chat_restore_malformed_no_leak") + secret = "sk-super-secret-token-do-not-log-me" + + async def _exercise() -> None: + with pytest.warns(UserWarning) as record: + # The invalid *value* here is the secret itself: content_type + # only accepts a fixed set of literals, so a bogus string + # fails validation with that string as the reported input. + await chat._restore_bookmark_message( + { + "role": "user", + "segments": [{"content": "hi", "content_type": secret}], + } + ) + assert secret not in str(record[0].message) + + run_async(_exercise) + + def test_user_input_reads_latest_stored(): from shiny import reactive from shinychat._chat import UserInput diff --git a/pkg-py/tests/pytest/test_greeting.py b/pkg-py/tests/pytest/test_greeting.py index 6d1ba34a4..e0c0c146a 100644 --- a/pkg-py/tests/pytest/test_greeting.py +++ b/pkg-py/tests/pytest/test_greeting.py @@ -9,6 +9,7 @@ from htmltools import HTML, HTMLDependency, tags from shiny.session import session_context from shinychat import Chat, chat_greeting, chat_ui +from shinychat._chat_client import resolve_greeting from shinychat._chat_types import ChatGreeting # --------------------------------------------------------------------------- @@ -431,3 +432,47 @@ def test_enable_bookmarking_excludes_greeting_dismissed(): assert "bm_chat_dis_greeting_dismissed" in bm_sess.bookmark.exclude +def test_resolve_greeting_static_string(): + chat, spy = _make_spy_chat() + + async def _run(): + await resolve_greeting(chat, "## Hi") + + _run_async(_run) + actions = _spy_actions(spy) + assert len(actions) == 1 + assert actions[0]["type"] == "greeting" + assert actions[0]["content"] == "## Hi" + + +def test_resolve_greeting_zero_arg_callable(): + chat, spy = _make_spy_chat() + called = False + + def _greeting(): + nonlocal called + called = True + return "## Generated" + + async def _run(): + await resolve_greeting(chat, _greeting) + + _run_async(_run) + assert called + actions = _spy_actions(spy) + assert actions[0]["content"] == "## Generated" + + +def test_resolve_greeting_awaitable_callable(): + chat, spy = _make_spy_chat() + + async def _greeting(): + return "## Async Generated" + + async def _run(): + await resolve_greeting(chat, _greeting) + + _run_async(_run) + actions = _spy_actions(spy) + assert actions[0]["content"] == "## Async Generated" + diff --git a/pkg-py/tests/test_chat_history.py b/pkg-py/tests/test_chat_history.py index 5b291d507..7ff364752 100644 --- a/pkg-py/tests/test_chat_history.py +++ b/pkg-py/tests/test_chat_history.py @@ -1,7 +1,7 @@ from __future__ import annotations -from typing import Any, cast -from unittest.mock import MagicMock, patch +from typing import Any, Awaitable, Callable, cast +from unittest.mock import AsyncMock, MagicMock, patch import pytest from shiny.module import ResolvedId @@ -170,3 +170,32 @@ def test_history_config_max_store_mb_default(): def test_history_config_max_store_mb_custom(): config = HistoryOptions(max_store_mb=50.0) assert config.max_store_mb == 50.0 + + +def test_controller_starts_none(): + chat = _make_chat() + assert chat.history._controller is None + + +@pytest.mark.anyio +async def test_setup_greeting_wires_on_settled(): + chat = _make_chat() + + class _FakeController: + on_settled: "Callable[[bool], Awaitable[None]] | None" = None + + fake_controller = _FakeController() + chat.history._controller = cast(Any, fake_controller) + + with patch( + "shinychat._chat_client.resolve_greeting", new=AsyncMock() + ) as mock_resolve: + chat.history.setup_greeting("## Hi") + assert fake_controller.on_settled is not None + + await fake_controller.on_settled(False) + mock_resolve.assert_awaited_once_with(chat, "## Hi") + + mock_resolve.reset_mock() + await fake_controller.on_settled(True) + mock_resolve.assert_not_awaited() diff --git a/pkg-py/tests/test_history_client.py b/pkg-py/tests/test_history_client.py index 0a14becca..af9dc623b 100644 --- a/pkg-py/tests/test_history_client.py +++ b/pkg-py/tests/test_history_client.py @@ -1,6 +1,7 @@ from __future__ import annotations import pytest +from htmltools import tags from shinychat._history_client import ( TurnsAdapter, as_turns_adapter, @@ -44,6 +45,22 @@ def test_chatlas_adapter_round_trip(): assert [t.role for t in client.get_turns()] == ["user", "assistant"] +def test_chatlas_adapter_serializes_dict_tool_result_display(): + chatlas = pytest.importorskip("chatlas") + result = chatlas.ContentToolResult( + value="done", + extra={"display": {"html": tags.div("Widget output")}}, + ) + client = chatlas.ChatOpenAI(api_key="fake") + client.set_turns([chatlas.Turn(role="user", contents=[result])]) + + dumped = as_turns_adapter(client).get_turns_json() + + display = dumped[0]["contents"][0]["extra"]["display"] + assert display["html"]["html"] == "
Widget output
" + assert isinstance(result.extra["display"], dict) + + def test_client_info_for_chatlas(): chatlas = pytest.importorskip("chatlas") client = chatlas.ChatOpenAI(api_key="fake") diff --git a/pkg-py/tests/test_history_controller.py b/pkg-py/tests/test_history_controller.py index f7cf6c041..8bdf1ce3c 100644 --- a/pkg-py/tests/test_history_controller.py +++ b/pkg-py/tests/test_history_controller.py @@ -4,7 +4,7 @@ import warnings from datetime import timedelta -from typing import Any +from typing import Any, cast from unittest.mock import AsyncMock import pytest @@ -18,7 +18,12 @@ ConversationStore, InMemoryConversationStore, ) -from shinychat._history_types import ConversationRecord, new_conversation_record +from shinychat._history_types import ( + MAX_SCHEMA_VERSION, + ConversationRecord, + UnsupportedSchemaVersionError, + new_conversation_record, +) def msg(role: str) -> dict[str, object]: @@ -150,6 +155,9 @@ def test_extend_with_no_new_ui_messages_leaves_ui_none(): class _FakeChat: + def __init__(self) -> None: + self.set_greeting_calls: list[Any] = [] + def _messages_for_bookmark(self) -> list[Any]: return [] @@ -162,6 +170,9 @@ async def clear_messages(self) -> None: async def _restore_bookmark_message(self, message_dict: Any) -> None: pass + async def set_greeting(self, greeting: Any) -> None: + self.set_greeting_calls.append(greeting) + class _FakeAdapter: def get_turns_json(self) -> list[Any]: @@ -247,6 +258,52 @@ def _make_controller( return controller, resolved_store +@pytest.mark.anyio +async def test_replay_ui_clears_greeting(): + controller, _store = _make_controller() + record = new_conversation_record(title="t") + + await controller.replay_ui(record) + + fake_chat = cast(Any, controller.chat) + assert fake_chat.set_greeting_calls == [None] + + +@pytest.mark.anyio +async def test_notify_settled_calls_on_settled_hook(): + controller, _store = _make_controller() + calls: list[bool] = [] + + async def _on_settled(restored: bool) -> None: + calls.append(restored) + + controller.on_settled = _on_settled + await controller.notify_settled(True) + await controller.notify_settled(False) + + assert calls == [True, False] + + +@pytest.mark.anyio +async def test_notify_settled_no_op_when_hook_unset(): + controller, _store = _make_controller() + await controller.notify_settled(True) + + +@pytest.mark.anyio +async def test_new_chat_notifies_settled_false(): + controller, _store = _make_controller() + calls: list[bool] = [] + + async def _on_settled(restored: bool) -> None: + calls.append(restored) + + controller.on_settled = _on_settled + await controller.new_chat() + + assert calls == [False] + + @pytest.mark.anyio async def test_controller_passes_partition_to_custom_store(): store = _PartitionCaptureStore() @@ -451,6 +508,7 @@ async def test_ui_offset_unchanged_when_save_current_store_put_raises(): class _NavFakeChat(_FakeChat): def __init__(self) -> None: + super().__init__() self.actions: list[dict[str, Any]] = [] self.cleared = 0 @@ -553,6 +611,42 @@ async def test_switch_to_nonexistent_id_raises(): await controller.switch_to("does-not-exist") +class _UnsupportedSchemaVersionStore(ConversationStore): + """Custom store whose get() returns a record from a newer, unsupported + schema version -- simulates a downgrade against a store written by a + future version of shinychat.""" + + async def list(self, partition: ConversationPartition) -> list[Any]: + return [] + + async def get( + self, partition: ConversationPartition, conv_id: str + ) -> ConversationRecord | None: + rec = new_conversation_record(title="from the future") + rec.id = conv_id + rec.schema_version = MAX_SCHEMA_VERSION + 1 + return rec + + async def put(self, partition: ConversationPartition, record: Any) -> None: + pass + + async def delete( + self, partition: ConversationPartition, conv_id: str + ) -> None: + pass + + +@pytest.mark.anyio +async def test_switch_to_rejects_record_with_unsupported_schema_version(): + store = _UnsupportedSchemaVersionStore() + controller, _store = _make_controller(store=store) + + with pytest.raises(UnsupportedSchemaVersionError): + await controller.switch_to("c_future") + + assert controller.record is None + + @pytest.mark.anyio async def test_new_chat_url_mode_sends_navigate_null(): controller, _store, chat = _make_nav_controller(with_url_mode=True) diff --git a/pkg-py/tests/test_history_store.py b/pkg-py/tests/test_history_store.py index b1dd9611b..451f49ac1 100644 --- a/pkg-py/tests/test_history_store.py +++ b/pkg-py/tests/test_history_store.py @@ -7,6 +7,8 @@ from typing import Any import pytest +from htmltools import HTMLDependency, tags +from shinychat._history_client import as_turns_adapter from shinychat._history_store import ( ConversationPartition, FileConversationStore, @@ -15,7 +17,12 @@ safe_conv_path, sanitize_scope, ) -from shinychat._history_types import ConversationRecord, new_conversation_record +from shinychat._history_types import ( + MAX_SCHEMA_VERSION, + ConversationRecord, + UnsupportedSchemaVersionError, + new_conversation_record, +) @pytest.fixture @@ -54,6 +61,40 @@ async def test_put_get_round_trip(store: FileConversationStore): assert got.nodes[nid].ui == rec.nodes[nid].ui +@pytest.mark.anyio +async def test_file_store_round_trips_dict_tool_result_display( + store: FileConversationStore, +): + chatlas = pytest.importorskip("chatlas") + + result = chatlas.ContentToolResult( + value="done", + extra={ + "display": { + "html": tags.div( + "Widget output", + HTMLDependency("my-dep", "1.0", source={"subdir": "."}), + ), + } + }, + ) + client = chatlas.ChatOpenAI(api_key="fake") + client.set_turns([chatlas.Turn(role="user", contents=[result])]) + adapter = as_turns_adapter(client) + + rec = new_conversation_record(title="Widget") + rec.append_linear(adapter.get_turns_json()) + + await store.put(part(), rec) + restored = await store.get(part(), rec.id) + + assert restored is not None + adapter.set_turns_json(restored.path_turns()) + display = adapter.get_turns_json()[0]["contents"][0]["extra"]["display"] + assert display["html"]["html"] == "
Widget output
" + assert display["html"]["dependencies"][0]["name"] == "my-dep" + + @pytest.mark.anyio async def test_put_get_round_trip_preserves_response_count( store: FileConversationStore, @@ -84,6 +125,24 @@ async def test_get_defaults_response_count_when_missing_from_disk( assert got.response_count == 0 +@pytest.mark.anyio +async def test_get_defaults_schema_version_when_missing_from_disk( + store: FileConversationStore, + tmp_path: Path, +): + rec = new_conversation_record(title="t") + await store.put(part(scope="alice"), rec) + scope_dir = tmp_path / sanitize_scope("chat") / sanitize_scope("alice") + record_file = scope_dir / rec.id / "record.json" + data = json.loads(record_file.read_text()) + del data["schema_version"] + record_file.write_text(json.dumps(data)) + + got = await store.get(part(scope="alice"), rec.id) + assert got is not None + assert got.schema_version == 1 + + @pytest.mark.anyio async def test_put_creates_directory_with_three_files( store: FileConversationStore, @@ -453,6 +512,126 @@ async def test_list_returns_independent_copy(store: FileConversationStore): assert len(store._meta_cache[part(scope="alice")]) == 1 +# --------------------------------------------------------------------------- +# schema_version rejection (issue #312) +# --------------------------------------------------------------------------- + + +def _corrupt_schema_version_on_disk( + tmp_path: Path, chat_id: str, scope: str, conv_id: str, version: int +) -> None: + scope_dir = tmp_path / sanitize_scope(chat_id) / sanitize_scope(scope) + record_file = scope_dir / conv_id / "record.json" + data = json.loads(record_file.read_text()) + data["schema_version"] = version + record_file.write_text(json.dumps(data)) + + +@pytest.mark.anyio +async def test_get_raises_on_unsupported_schema_version_on_disk( + store: FileConversationStore, + tmp_path: Path, +): + rec = new_conversation_record(title="t") + await store.put(part(scope="alice"), rec) + _corrupt_schema_version_on_disk( + tmp_path, "chat", "alice", rec.id, MAX_SCHEMA_VERSION + 1 + ) + + with pytest.raises(UnsupportedSchemaVersionError): + await store.get(part(scope="alice"), rec.id) + + +@pytest.mark.anyio +async def test_list_raises_on_unsupported_schema_version_on_disk( + store: FileConversationStore, + tmp_path: Path, +): + good = new_conversation_record(title="good") + bad = new_conversation_record(title="bad") + await store.put(part(scope="alice"), good) + await store.put(part(scope="alice"), bad) + _corrupt_schema_version_on_disk( + tmp_path, "chat", "alice", bad.id, MAX_SCHEMA_VERSION + 1 + ) + + with pytest.raises(UnsupportedSchemaVersionError): + await store.list(part(scope="alice")) + + +@pytest.mark.anyio +async def test_list_checks_schema_version_before_decoding_record( + store: FileConversationStore, + tmp_path: Path, +): + rec = new_conversation_record(title="future") + await store.put(part(scope="alice"), rec) + scope_dir = tmp_path / sanitize_scope("chat") / sanitize_scope("alice") + record_file = scope_dir / rec.id / "record.json" + data = json.loads(record_file.read_text()) + data["schema_version"] = MAX_SCHEMA_VERSION + 1 + data["nodes"] = [] + record_file.write_text(json.dumps(data)) + + with pytest.raises(UnsupportedSchemaVersionError): + await store.list(part(scope="alice")) + + +@pytest.mark.anyio +async def test_put_rejects_unsupported_schema_version_into_empty_store( + store: FileConversationStore, + tmp_path: Path, +): + rec = new_conversation_record(title="t") + rec.schema_version = MAX_SCHEMA_VERSION + 1 + + with pytest.raises(UnsupportedSchemaVersionError): + await store.put(part(scope="alice"), rec) + + scope_dir = tmp_path / sanitize_scope("chat") / sanitize_scope("alice") + assert not (scope_dir / rec.id).exists() + + +@pytest.mark.anyio +async def test_put_rejects_valid_record_over_unsupported_on_disk_record( + store: FileConversationStore, + tmp_path: Path, +): + rec = new_conversation_record(title="t") + rec.append_linear( + [{"role": "user", "content": "hi"}], + ui=[{"role": "user", "segments": []}], + ) + await store.put(part(scope="alice"), rec) + _corrupt_schema_version_on_disk( + tmp_path, "chat", "alice", rec.id, MAX_SCHEMA_VERSION + 1 + ) + + scope_dir = tmp_path / sanitize_scope("chat") / sanitize_scope("alice") + conv_dir = scope_dir / rec.id + before = {f.name: f.read_bytes() for f in conv_dir.iterdir() if f.is_file()} + + rec.append_linear([{"role": "assistant", "content": "hello"}]) + with pytest.raises(UnsupportedSchemaVersionError): + await store.put(part(scope="alice"), rec) + + after = {f.name: f.read_bytes() for f in conv_dir.iterdir() if f.is_file()} + assert after == before + + +@pytest.mark.anyio +async def test_memory_put_rejects_unsupported_schema_version( + mem_store: InMemoryConversationStore, +): + rec = new_conversation_record(title="t") + rec.schema_version = MAX_SCHEMA_VERSION + 1 + + with pytest.raises(UnsupportedSchemaVersionError): + await mem_store.put(part(scope="alice"), rec) + + assert await mem_store.get(part(scope="alice"), rec.id) is None + + # --------------------------------------------------------------------------- # InMemoryConversationStore # --------------------------------------------------------------------------- diff --git a/pkg-py/tests/test_history_types.py b/pkg-py/tests/test_history_types.py index 9d35eb862..e7ae80a24 100644 --- a/pkg-py/tests/test_history_types.py +++ b/pkg-py/tests/test_history_types.py @@ -3,6 +3,8 @@ ConversationMeta, ConversationNode, ConversationRecord, + UnsupportedSchemaVersionError, + check_schema_version, new_conversation_record, ) @@ -23,6 +25,12 @@ def test_new_record_is_empty_draft(): assert rec.response_count == 0 +@pytest.mark.parametrize("version", [True, 1.0, "1", [], [1], float("nan")]) +def test_check_schema_version_rejects_non_integer_values(version: object): + with pytest.raises(UnsupportedSchemaVersionError): + check_schema_version(version) + + def test_append_linear_builds_chain(): rec = new_conversation_record(title="t") n1 = rec.append_linear(turn("user", "hi"))