Skip to content
Merged
Show file tree
Hide file tree
Changes from 4 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 8 additions & 6 deletions openrag/api/routers/user/source_links.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,9 +33,11 @@ def build_document_source_link(
encoded_url = None
if filename:
encoded_url = quote(static_url_builder(doc_metadata["_id"]), safe=":/")
return {
"source_type": "document",
**({"file_url": encoded_url} if encoded_url else {}),
"chunk_url": chunk_url_builder(doc_metadata["_id"]),
**doc_metadata,
}
link = dict(doc_metadata)
link["source_type"] = "document"
link["chunk_url"] = chunk_url_builder(doc_metadata["_id"])
if encoded_url:
link["file_url"] = encoded_url
else:
link.pop("file_url", None)
return link
136 changes: 99 additions & 37 deletions openrag/app_front.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,11 @@
import json
import os
import secrets
import string
import time
from functools import lru_cache
from pathlib import Path
from urllib.parse import urlparse
from urllib.parse import quote, urlparse

import chainlit as cl
import httpx
Expand Down Expand Up @@ -42,12 +43,31 @@
OPENRAG_CHAT_PROFILES_METADATA_KEY = "openrag_chat_profiles"
OPENRAG_SESSION_COOKIE_NAME = "openrag_session"
_OPENRAG_TOKEN_STORE: dict[str, tuple[str, float]] = {}
_MARKDOWN_ESCAPE_TABLE = str.maketrans({char: f"\\{char}" for char in string.punctuation})
_MARKDOWN_URL_SAFE_CHARS = ":/?#[]@!$&'+,;=%"
_MARKDOWN_UNSAFE_SOURCE_NAME_CHARS = str.maketrans(dict.fromkeys("[]()*_`~#>|\\", " "))


class MissingOpenRAGCredentialError(RuntimeError):
pass


def _escape_markdown_text(value: str) -> str:
"""Render untrusted source metadata as literal Markdown text."""
return value.translate(_MARKDOWN_ESCAPE_TABLE)


def _safe_source_name(value: str, existing: dict) -> str:
"""Build a Markdown-inert name that Chainlit can match to its element."""
base = " ".join(value.translate(_MARKDOWN_UNSAFE_SOURCE_NAME_CHARS).split()) or "source"
Comment thread
hedhoud marked this conversation as resolved.
candidate = base
suffix = 2
while candidate in existing:
candidate = f"{base} {suffix}"
suffix += 1
return candidate


def get_user_language() -> str:
"""Return the active language: env override if set, otherwise browser's Accept-Language."""
if DEFAULT_LANGUAGE:
Expand Down Expand Up @@ -465,26 +485,53 @@ async def __fetch_page_content(chunk_url, headers=None):


async def _format_sources(metadata_sources, only_txt=False, api_key=None):
external_url = get_external_url() # used to override the base URL when the front-end requests a file resource
if not metadata_sources:
return None, None
if not isinstance(metadata_sources, list) or not metadata_sources:
return [], []

d = {}
headers = get_headers(api_key)
external_url = get_external_url() # used to override the base URL when the front-end requests a file resource
for i, s in enumerate(metadata_sources):
if not isinstance(s, dict):
continue

if s.get("source_type") == "web":
title = s.get("title") or s.get("url", f"Web source {i + 1}")
title = s.get("title", "")
url = s.get("url", "")
snippet = s.get("snippet", "")
content = f"**[{title}]({url})**\n\n{snippet}"
source_name = title
if source_name in d:
source_name = f"{title} ({i})"
title = title.strip() if isinstance(title, str) else ""
url = url.strip() if isinstance(url, str) else ""
snippet = snippet.strip() if isinstance(snippet, str) else ""
try:
parsed_url = urlparse(url)
if parsed_url.scheme not in {"http", "https"} or not parsed_url.netloc:
continue
markdown_url = quote(str(httpx.URL(url)), safe=_MARKDOWN_URL_SAFE_CHARS)
except (ValueError, httpx.InvalidURL):
continue

source_label = title or url
source_name = _safe_source_name(source_label, d)
content = f"**[{_escape_markdown_text(source_label)}]({markdown_url})**"
if snippet:
content += f"\n\n{_escape_markdown_text(snippet)}"
d[source_name] = cl.Text(content=content, name=source_name, display="side")
continue
Comment thread
coderabbitai[bot] marked this conversation as resolved.

filename = Path(s["filename"])
file_url = s["file_url"]
filename_value = s.get("filename")
file_url = s.get("file_url")
page = s.get("page")
if (
not isinstance(filename_value, str)
or not filename_value.strip()
or not isinstance(file_url, str)
or not file_url.strip()
):
continue

filename = Path(filename_value.strip())
suffix = filename.suffix.lower()
file_url = file_url.strip()
file_url = file_url.replace(INTERNAL_BASE_URL, external_url) # put the correct base url
# Avoid leaking the credential in the URL (browser history, proxy logs,
# Referer headers). In OIDC mode the browser already sends the
Expand All @@ -495,36 +542,51 @@ async def _format_sources(metadata_sources, only_txt=False, api_key=None):
# authenticate the fetch.
if api_key and (AUTH_MODE != "oidc" or _current_openrag_auth_provider() == "credentials"):
file_url = f"{file_url}?token={api_key}"
page = s["page"]
source_name = f"{filename}" + (
f" (page: {page})" if filename.suffix in [".pdf", ".pptx", ".docx", ".doc"] else ""
page_label = str(page).strip() if page is not None else ""
source_label = f"{filename}" + (
f" (page: {page_label})" if suffix in [".pdf", ".pptx", ".docx", ".doc"] and page_label else ""
)
source_name = _safe_source_name(source_label, d)

if only_txt:
chunk_content = await __fetch_page_content(chunk_url=s["chunk_url"], headers=headers)
elem = cl.Text(content=chunk_content, name=source_name, display="side")
else:
match filename.suffix.lower():
case ".pdf":
elem = cl.Pdf(
name=source_name,
url=file_url,
page=int(s["page"]),
display="side",
)
case suffix if suffix in [".png", ".jpg", ".jpeg"]:
elem = cl.Image(name=source_name, url=file_url, display="side")
case ".mp4":
elem = cl.Video(name=source_name, url=file_url, display="side")
case ".mp3":
elem = cl.Audio(name=source_name, url=file_url, display="side")
case _:
chunk_content = await __fetch_page_content(chunk_url=s["chunk_url"], headers=headers)
elem = cl.Text(content=chunk_content, name=source_name, display="side")
try:
if only_txt:
chunk_url = s.get("chunk_url")
if not isinstance(chunk_url, str) or not chunk_url.strip():
continue
chunk_content = await __fetch_page_content(chunk_url=chunk_url, headers=headers)
if not isinstance(chunk_content, str) or not chunk_content.strip():
continue
elem = cl.Text(content=chunk_content, name=source_name, display="side")
else:
match suffix:
case ".pdf":
elem = cl.Pdf(
name=source_name,
url=file_url,
page=int(page) if page_label else None,
display="side",
)
case suffix if suffix in [".png", ".jpg", ".jpeg"]:
elem = cl.Image(name=source_name, url=file_url, display="side")
case ".mp4":
elem = cl.Video(name=source_name, url=file_url, display="side")
case ".mp3":
elem = cl.Audio(name=source_name, url=file_url, display="side")
case _:
chunk_url = s.get("chunk_url")
if not isinstance(chunk_url, str) or not chunk_url.strip():
continue
chunk_content = await __fetch_page_content(chunk_url=chunk_url, headers=headers)
if not isinstance(chunk_content, str) or not chunk_content.strip():
continue
elem = cl.Text(content=chunk_content, name=source_name, display="side")
except (httpx.HTTPError, httpx.InvalidURL, TypeError, ValueError, AttributeError):
logger.warning("Skipping an unavailable source", source_index=i)
continue
Comment thread
coderabbitai[bot] marked this conversation as resolved.

d[source_name] = elem

source_names = list(d.keys())
source_names = list(d)
elements = list(d.values())

return elements, source_names
Expand Down Expand Up @@ -582,7 +644,7 @@ async def on_message(message: cl.Message):
# Show sources
elements, source_names = await _format_sources(sources, api_key=api_key, only_txt=False)
msg.elements = elements if elements else []
if source_names:
if elements and source_names:
s = "\n\n" + "-" * 50 + f"\n\n{t('sources_label')}: \n" + "\n".join(source_names)
await msg.stream_token(s)
await msg.update()
Expand Down
22 changes: 22 additions & 0 deletions tests/unit/api/routers/user/test_source_links.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,3 +61,25 @@ def test_metadata_is_passed_through():
link = _build({"_id": "x", "source": "doc.pdf", "author": "alice"})
assert link["author"] == "alice"
assert link["chunk_url"] == "https://host/extract/x"


def test_metadata_cannot_override_authoritative_source_fields():
link = _build(
{
"_id": "42",
"source": "diagram.png",
"file_url": "https://attacker.example/file",
"chunk_url": "https://attacker.example/chunk",
"source_type": "web",
}
)

assert link["source_type"] == "document"
assert link["file_url"] == "https://host/static/42"
assert link["chunk_url"] == "https://host/extract/42"


def test_metadata_file_url_is_removed_when_source_is_missing():
link = _build({"_id": "42", "file_url": "https://attacker.example/file"})

assert "file_url" not in link
Loading
Loading