From f14eb56402db730af99fd1ee56714e7201ad7877 Mon Sep 17 00:00:00 2001 From: alkinun Date: Sat, 18 Jul 2026 16:09:11 +0300 Subject: [PATCH] Studio: harden research sources and limits --- studio/backend/core/inference/tools.py | 69 +++++++- studio/backend/core/research_runs.py | 129 ++++++++++++-- studio/backend/routes/research_runs.py | 4 +- studio/backend/storage/research_runs_db.py | 96 ++++++++++- studio/backend/storage/studio_db.py | 69 +++++++- studio/backend/tests/test_rag_retrieval.py | 51 ++++++ .../tests/test_research_runs_storage.py | 158 +++++++++++++++++- .../components/assistant-ui/markdown-text.tsx | 23 +-- .../components/assistant-ui/rag-sources.tsx | 38 +++-- .../components/markdown/markdown-preview.tsx | 2 + .../chat/components/research-message.tsx | 28 +++- .../src/features/chat/types/research.ts | 7 + studio/frontend/src/lib/safe-markdown-url.ts | 33 ++++ .../test_deep_research_frontend_contract.py | 5 + 14 files changed, 640 insertions(+), 72 deletions(-) create mode 100644 studio/frontend/src/lib/safe-markdown-url.ts diff --git a/studio/backend/core/inference/tools.py b/studio/backend/core/inference/tools.py index 39d0645f03..7f02392731 100644 --- a/studio/backend/core/inference/tools.py +++ b/studio/backend/core/inference/tools.py @@ -48,6 +48,7 @@ from loggers import get_logger logger = get_logger(__name__) _EXEC_TIMEOUT = 300 # 5 minutes +_RAG_SEARCH_SLOT = threading.BoundedSemaphore(1) # Splits the UI source-map from the result; loops strip it (like __IMAGES__). RAG_SOURCES_SENTINEL = "\n__RAG_SOURCES__:" @@ -3211,7 +3212,12 @@ def execute_tool( logger.info(f"execute_tool: name={name}, session_id={session_id}, timeout={timeout}") effective_timeout = _EXEC_TIMEOUT if timeout is _TIMEOUT_UNSET else timeout if name == "search_knowledge_base": - return _search_knowledge_base(arguments, rag_scope) + return _search_knowledge_base_with_budget( + arguments, + rag_scope, + effective_timeout, + cancel_event, + ) if name == "render_html": return _render_html_result(arguments) if name.startswith(MCP_TOOL_PREFIX): @@ -3337,6 +3343,65 @@ def _search_knowledge_base(arguments: dict, rag_scope: dict | None) -> str: return text +def _search_knowledge_base_with_budget( + arguments: dict, + rag_scope: dict | None, + timeout: int | None, + cancel_event = None, +) -> str: + if cancel_event is not None and cancel_event.is_set(): + return "Error: knowledge base search cancelled." + deadline = time.monotonic() + timeout if timeout is not None else None + while not _RAG_SEARCH_SLOT.acquire(timeout = 0.05): + if cancel_event is not None and cancel_event.is_set(): + return "Error: knowledge base search cancelled." + if deadline is not None and time.monotonic() >= deadline: + return "Error: knowledge base search timed out." + if cancel_event is not None and cancel_event.is_set(): + _RAG_SEARCH_SLOT.release() + return "Error: knowledge base search cancelled." + if deadline is not None and time.monotonic() >= deadline: + _RAG_SEARCH_SLOT.release() + return "Error: knowledge base search timed out." + + if timeout is None and cancel_event is None: + try: + return _search_knowledge_base(arguments, rag_scope) + finally: + _RAG_SEARCH_SLOT.release() + + result: queue.Queue = queue.Queue(maxsize = 1) + + def search() -> None: + try: + result.put((True, _search_knowledge_base(arguments, rag_scope))) + except BaseException as exc: + result.put((False, exc)) + finally: + _RAG_SEARCH_SLOT.release() + + try: + threading.Thread(target = search, name = "rag-tool-search", daemon = True).start() + except Exception: + _RAG_SEARCH_SLOT.release() + raise + while True: + if cancel_event is not None and cancel_event.is_set(): + return "Error: knowledge base search cancelled." + if deadline is not None and time.monotonic() >= deadline: + return "Error: knowledge base search timed out." + wait = 0.05 + if deadline is not None: + wait = min(wait, max(0.001, deadline - time.monotonic())) + try: + ok, value = result.get(timeout = wait) + except queue.Empty: + continue + if ok: + return value + raise value + + # Forced first-pass RAG retrieval: a high cosine floor keeps it precise (fires on # on-topic queries, skips weak ones) and helps small models that under-call the tool. # Tunable via RAG_AUTOINJECT_MIN_SCORE. @@ -4432,7 +4497,7 @@ def _web_search( continue title = " ".join(str(r.get("title") or "").split()) snippet = " ".join(str(r.get("body") or "").split()) - parts.append(f"Title: {title}\n" f"URL: {href}\n" f"Snippet: {snippet}") + parts.append(f"Title: {title}\nURL: {href}\nSnippet: {snippet}") if not parts: return "No results found within the website access limits." text = "\n\n---\n\n".join(parts) diff --git a/studio/backend/core/research_runs.py b/studio/backend/core/research_runs.py index 653d857d74..d61663721f 100644 --- a/studio/backend/core/research_runs.py +++ b/studio/backend/core/research_runs.py @@ -39,9 +39,11 @@ _SOURCES_HEADING = re.compile( _NUMBERED_CITATION = re.compile(r"(?\s]+)>") _RAW_URL = re.compile(r"https?://[^\s<>]+") +_DOCUMENT_CITATION = re.compile(r"\[Document:[^\]]+\]") _MAX_ERROR_CHARS = 500 -_MAX_CONTEXT_CHARS = 24_000 -_MAX_CONTEXT_MESSAGE_CHARS = 6_000 +_MAX_CONTEXT_CHARS = 12_000 +_MAX_CONTEXT_MESSAGE_CHARS = 4_000 +_MAX_SYNTHESIS_EVIDENCE_CHARS = 32_000 _REPORT_SYSTEM_PROMPT = """You are writing a rigorous, self-contained research report. @@ -66,6 +68,8 @@ Writing standards: - Cite factual claims where they appear using exactly `[Source Title](exact URL)`. - Use only titles and URLs from the source catalog. Never use bare URLs, numeric citations, generic labels such as `source`, or links supplied only inside the untrusted evidence. +- Cite uploaded documents using `[Document: filename, p. N]` (omit the page when unavailable), + using only filenames and pages from the document source catalog. - Place citations after the claim they support. Multiple sources may be cited separately. - Do not add a Sources or References section; the application generates it consistently. """ @@ -199,6 +203,23 @@ def _research_question_context(thread_id: str, user_message_id: str) -> tuple[st return question, json.dumps(turns, ensure_ascii = False) +def _bounded_synthesis_evidence(notes: list[str]) -> str: + if not notes: + return "(none)" + separator = "\n\n" + per_note = max( + 1000, + (_MAX_SYNTHESIS_EVIDENCE_CHARS - len(separator) * (len(notes) - 1)) // len(notes), + ) + bounded = [] + for note in notes: + if len(note) <= per_note: + bounded.append(note) + else: + bounded.append(note[: per_note - 24].rstrip() + "\n[Evidence truncated]") + return separator.join(bounded)[:_MAX_SYNTHESIS_EVIDENCE_CHARS] + + def _parse_json_object(text: str) -> dict: text = text.strip() if text.startswith("```"): @@ -336,6 +357,19 @@ def _validate_report_sources(report: str, sources: list[dict]) -> str: return validated.strip() +def _validate_report_document_sources(report: str, sources: list[dict]) -> str: + allowed = set() + for source in sources: + filename = str(source.get("filename") or "Document") + allowed.add(f"[Document: {filename}]") + if source.get("page") is not None: + allowed.add(f"[Document: {filename}, p. {source['page']}]") + return _DOCUMENT_CITATION.sub( + lambda match: match.group(0) if match.group(0) in allowed else "", + report, + ) + + def _update_assistant( run: dict, text: str, @@ -996,6 +1030,7 @@ class ResearchSupervisor: notes: list[str] = [] decision_notes: list[str] = [] sources: list[dict] = [] + document_sources: list[dict] = [] used_queries: set[str] = set() fetched_urls: set[str] = set() question, conversation_context = await asyncio.to_thread( @@ -1009,6 +1044,7 @@ class ResearchSupervisor: raise LeaseLost() if resuming: sources = list(run.get("sources") or [])[:max_sources] + document_sources = list(run.get("documentSources") or [])[:max_sources] for step in run.get("steps") or []: result = step.get("result") if isinstance(step.get("result"), dict) else {} @@ -1031,10 +1067,37 @@ class ResearchSupervisor: f"Snippet: {source.get('snippet') or ''}" for source in step_sources ) + restored_rag_sources = [ + item for item in result.get("evidenceSources") or [] if isinstance(item, dict) + ] + document_source_keys = { + str( + source.get("chunkId") + or f"{source.get('documentId') or source.get('filename')}:{source.get('page') or ''}" + ) + for source in document_sources + } + for source in restored_rag_sources: + source_key = str( + source.get("chunkId") + or f"{source.get('documentId') or source.get('filename')}:{source.get('page') or ''}" + ) + if source_key in document_source_keys or len(document_sources) >= max_sources: + continue + written = await asyncio.to_thread( + db.upsert_document_source, + run["id"], + int(step["position"]), + source, + self.worker_id, + ) + await self._check_worker_write(run["id"], written) + document_source_keys.add(source_key) + document_sources.append({**source, "stepPosition": step["position"]}) rag_evidence = "\n".join( - f"{item.get('filename') or 'Document'}: {item.get('snippet') or ''}" - for item in result.get("evidenceSources") or [] - if isinstance(item, dict) + f"{item.get('filename') or 'Document'}: " + f"{item.get('text') or item.get('snippet') or ''}" + for item in restored_rag_sources ) title = str(step.get("title") or "Recovered research step") notes.append( @@ -1184,6 +1247,41 @@ class ResearchSupervisor: ) rag_result, rag_sources = _split_rag_result(rag_result) await self._check_active(run["id"]) + document_source_keys = { + str( + source.get("chunkId") + or f"{source.get('documentId') or source.get('filename')}:{source.get('page') or ''}" + ) + for source in document_sources + } + accepted_rag_sources = [] + for source in rag_sources: + source_key = str( + source.get("chunkId") + or f"{source.get('documentId') or source.get('filename')}:{source.get('page') or ''}" + ) + if source_key not in document_source_keys: + if len(document_sources) >= max_sources: + continue + written = await asyncio.to_thread( + db.upsert_document_source, + run["id"], + position, + source, + self.worker_id, + ) + await self._check_worker_write(run["id"], written) + document_source_keys.add(source_key) + document_sources.append({**source, "stepPosition": position}) + accepted_rag_sources.append(source) + if accepted_rag_sources: + rag_result = "\n\n".join( + f"Document: {source.get('filename') or 'Document'}" + f"{', page ' + str(source.get('page')) if source.get('page') is not None else ''}\n" + f"{source.get('text') or source.get('snippet') or ''}" + for source in accepted_rag_sources + ) + rag_sources = accepted_rag_sources step_sources = [] for match in _URL_BLOCK.finditer(result if action["action"] == "search" else ""): if len(sources) >= max_sources: @@ -1225,7 +1323,7 @@ class ResearchSupervisor: step_result = { "action": action["action"], "input": argument, - "sourceCount": len(step_sources), + "sourceCount": len(step_sources) + len(rag_sources), "sourceUrls": [source["url"] for source in step_sources], "evidenceSources": rag_sources, **({"excerpt": clean_result[:2000]} if action["action"] == "fetch" else {}), @@ -1254,19 +1352,24 @@ class ResearchSupervisor: "title": action["title"], "action": action["action"], "input": argument, - "sourceCount": len(step_sources), + "sourceCount": len(step_sources) + len(rag_sources), **({"error": clean_result[:500]} if tool_failed else {}), }, ) await self._check_worker_write(run["id"], seq is not None) await self._check_active(run["id"]) source_catalog = "\n".join( - f"{index}. Title: {source.get('title') or source['url']}\n" - f" URL: {source['url']}\n" - f" Search snippet: {source.get('snippet') or '(none)'}" + f"{index}. Title: {source.get('title') or source['url']}\n URL: {source['url']}" for index, source in enumerate(sources, 1) ) - evidence_text = "\n\n".join(notes) + document_source_catalog = "\n".join( + f"{index}. Filename: {source.get('filename') or 'Document'}\n" + f" Page: {source.get('page') if source.get('page') is not None else '(unknown)'}\n" + f" Document ID: {source.get('documentId') or '(unknown)'}\n" + f" Chunk ID: {source.get('chunkId') or '(unknown)'}" + for index, source in enumerate(document_sources, 1) + ) + evidence_text = _bounded_synthesis_evidence(notes) report, synthesis_reasoning, synthesis_finish_reason = await self._stream_completion( run, [ @@ -1282,6 +1385,9 @@ class ResearchSupervisor: f"\n\n" f"\n{source_catalog or '(no web sources gathered)'}\n" f"\n\n" + f"\n" + f"{document_source_catalog or '(no document sources gathered)'}\n" + f"\n\n" f"\n{evidence_text}\n" f"" ), @@ -1298,6 +1404,7 @@ class ResearchSupervisor: if not report: raise ValueError("Local model returned an empty report") report = _validate_report_sources(report, sources) + report = _validate_report_document_sources(report, document_sources) reasoning = await asyncio.to_thread(db.get_reasoning_text, run["id"]) if synthesis_reasoning and synthesis_reasoning not in reasoning: reasoning += synthesis_reasoning diff --git a/studio/backend/routes/research_runs.py b/studio/backend/routes/research_runs.py index b3a8d0f2a6..479b634544 100644 --- a/studio/backend/routes/research_runs.py +++ b/studio/backend/routes/research_runs.py @@ -239,7 +239,7 @@ async def create_research_run( raise HTTPException( status_code = 400, detail = "userMessageId must identify a user message in the thread" ) - if db.has_thread_claim(current_subject, payload.threadId): + if db.has_thread_claim(payload.threadId): raise HTTPException( status_code = 409, detail = "This thread already has a Deep Research run", @@ -271,7 +271,7 @@ async def active_research_runs( ): return { "runs": db.list_active(current_subject, thread_id), - "hasRun": db.has_thread_claim(current_subject, thread_id), + "hasRun": db.has_thread_claim(thread_id), } diff --git a/studio/backend/storage/research_runs_db.py b/studio/backend/storage/research_runs_db.py index 0700019e8f..fcdf8a1d95 100644 --- a/studio/backend/storage/research_runs_db.py +++ b/studio/backend/storage/research_runs_db.py @@ -151,8 +151,8 @@ def create_run( ) except sqlite3.IntegrityError as exc: claim = conn.execute( - "SELECT 1 FROM research_thread_claims WHERE owner_subject=? AND thread_id=?", - (owner_subject, thread_id), + "SELECT 1 FROM research_thread_claims WHERE thread_id=?", + (thread_id,), ).fetchone() if claim is not None: raise ResearchConflictError("This thread already has a Deep Research run") from exc @@ -295,6 +295,16 @@ def get_run(run_id: str, owner_subject: str | None = None) -> dict | None: (run_id,), ).fetchall() ] + result["documentSources"] = [ + dict(r) + for r in conn.execute( + "SELECT id, step_position AS stepPosition, document_id AS documentId, " + "chunk_id AS chunkId, filename, page, score, snippet, " + "fetched_at AS fetchedAt FROM research_document_sources " + "WHERE run_id = ? ORDER BY id", + (run_id,), + ).fetchall() + ] return result finally: conn.close() @@ -314,13 +324,13 @@ def list_active(owner_subject: str, thread_id: str) -> list[dict]: return [run for row in rows if (run := get_run(row["id"], owner_subject)) is not None] -def has_thread_claim(owner_subject: str, thread_id: str) -> bool: +def has_thread_claim(thread_id: str) -> bool: conn = get_connection() try: return ( conn.execute( - "SELECT 1 FROM research_thread_claims WHERE owner_subject=? AND thread_id=?", - (owner_subject, thread_id), + "SELECT 1 FROM research_thread_claims WHERE thread_id=?", + (thread_id,), ).fetchone() is not None ) @@ -607,6 +617,12 @@ def retry(run_id: str, max_retries: int = 3) -> str: raise ResearchConflictError("Only failed or cancelled runs can be retried") if int(row["retry_count"]) >= max_retries: raise ResearchConflictError("Retry budget exhausted") + claim = conn.execute( + "SELECT owner_subject FROM research_thread_claims WHERE thread_id=?", + (row["thread_id"],), + ).fetchone() + if claim is None or claim["owner_subject"] != row["owner_subject"]: + raise ResearchConflictError("This run does not own the thread research claim") placeholders = ",".join("?" for _ in ACTIVE_STATUSES) active = conn.execute( f"SELECT id FROM research_runs WHERE owner_subject=? AND thread_id=? AND id<>? " @@ -640,6 +656,7 @@ def retry(run_id: str, max_retries: int = 3) -> str: if status != "awaiting_approval": conn.execute("DELETE FROM research_plan_steps WHERE run_id = ?", (run_id,)) conn.execute("DELETE FROM research_sources WHERE run_id = ?", (run_id,)) + conn.execute("DELETE FROM research_document_sources WHERE run_id = ?", (run_id,)) _event_locked(conn, run_id, "run.retried", {"status": status}) _commit_event(conn) return status @@ -656,10 +673,12 @@ def claim_next(worker_id: str, lease_ms: int = 120_000) -> dict | None: conn.execute("BEGIN IMMEDIATE") now = now_ms() row = conn.execute( - """SELECT * FROM research_runs - WHERE status IN ('planning','queued','running','cancelling') - AND (lease_owner IS NULL OR lease_expires_at < ?) - ORDER BY created_at LIMIT 1""", + """SELECT r.* FROM research_runs r + JOIN research_thread_claims c ON c.thread_id=r.thread_id + WHERE r.owner_subject=c.owner_subject + AND r.status IN ('planning','queued','running','cancelling') + AND (r.lease_owner IS NULL OR r.lease_expires_at < ?) + ORDER BY r.created_at LIMIT 1""", (now,), ).fetchone() if row is None: @@ -875,6 +894,7 @@ def reset_execution_steps(run_id: str, worker_id: str | None = None) -> bool: return False conn.execute("DELETE FROM research_plan_steps WHERE run_id = ?", (run_id,)) conn.execute("DELETE FROM research_sources WHERE run_id = ?", (run_id,)) + conn.execute("DELETE FROM research_document_sources WHERE run_id = ?", (run_id,)) conn.commit() return True except Exception: @@ -901,6 +921,10 @@ def prepare_execution_resume(run_id: str, worker_id: str) -> bool: "DELETE FROM research_sources WHERE run_id = ? AND step_position = ?", [(run_id, int(row["position"])) for row in interrupted], ) + conn.executemany( + "DELETE FROM research_document_sources WHERE run_id = ? AND step_position = ?", + [(run_id, int(row["position"])) for row in interrupted], + ) conn.execute( "DELETE FROM research_plan_steps WHERE run_id = ? " "AND status NOT IN ('completed','failed')", @@ -1050,6 +1074,60 @@ def upsert_source( conn.close() +def upsert_document_source( + run_id: str, + position: int, + source: dict[str, Any], + worker_id: str | None = None, +) -> bool: + filename = str(source.get("filename") or "Document")[:500] + document_id = source.get("documentId") + chunk_id = source.get("chunkId") + page = source.get("page") + source_key = str(chunk_id or f"{document_id or filename}:{page or ''}")[:1000] + conn = get_connection() + try: + conn.execute("BEGIN IMMEDIATE") + if worker_id is not None and not _worker_can_write_locked( + conn, + run_id, + worker_id, + {"running"}, + ): + conn.commit() + return False + fetched_at = now_ms() + conn.execute( + """INSERT INTO research_document_sources + (run_id, step_position, source_key, document_id, chunk_id, filename, + page, score, snippet, fetched_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(run_id, source_key) DO UPDATE SET + step_position=excluded.step_position, document_id=excluded.document_id, + chunk_id=excluded.chunk_id, filename=excluded.filename, page=excluded.page, + score=excluded.score, snippet=excluded.snippet, fetched_at=excluded.fetched_at""", + ( + run_id, + position, + source_key, + str(document_id)[:500] if document_id is not None else None, + str(chunk_id)[:500] if chunk_id is not None else None, + filename, + int(page) if isinstance(page, (int, float)) else None, + float(source["score"]) if isinstance(source.get("score"), (int, float)) else None, + str(source.get("text") or source.get("snippet") or "")[:4000], + fetched_at, + ), + ) + conn.commit() + return True + except Exception: + conn.rollback() + raise + finally: + conn.close() + + def list_events( run_id: str, owner_subject: str, diff --git a/studio/backend/storage/studio_db.py b/studio/backend/storage/studio_db.py index fcadb49b4f..d19270e23f 100644 --- a/studio/backend/storage/studio_db.py +++ b/studio/backend/storage/studio_db.py @@ -431,17 +431,54 @@ def _ensure_schema(conn: sqlite3.Connection) -> None: """ CREATE TABLE IF NOT EXISTS research_thread_claims ( owner_subject TEXT NOT NULL, - thread_id TEXT NOT NULL REFERENCES chat_threads(id) ON DELETE CASCADE, - created_at INTEGER NOT NULL, - PRIMARY KEY(owner_subject, thread_id) + thread_id TEXT NOT NULL PRIMARY KEY REFERENCES chat_threads(id) ON DELETE CASCADE, + created_at INTEGER NOT NULL ) WITHOUT ROWID """ ) + claim_pk = [ + row[1] + for row in sorted( + conn.execute("PRAGMA table_info(research_thread_claims)").fetchall(), + key = lambda row: int(row[5] or 0), + ) + if int(row[5] or 0) > 0 + ] + if claim_pk != ["thread_id"]: + conn.execute("ALTER TABLE research_thread_claims RENAME TO research_thread_claims_legacy") + conn.execute( + """ + CREATE TABLE research_thread_claims ( + owner_subject TEXT NOT NULL, + thread_id TEXT NOT NULL PRIMARY KEY REFERENCES chat_threads(id) ON DELETE CASCADE, + created_at INTEGER NOT NULL + ) WITHOUT ROWID + """ + ) + conn.execute( + """INSERT OR IGNORE INTO research_thread_claims + (owner_subject, thread_id, created_at) + SELECT owner_subject, thread_id, created_at + FROM research_thread_claims_legacy + ORDER BY created_at, owner_subject""" + ) + conn.execute("DROP TABLE research_thread_claims_legacy") conn.execute( """INSERT OR IGNORE INTO research_thread_claims (owner_subject, thread_id, created_at) - SELECT owner_subject, thread_id, MIN(created_at) - FROM research_runs GROUP BY owner_subject, thread_id""" + SELECT owner_subject, thread_id, created_at + FROM research_runs ORDER BY created_at, id""" + ) + conn.execute( + """UPDATE research_runs + SET status='failed', error_message='Superseded by the global thread research claim', + lease_owner=NULL, lease_expires_at=NULL, completed_at=COALESCE(completed_at, updated_at) + WHERE status IN ('planning','awaiting_approval','queued','running','paused','cancelling') + AND EXISTS ( + SELECT 1 FROM research_thread_claims c + WHERE c.thread_id=research_runs.thread_id + AND c.owner_subject<>research_runs.owner_subject + )""" ) conn.execute( """ @@ -472,6 +509,24 @@ def _ensure_schema(conn: sqlite3.Connection) -> None: ) """ ) + conn.execute( + """ + CREATE TABLE IF NOT EXISTS research_document_sources ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + run_id TEXT NOT NULL REFERENCES research_runs(id) ON DELETE CASCADE, + step_position INTEGER, + source_key TEXT NOT NULL, + document_id TEXT, + chunk_id TEXT, + filename TEXT NOT NULL, + page INTEGER, + score REAL, + snippet TEXT, + fetched_at INTEGER NOT NULL, + UNIQUE(run_id, source_key) + ) + """ + ) conn.execute( """ CREATE TABLE IF NOT EXISTS research_events ( @@ -495,6 +550,10 @@ def _ensure_schema(conn: sqlite3.Connection) -> None: conn.execute( "CREATE INDEX IF NOT EXISTS idx_research_sources_run ON research_sources(run_id, id)" ) + conn.execute( + "CREATE INDEX IF NOT EXISTS idx_research_document_sources_run " + "ON research_document_sources(run_id, id)" + ) def _prompt_entry_from_row(row: sqlite3.Row) -> dict: diff --git a/studio/backend/tests/test_rag_retrieval.py b/studio/backend/tests/test_rag_retrieval.py index 69d9e90871..e630711c10 100644 --- a/studio/backend/tests/test_rag_retrieval.py +++ b/studio/backend/tests/test_rag_retrieval.py @@ -4,6 +4,8 @@ """Retrieval + tool tests: RRF fusion, min-score floor, scope, source-map.""" import math +import threading +import time import pytest @@ -192,6 +194,55 @@ def test_dispatcher_no_sentinel_when_no_hits(rag_home, monkeypatch): assert tools.RAG_SOURCES_SENTINEL not in out +def test_knowledge_search_honors_cancellation_and_timeout(monkeypatch): + from core.inference import tools + + started = threading.Event() + release = threading.Event() + calls = 0 + + def stalled_search(arguments, rag_scope): + nonlocal calls + calls += 1 + started.set() + release.wait() + return "late" + + monkeypatch.setattr(tools, "_search_knowledge_base", stalled_search) + cancel = threading.Event() + + def cancel_after_start(): + started.wait() + cancel.set() + + threading.Thread(target = cancel_after_start, daemon = True).start() + began = time.monotonic() + try: + cancelled = tools.execute_tool( + "search_knowledge_base", + {"query": "q"}, + cancel_event = cancel, + timeout = 30, + rag_scope = {"kb_id": "a"}, + ) + assert "cancelled" in cancelled.lower() + assert time.monotonic() - began < 1 + + started.clear() + timed_out = tools.execute_tool( + "search_knowledge_base", + {"query": "q"}, + timeout = 0, + rag_scope = {"kb_id": "a"}, + ) + assert "timed out" in timed_out.lower() + assert calls == 1 + finally: + release.set() + assert tools._RAG_SEARCH_SLOT.acquire(timeout = 1) + tools._RAG_SEARCH_SLOT.release() + + def test_search_for_autoinject_gates_on_dense_score(rag_conn, bow_embeddings, monkeypatch): _add_doc(rag_conn, "kb_a", "d1", "paper.pdf", "h1", "body text here", page = 3) diff --git a/studio/backend/tests/test_research_runs_storage.py b/studio/backend/tests/test_research_runs_storage.py index 289e5ec4f6..6668794a3e 100644 --- a/studio/backend/tests/test_research_runs_storage.py +++ b/studio/backend/tests/test_research_runs_storage.py @@ -136,12 +136,41 @@ def test_planner_uses_last_valid_plan_when_reasoning_contains_a_draft(): assert worker._parse_and_validate_plan("", reasoning, 5) == _plan() +def test_synthesis_evidence_is_bounded_across_all_steps(): + from core import research_runs as worker + + evidence = worker._bounded_synthesis_evidence( + [f"### Step {index}\n" + "x" * 20_000 for index in range(12)] + ) + + assert len(evidence) <= worker._MAX_SYNTHESIS_EVIDENCE_CHARS + assert all(f"### Step {index}" in evidence for index in range(12)) + + def test_report_is_recovered_from_substantial_synthesis_reasoning(): from core import research_runs as worker report = "**Executive Summary**\n\n" + ("Evidence-based conclusion. " * 30) reasoning = "I will organize the final answer.\n" + report assert worker._recover_report_from_reasoning(reasoning) == report.strip() + + +def test_document_citations_are_restricted_to_persisted_sources(): + from core import research_runs as worker + + report = ( + "Supported [Document: private.pdf, p. 2]. " + "Fabricated [Document: invented.pdf, p. 9] and " + "[Document: multiline.pdf,\np. 3]." + ) + validated = worker._validate_report_document_sources( + report, + [{"filename": "private.pdf", "page": 2}], + ) + + assert "[Document: private.pdf, p. 2]" in validated + assert "invented.pdf" not in validated + assert "multiline.pdf" not in validated assert worker._recover_report_from_reasoning("Too short") == "" assert worker._recover_report_from_reasoning("Internal analysis. " * 50) == "" assert ( @@ -288,10 +317,81 @@ def test_schema_and_state_transitions(research_home): "research_thread_claims", "research_plan_steps", "research_sources", + "research_document_sources", "research_events", } +def test_owner_scoped_claim_schema_migrates_to_global(tmp_path, monkeypatch): + monkeypatch.setenv("UNSLOTH_STUDIO_HOME", str(tmp_path)) + monkeypatch.setattr(studio_db, "_schema_ready", False) + studio_db.upsert_chat_thread( + { + "id": "shared-thread", + "title": "Shared", + "modelType": "base", + "modelId": "model", + "createdAt": 1, + } + ) + studio_db.upsert_chat_message( + { + "id": "shared-user", + "threadId": "shared-thread", + "role": "user", + "content": [{"type": "text", "text": "Question"}], + "createdAt": 2, + } + ) + conn = studio_db.get_connection() + try: + conn.execute("DROP TABLE research_thread_claims") + conn.execute( + """CREATE TABLE research_thread_claims ( + owner_subject TEXT NOT NULL, + thread_id TEXT NOT NULL REFERENCES chat_threads(id) ON DELETE CASCADE, + created_at INTEGER NOT NULL, + PRIMARY KEY(owner_subject, thread_id) + ) WITHOUT ROWID""" + ) + conn.executemany( + "INSERT INTO research_thread_claims VALUES (?, 'shared-thread', ?)", + [("bob", 20), ("alice", 10)], + ) + conn.executemany( + """INSERT INTO research_runs + (id, owner_subject, thread_id, user_message_id, status, config_json, + created_at, updated_at) + VALUES (?, ?, 'shared-thread', 'shared-user', 'queued', '{}', ?, ?)""", + [("bob-run", "bob", 20, 20), ("alice-run", "alice", 10, 10)], + ) + conn.commit() + finally: + conn.close() + + studio_db._schema_ready = False + conn = studio_db.get_connection() + try: + primary_key = [ + row["name"] + for row in conn.execute("PRAGMA table_info(research_thread_claims)").fetchall() + if row["pk"] + ] + claims = conn.execute( + "SELECT owner_subject, thread_id FROM research_thread_claims" + ).fetchall() + runs = conn.execute("SELECT id, status FROM research_runs ORDER BY id").fetchall() + finally: + conn.close() + + assert primary_key == ["thread_id"] + assert [tuple(row) for row in claims] == [("alice", "shared-thread")] + assert [tuple(row) for row in runs] == [("alice-run", "queued"), ("bob-run", "failed")] + with pytest.raises(research_db.ResearchConflictError, match = "does not own"): + research_db.retry("bob-run") + assert research_db.claim_next("migration-worker")["id"] == "alice-run" + + def test_pruning_messages_preserves_runs_whose_user_message_survives(research_home): _create() studio_db.upsert_chat_message( @@ -313,7 +413,7 @@ def test_pruning_messages_preserves_runs_whose_user_message_survives(research_ho studio_db.sync_chat_messages("thread-1", survivors, prune_missing = True) assert research_db.get_run("run-1") is not None - assert research_db.has_thread_claim("alice", "thread-1") is True + assert research_db.has_thread_claim("thread-1") is True assert studio_db.get_chat_message("thread-1", "temporary") is None @@ -482,11 +582,23 @@ def test_execution_reset_clears_steps_and_sources(research_home): "run-1", 0, "Old step", "old query", "completed", worker_id = "worker-1" ) research_db.upsert_source("run-1", 0, "https://old.example", "Old", "Stale", "worker-1") + research_db.upsert_document_source( + "run-1", + 0, + { + "documentId": "doc-old", + "chunkId": "chunk-old", + "filename": "old.pdf", + "text": "Stale private evidence", + }, + "worker-1", + ) assert research_db.reset_execution_steps("run-1", "worker-1") is True run = research_db.get_run("run-1") assert run["steps"] == [] assert run["sources"] == [] + assert run["documentSources"] == [] def test_supervisor_stop_signals_tool_cancellation_before_task_cancelled(research_home): @@ -848,6 +960,8 @@ def test_supervisor_planning_and_research_are_durable_with_mocked_io(research_ho return json.dumps(_plan()), "Planned several lines of inquiry.", "stop" if "iterative research process" in system: return next(decisions), "Evaluated the evidence and selected the next action.", "stop" + assert "" in prompt + assert "private.pdf" in prompt report = report_response research_db.set_report_progress(run["id"], report) return report, "Checked the available evidence.", "stop" @@ -857,7 +971,22 @@ def test_supervisor_planning_and_research_are_durable_with_mocked_io(research_ho def fake_tool(name, arguments, *args, **kwargs): tool_calls.append((name, kwargs)) if name == "search_knowledge_base": - return "Private evidence" + return ( + "Private evidence" + + worker.RAG_SOURCES_SENTINEL + + json.dumps( + [ + { + "chunkId": "doc-1:0", + "documentId": "doc-1", + "filename": "private.pdf", + "page": 2, + "text": "Private durable evidence", + "score": 0.9, + } + ] + ) + ) if arguments.get("url"): return "Full page evidence." return "Title: Example\nURL: https://example.com\nSnippet: Evidence snippet." @@ -882,6 +1011,8 @@ def test_supervisor_planning_and_research_are_durable_with_mocked_io(research_ho assert completed["status"] == "completed" assert completed["report"].startswith("# Final report") assert completed["sources"][0]["url"] == "https://example.com" + assert completed["documentSources"][0]["documentId"] == "doc-1" + assert completed["documentSources"][0]["filename"] == "private.pdf" assert completed["steps"][0]["query"] == "example evidence" assert completed["steps"][0]["input"] == "example evidence" assert completed["steps"][0]["result"]["input"] == "example evidence" @@ -991,6 +1122,7 @@ def test_recovered_running_research_resumes_durable_progress(research_home, monk assert completed["status"] == "completed" assert [step["position"] for step in completed["steps"]] == [0] assert [source["url"] for source in completed["sources"]] == ["https://saved.example/source"] + assert [source["filename"] for source in completed["documentSources"]] == ["private.txt"] assert completed["report"].startswith("# Resumed report") @@ -1096,7 +1228,7 @@ def test_assistant_discovery_binding_and_terminal_fallback_are_idempotent(resear def test_research_claim_lasts_for_thread_lifetime(research_home): _create() - assert research_db.has_thread_claim("alice", "thread-1") is True + assert research_db.has_thread_claim("thread-1") is True conn = studio_db.get_connection() try: @@ -1105,7 +1237,7 @@ def test_research_claim_lasts_for_thread_lifetime(research_home): finally: conn.close() assert research_db.get_run("run-1") is None - assert research_db.has_thread_claim("alice", "thread-1") is True + assert research_db.has_thread_claim("thread-1") is True studio_db.upsert_chat_message( { @@ -1124,7 +1256,23 @@ def test_research_claim_lasts_for_thread_lifetime(research_home): ) studio_db.delete_chat_threads(["thread-1"]) - assert research_db.has_thread_claim("alice", "thread-1") is False + assert research_db.has_thread_claim("thread-1") is False + + +def test_research_claim_is_global_across_authenticated_subjects(research_home): + first = _create() + + with pytest.raises(research_db.ResearchConflictError, match = "already has"): + research_db.create_run( + run_id = "run-2", + owner_subject = "bob", + thread_id = "thread-1", + user_message_id = "user-1", + assistant_message_id = None, + config = first["config"], + ) + + assert research_db.has_thread_claim("thread-1") is True def test_list_active_returns_complete_snapshots(research_home): diff --git a/studio/frontend/src/components/assistant-ui/markdown-text.tsx b/studio/frontend/src/components/assistant-ui/markdown-text.tsx index 40fc8b8da6..9722018ba4 100644 --- a/studio/frontend/src/components/assistant-ui/markdown-text.tsx +++ b/studio/frontend/src/components/assistant-ui/markdown-text.tsx @@ -14,14 +14,15 @@ import { import { copyToClipboard } from "@/lib/copy-to-clipboard"; import { preprocessLaTeX } from "@/lib/latex"; import { openLink } from "@/lib/open-link"; -import { INTERNAL, useAuiState, useMessagePartText } from "@assistant-ui/react"; +import { safeMarkdownUrl } from "@/lib/safe-markdown-url"; import { Tick02Icon } from "@/lib/tick-icon"; +import { INTERNAL, useAuiState, useMessagePartText } from "@assistant-ui/react"; import { Copy01Icon, Download01Icon } from "@hugeicons/core-free-icons"; import { HugeiconsIcon } from "@hugeicons/react"; import { createMathPlugin } from "@streamdown/math"; import { mermaid } from "@streamdown/mermaid"; import { useEffect, useMemo, useRef, useState } from "react"; -import { Block, type BlockProps, Streamdown, defaultUrlTransform, type UrlTransform } from "streamdown"; +import { Block, type BlockProps, Streamdown } from "streamdown"; import { createCodePlugin } from "./code-plugin"; import "katex/dist/katex.min.css"; import { AudioPlayer } from "./audio-player"; @@ -368,22 +369,6 @@ function useRafCoalescedText(text: string, isStreaming: boolean): string { return text; } -const safeImageUrl: UrlTransform = (url, _key, node) => { - // Only images are restricted; links/other nodes use the default transform. - if (node.tagName !== "img") return defaultUrlTransform(url, _key, node); - - // Strip ASCII controls first: browsers drop them mid-parse, so a value like - // "\t//attacker.com" would otherwise slip past the guards below. - // eslint-disable-next-line no-control-regex - const normalized = url.replace(/[\x00-\x1f\x7f]/g, "").trim(); - const lower = normalized.toLowerCase(); - - if (lower.startsWith("data:") || lower.startsWith("blob:")) return normalized; - if (/^[/\\]{2}/.test(normalized)) return null; // protocol-relative: // \\ /\ \/ - if (/^[a-zA-Z][a-zA-Z0-9+\-.]*:/.test(normalized)) return null; // scheme prefix (colon later in path is fine) - return normalized; // relative -> same-origin -}; - const MarkdownTextImpl = () => { const { text, status } = useMessagePartText(); const displayText = useRafCoalescedText(text, status.type === "running"); @@ -404,7 +389,7 @@ const MarkdownTextImpl = () => { isAnimating={status.type === "running"} plugins={{ code, math, mermaid }} components={STREAMDOWN_COMPONENTS} - urlTransform={safeImageUrl} + urlTransform={safeMarkdownUrl} controls={{ code: false, mermaid: { diff --git a/studio/frontend/src/components/assistant-ui/rag-sources.tsx b/studio/frontend/src/components/assistant-ui/rag-sources.tsx index ab7a572e52..27e26ca8e9 100644 --- a/studio/frontend/src/components/assistant-ui/rag-sources.tsx +++ b/studio/frontend/src/components/assistant-ui/rag-sources.tsx @@ -9,27 +9,26 @@ import type { FC } from "react"; import { type Citation, parseCitations } from "./citation-utils"; import { CitationBadge } from "./tool-ui-knowledge-base"; -export const RagSourcesGroup: FC = () => { - const message = useMessage(); - - const all: Citation[] = []; - for (const part of message.content ?? []) { - if (part.type === "tool-call" && part.toolName === "search_knowledge_base") { - all.push(...parseCitations(part.result)); - } - } - +export const DocumentSourcesGroup: FC<{ sources: Citation[] }> = ({ + sources: all, +}) => { // Map updates keep first-seen order, so dedup to best-scoring chunk per doc. const byDoc = new Map(); for (const c of all) { const key = c.documentId ?? c.filename; const prev = byDoc.get(key); - if (!prev || (c.score ?? -Infinity) > (prev.score ?? -Infinity)) { + if ( + !prev || + (c.score ?? Number.NEGATIVE_INFINITY) > + (prev.score ?? Number.NEGATIVE_INFINITY) + ) { byDoc.set(key, c); } } const sources = Array.from(byDoc.values()); - if (sources.length === 0) return null; + if (sources.length === 0) { + return null; + } return (
@@ -44,3 +43,18 @@ export const RagSourcesGroup: FC = () => {
); }; + +export const RagSourcesGroup: FC = () => { + const message = useMessage(); + + const sources: Citation[] = []; + for (const part of message.content ?? []) { + if ( + part.type === "tool-call" && + part.toolName === "search_knowledge_base" + ) { + sources.push(...parseCitations(part.result)); + } + } + return ; +}; diff --git a/studio/frontend/src/components/markdown/markdown-preview.tsx b/studio/frontend/src/components/markdown/markdown-preview.tsx index 34e516e74d..6421bc0129 100644 --- a/studio/frontend/src/components/markdown/markdown-preview.tsx +++ b/studio/frontend/src/components/markdown/markdown-preview.tsx @@ -2,6 +2,7 @@ // Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 import { openLink } from "@/lib/open-link"; +import { safeMarkdownUrl } from "@/lib/safe-markdown-url"; import { cn } from "@/lib/utils"; import { code } from "@streamdown/code"; import { math } from "@streamdown/math"; @@ -56,6 +57,7 @@ function MarkdownPreviewImpl({ mode="static" plugins={MARKDOWN_PLUGINS} components={MARKDOWN_COMPONENTS} + urlTransform={safeMarkdownUrl} controls={false} className={markdownClassName} > diff --git a/studio/frontend/src/features/chat/components/research-message.tsx b/studio/frontend/src/features/chat/components/research-message.tsx index 563cacd009..a2f5520637 100644 --- a/studio/frontend/src/features/chat/components/research-message.tsx +++ b/studio/frontend/src/features/chat/components/research-message.tsx @@ -1,19 +1,17 @@ // SPDX-License-Identifier: AGPL-3.0-only -import { MarkdownPreview } from "@/components/markdown/markdown-preview"; +import type { Citation } from "@/components/assistant-ui/citation-utils"; +import { DocumentSourcesGroup } from "@/components/assistant-ui/rag-sources"; import { type SourceData, SourcesGroup, } from "@/components/assistant-ui/sources"; +import { MarkdownPreview } from "@/components/markdown/markdown-preview"; import { Button } from "@/components/ui/button"; import { Spinner } from "@/components/ui/spinner"; import { cn } from "@/lib/utils"; import { useAuiState } from "@assistant-ui/react"; -import { - Check, - Telescope, - TriangleAlert, -} from "lucide-react"; +import { Check, Telescope, TriangleAlert } from "lucide-react"; import { type ReactElement, useEffect } from "react"; import { ensureResearchRunFollowed, @@ -76,6 +74,21 @@ export function ResearchMessage(): ReactElement { title: source.title || source.url, description: source.snippet ?? undefined, })); + const documentSources: Citation[] = (run.documentSources ?? []).map( + (source, index) => ({ + id: source.chunkId ?? String(source.id ?? index), + filename: source.filename, + page: source.page, + score: source.score, + text: source.snippet ?? "", + documentId: source.documentId, + chunkId: source.chunkId, + }), + ); + const documentCount = new Set( + documentSources.map((source) => source.documentId ?? source.filename), + ).size; + const sourceCount = sources.length + documentCount; return (
+
); } diff --git a/studio/frontend/src/features/chat/types/research.ts b/studio/frontend/src/features/chat/types/research.ts index d3bf68efa8..0fd42c3a14 100644 --- a/studio/frontend/src/features/chat/types/research.ts +++ b/studio/frontend/src/features/chat/types/research.ts @@ -62,6 +62,12 @@ export interface ResearchSource { fetchedAt?: number; } +export interface ResearchDocumentSource extends ResearchEvidenceSource { + id?: string | number; + stepPosition?: number | null; + fetchedAt?: number; +} + export interface ResearchInferenceRequest { model: string; temperature?: number; @@ -104,6 +110,7 @@ export interface ResearchRun { planHash: string | null; steps: ResearchStepSnapshot[]; sources: ResearchSource[]; + documentSources?: ResearchDocumentSource[]; config?: { model?: string; inferenceRequest?: Record; diff --git a/studio/frontend/src/lib/safe-markdown-url.ts b/studio/frontend/src/lib/safe-markdown-url.ts new file mode 100644 index 0000000000..6f4a175e37 --- /dev/null +++ b/studio/frontend/src/lib/safe-markdown-url.ts @@ -0,0 +1,33 @@ +import { type UrlTransform, defaultUrlTransform } from "streamdown"; + +const PROTOCOL_RELATIVE_RE = /^[/\\]{2}/; +const SCHEME_RE = /^[a-zA-Z][a-zA-Z0-9+\-.]*:/; + +function stripAsciiControls(value: string): string { + return Array.from(value, (character) => { + const code = character.charCodeAt(0); + return code <= 0x1f || code === 0x7f ? "" : character; + }).join(""); +} + +export const safeMarkdownUrl: UrlTransform = (url, key, node) => { + if (node.tagName !== "img") { + return defaultUrlTransform(url, key, node); + } + + // Browsers discard ASCII controls while parsing URLs, so strip them before + // rejecting remote schemes and protocol-relative image locations. + const normalized = stripAsciiControls(url).trim(); + const lower = normalized.toLowerCase(); + + if (lower.startsWith("data:") || lower.startsWith("blob:")) { + return normalized; + } + if (PROTOCOL_RELATIVE_RE.test(normalized)) { + return null; + } + if (SCHEME_RE.test(normalized)) { + return null; + } + return normalized; +}; diff --git a/tests/studio/test_deep_research_frontend_contract.py b/tests/studio/test_deep_research_frontend_contract.py index 2a22257dd9..55a33ff24e 100644 --- a/tests/studio/test_deep_research_frontend_contract.py +++ b/tests/studio/test_deep_research_frontend_contract.py @@ -86,6 +86,8 @@ def test_research_presentation_is_integrated() -> None: store = source("features/chat/stores/chat-runtime-store.ts") activity = source("features/chat/components/research-activity-panel.tsx") message = source("features/chat/components/research-message.tsx") + markdown_preview = source("components/markdown/markdown-preview.tsx") + safe_markdown_url = source("lib/safe-markdown-url.ts") coordinator = source("features/chat/stores/research-run-store.ts") assert "DeepResearchComposerButton" in thread assert "Deep research" in thread @@ -103,6 +105,9 @@ def test_research_presentation_is_integrated() -> None: assert "Stop research" not in activity assert "retryResearchRun" in activity assert "Deep research completed" in message + assert "