diff --git a/studio/backend/core/research_runs.py b/studio/backend/core/research_runs.py index d61663721f..15e451f6a3 100644 --- a/studio/backend/core/research_runs.py +++ b/studio/backend/core/research_runs.py @@ -17,6 +17,7 @@ from typing import Any, AsyncIterator import httpx from auth import storage as auth_storage +from core.inference.message_content import content_to_text from core.inference.tool_loop_controller import is_tool_error, strip_result_for_model from core.inference.tools import RAG_SOURCES_SENTINEL, execute_tool from core.inference.web_access_policy import check_url_access, website_policy_prompt @@ -151,16 +152,7 @@ def _safe_error(exc: BaseException) -> str: def _extract_text(message: dict) -> str: - content = message.get("content") - if isinstance(content, str): - return content - if isinstance(content, list): - return "\n".join( - str(part.get("text") or "") - for part in content - if isinstance(part, dict) and part.get("type") == "text" - ).strip() - return "" + return content_to_text(message.get("content")).strip() def _research_question_context(thread_id: str, user_message_id: str) -> tuple[str, str]: diff --git a/studio/backend/routes/research_runs.py b/studio/backend/routes/research_runs.py index 837ae7a370..da10a9492e 100644 --- a/studio/backend/routes/research_runs.py +++ b/studio/backend/routes/research_runs.py @@ -16,6 +16,7 @@ from fastapi.responses import StreamingResponse from pydantic import AliasChoices, BaseModel, ConfigDict, Field from auth.authentication import get_current_subject +from core.inference.message_content import content_to_text from core.inference.web_access_policy import normalize_website_policy from storage import research_runs_db as db from storage.studio_db import get_chat_message, get_chat_thread, upsert_chat_message @@ -239,6 +240,11 @@ async def create_research_run( raise HTTPException( status_code = 400, detail = "userMessageId must identify a user message in the thread" ) + if not content_to_text(user_message.get("content")).strip(): + raise HTTPException( + status_code = 400, + detail = "Deep research requires a user message with non-empty text", + ) if db.has_thread_claim(payload.threadId): raise HTTPException( status_code = 409, diff --git a/studio/backend/tests/test_research_runs_storage.py b/studio/backend/tests/test_research_runs_storage.py index 4a3c981f08..9613400116 100644 --- a/studio/backend/tests/test_research_runs_storage.py +++ b/studio/backend/tests/test_research_runs_storage.py @@ -1144,6 +1144,87 @@ def test_create_without_assistant_id_does_not_eagerly_create_message(research_ho assert studio_db.list_chat_messages("thread-1") == before +@pytest.mark.parametrize( + ("content", "attachments"), + [ + ([{"type": "text", "text": " \n\t"}], None), + ( + [{"type": "file", "filename": "notes.pdf"}], + [{"name": "notes.pdf", "contentType": "application/pdf"}], + ), + ], +) +def test_route_rejects_textless_research_before_claim(research_home, content, attachments): + from fastapi import HTTPException + from routes.research_runs import CreateResearchRun, create_research_run + + studio_db.upsert_chat_message( + { + "id": "user-1", + "threadId": "thread-1", + "role": "user", + "content": content, + "attachments": attachments, + "createdAt": 2, + } + ) + request = SimpleNamespace(app = SimpleNamespace(state = SimpleNamespace())) + + with pytest.raises(HTTPException, match = "non-empty text") as caught: + asyncio.run( + create_research_run( + CreateResearchRun( + threadId = "thread-1", + userMessageId = "user-1", + inferenceRequest = {"model": "local-model"}, + ), + request, + current_subject = "alice", + ) + ) + + assert caught.value.status_code == 400 + assert research_db.has_thread_claim("thread-1") is False + assert research_db.get_run("run-1") is None + + +@pytest.mark.parametrize( + "content", + [ + ["Research this question"], + [{"text": "Research this question"}], + ], +) +def test_route_accepts_canonical_text_content_shapes(research_home, content): + from core import research_runs as worker + from routes.research_runs import CreateResearchRun, create_research_run + + studio_db.upsert_chat_message( + { + "id": "user-1", + "threadId": "thread-1", + "role": "user", + "content": content, + "createdAt": 2, + } + ) + run = asyncio.run( + create_research_run( + CreateResearchRun( + threadId = "thread-1", + userMessageId = "user-1", + inferenceRequest = {"model": "local-model"}, + ), + SimpleNamespace(app = SimpleNamespace(state = SimpleNamespace())), + current_subject = "alice", + ) + ) + + assert run["status"] == "planning" + assert research_db.has_thread_claim("thread-1") is True + assert worker._extract_text({"content": content}) == "Research this question" + + def test_route_rejects_overlapping_active_run_for_thread(research_home): from fastapi import HTTPException from routes.research_runs import CreateResearchRun, create_research_run diff --git a/studio/frontend/src/components/assistant-ui/thread.tsx b/studio/frontend/src/components/assistant-ui/thread.tsx index 5075d81ec5..612cf1c1c1 100644 --- a/studio/frontend/src/components/assistant-ui/thread.tsx +++ b/studio/frontend/src/components/assistant-ui/thread.tsx @@ -3614,20 +3614,23 @@ const ComposerRightControls: FC<{ }; const MessageError: FC = () => { + const research = useResearchMessageState(); return ( {/* Recovery path for interrupted/failed turns: regenerate in place. */} - - - + {!research.runId && ( + + + + )} ); @@ -3975,10 +3978,50 @@ const ForkMessageButton: FC = () => { ); }; +const TERMINAL_RESEARCH_MESSAGE_STATUSES = new Set([ + "cancelled", + "completed", + "failed", +]); + +const useResearchMessageState = () => { + const metadata = useAuiState(({ message }) => + ( + message.metadata as + | { + custom?: { + researchRunId?: unknown; + researchStatus?: unknown; + researchRun?: { id?: unknown; status?: unknown }; + }; + } + | undefined + )?.custom, + ); + const metadataRunId = metadata?.researchRunId ?? metadata?.researchRun?.id; + const runId = typeof metadataRunId === "string" ? metadataRunId : null; + const followedStatus = useResearchRunStore((state) => + runId ? state.sessions[runId]?.run.status : undefined, + ); + const metadataStatus = metadata?.researchStatus ?? metadata?.researchRun?.status; + return { + runId, + status: + followedStatus ?? + (typeof metadataStatus === "string" ? metadataStatus : null), + }; +}; + const DeleteMessageButton: FC = () => { const aui = useAui(); const messageId = useAuiState(({ message }) => message.id); const isRunning = useAuiState(({ thread }) => thread.isRunning); + const research = useResearchMessageState(); + const isActiveResearchMessage = Boolean( + research.runId && + (!research.status || + !TERMINAL_RESEARCH_MESSAGE_STATUSES.has(research.status)), + ); const handleDelete = async () => { const thread = aui.thread(); @@ -4023,6 +4066,10 @@ const DeleteMessageButton: FC = () => { } }; + if (isActiveResearchMessage) { + return null; + } + return ( { const EditAssistantMessageButton: FC = () => { const messageId = useAuiState(({ message }) => message.id); - const isResearchMessage = useAuiState(({ message }) => { - const custom = ( - message.metadata as - | { custom?: { researchRunId?: unknown } } - | undefined - )?.custom; - return typeof custom?.researchRunId === "string"; - }); + const research = useResearchMessageState(); const isRunning = useAuiState(({ thread }) => thread.isRunning); const setEditingId = useChatRuntimeStore((s) => s.setEditingMessageId); - if (isResearchMessage) return null; + if (research.runId) return null; return ( { const AssistantActionBar: FC = () => { const { forkMessage, forkDisabled } = useForkMessageAction(); + const research = useResearchMessageState(); const [detailsOpen, setDetailsOpen] = useState(false); const ttsEnabled = useVoiceSettingsStore((s) => s.ttsEnabled); // hideWhenRunning is thread-level, so a new run would hide this bar and its @@ -4115,11 +4156,13 @@ const AssistantActionBar: FC = () => { > - - - - - + {!research.runId && ( + + + + + + )} {ttsEnabled && ( diff --git a/tests/studio/test_deep_research_frontend_contract.py b/tests/studio/test_deep_research_frontend_contract.py index 8adec878a5..7786bfe2b9 100644 --- a/tests/studio/test_deep_research_frontend_contract.py +++ b/tests/studio/test_deep_research_frontend_contract.py @@ -94,7 +94,13 @@ def test_research_presentation_is_integrated() -> None: research_gate = thread.split("const researchDisabled =", 1)[1].split(";", 1)[0] assert "!modelLoaded" not in research_gate assert "" in thread - assert "if (isResearchMessage) return null" in thread + assert "if (research.runId) return null" in thread + assert "!research.runId &&" in thread + assert "if (isActiveResearchMessage)" in thread + message_error = thread.split("const MessageError: FC = () =>", 1)[1].split( + "const GeneratingIndicator:", 1 + )[0] + assert "!research.runId &&" in message_error assert "ResearchActivityPanel" in page assert "ResearchActivitySheet" in page assert "ResearchActivityPanel" in chat_index