fix(backend): limit uploaded file context manifest (#3917)

* limit uploaded file context manifest

* fix: address uploaded file context review
This commit is contained in:
thefoolgy 2026-07-04 08:58:12 +08:00 committed by GitHub
parent 48477d868b
commit 76aa599107
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 258 additions and 12 deletions

View File

@ -1,6 +1,8 @@
"""Middleware to inject uploaded files information into agent context.""" """Middleware to inject uploaded files information into agent context."""
import logging import logging
import re
from collections import Counter
from pathlib import Path from pathlib import Path
from typing import NotRequired, override from typing import NotRequired, override
@ -14,12 +16,53 @@ from deerflow.config.paths import Paths, get_paths
from deerflow.runtime.user_context import get_effective_user_id from deerflow.runtime.user_context import get_effective_user_id
from deerflow.uploads.manager import is_upload_staging_file from deerflow.uploads.manager import is_upload_staging_file
from deerflow.utils.file_conversion import extract_outline from deerflow.utils.file_conversion import extract_outline
from deerflow.utils.messages import ORIGINAL_USER_CONTENT_KEY, message_content_to_text from deerflow.utils.messages import ORIGINAL_USER_CONTENT_KEY, get_original_user_content_text, message_content_to_text
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_OUTLINE_PREVIEW_LINES = 5 _OUTLINE_PREVIEW_LINES = 5
_MAX_FILES_PER_CONTEXT_SECTION = 10
_QUERY_TOKEN_RE = re.compile(r"[a-z0-9]+")
def _extension_label(file: dict) -> str:
extension = str(file.get("extension") or Path(str(file.get("filename") or "")).suffix).lower()
return extension or "(no extension)"
def _format_omitted_file_types(files: list[dict]) -> str:
counts = Counter(_extension_label(file) for file in files)
parts = [f"{count} {extension}" for extension, count in sorted(counts.items())]
return ", ".join(parts)
def _query_match_strength(file: dict, query_text: str) -> int:
query = query_text.lower()
if not query:
return 0
filename = str(file.get("filename") or "").lower()
stem = Path(filename).stem
extension_label = _extension_label(file)
extension = extension_label[1:] if extension_label.startswith(".") else ""
if filename and filename in query:
return 3
if len(stem) >= 3 and stem in query:
return 3
token_match = False
for token in _QUERY_TOKEN_RE.findall(stem):
if len(token) >= 3 and token in query:
token_match = True
break
if token_match:
return 2
if extension and re.search(rf"\b{re.escape(extension)}s?\b", query):
return 1
return 0
def _extract_outline_for_file(file_path: Path) -> tuple[list[dict], list[str]]: def _extract_outline_for_file(file_path: Path) -> tuple[list[dict], list[str]]:
@ -76,14 +119,24 @@ class UploadsMiddleware(AgentMiddleware[UploadsMiddlewareState]):
state_schema = UploadsMiddlewareState state_schema = UploadsMiddlewareState
def __init__(self, base_dir: str | None = None): def __init__(
self,
base_dir: str | None = None,
*,
max_files_per_context_section: int = _MAX_FILES_PER_CONTEXT_SECTION,
):
"""Initialize the middleware. """Initialize the middleware.
Args: Args:
base_dir: Base directory for thread data. Defaults to Paths resolution. base_dir: Base directory for thread data. Defaults to Paths resolution.
max_files_per_context_section: Maximum number of files listed in
each uploaded-files prompt section.
""" """
super().__init__() super().__init__()
if max_files_per_context_section < 1:
raise ValueError("max_files_per_context_section must be at least 1")
self._paths = Paths(base_dir) if base_dir else get_paths() self._paths = Paths(base_dir) if base_dir else get_paths()
self._max_files_per_context_section = max_files_per_context_section
def _format_file_entry(self, file: dict, lines: list[str]) -> None: def _format_file_entry(self, file: dict, lines: list[str]) -> None:
"""Append a single file entry (name, size, path, optional outline) to lines.""" """Append a single file entry (name, size, path, optional outline) to lines."""
@ -91,6 +144,8 @@ class UploadsMiddleware(AgentMiddleware[UploadsMiddlewareState]):
size_str = f"{size_kb:.1f} KB" if size_kb < 1024 else f"{size_kb / 1024:.1f} MB" size_str = f"{size_kb:.1f} KB" if size_kb < 1024 else f"{size_kb / 1024:.1f} MB"
lines.append(f"- {file['filename']} ({size_str})") lines.append(f"- {file['filename']} ({size_str})")
lines.append(f" Path: {file['path']}") lines.append(f" Path: {file['path']}")
if file.get("selection_reason") == "query_match":
lines.append(" Selected because: matched the current query.")
outline = file.get("outline") or [] outline = file.get("outline") or []
if outline: if outline:
truncated = outline[-1].get("truncated", False) truncated = outline[-1].get("truncated", False)
@ -109,7 +164,41 @@ class UploadsMiddleware(AgentMiddleware[UploadsMiddlewareState]):
lines.append(" Use `grep` to search for keywords (e.g. `grep(pattern='keyword', path='/mnt/user-data/uploads/')`).") lines.append(" Use `grep` to search for keywords (e.g. `grep(pattern='keyword', path='/mnt/user-data/uploads/')`).")
lines.append("") lines.append("")
def _create_files_message(self, new_files: list[dict], historical_files: list[dict]) -> str: def _select_files_for_context(
self,
files: list[dict],
query_text: str,
*,
recency_key: str | None = None,
) -> tuple[list[dict], list[dict]]:
"""Return bounded context files, prioritizing current-query matches."""
ranked: list[tuple[tuple, dict]] = []
for index, file in enumerate(files):
selected_file = dict(file)
match_strength = _query_match_strength(selected_file, query_text)
query_match = match_strength > 0
if query_match:
selected_file["selection_reason"] = "query_match"
if recency_key:
sort_key = (-match_strength, -float(selected_file.get(recency_key) or 0), selected_file["filename"])
else:
sort_key = (-match_strength, index)
ranked.append((sort_key, selected_file))
ranked.sort(key=lambda item: item[0])
selected = [file for _, file in ranked[: self._max_files_per_context_section]]
omitted = [file for _, file in ranked[self._max_files_per_context_section :]]
return selected, omitted
def _create_files_message(
self,
new_files: list[dict],
historical_files: list[dict],
*,
omitted_new_files: list[dict] | None = None,
omitted_historical_files: list[dict] | None = None,
) -> str:
"""Create a formatted message listing uploaded files. """Create a formatted message listing uploaded files.
Args: Args:
@ -117,6 +206,8 @@ class UploadsMiddleware(AgentMiddleware[UploadsMiddlewareState]):
historical_files: Files uploaded in previous messages. historical_files: Files uploaded in previous messages.
Each file dict may contain an optional ``outline`` key a list of Each file dict may contain an optional ``outline`` key a list of
``{title, line}`` dicts extracted from the converted Markdown file. ``{title, line}`` dicts extracted from the converted Markdown file.
omitted_new_files: Current-message files omitted from the prompt context.
omitted_historical_files: Older historical files omitted from the prompt context.
Returns: Returns:
Formatted string inside <uploaded_files> tags. Formatted string inside <uploaded_files> tags.
@ -128,6 +219,12 @@ class UploadsMiddleware(AgentMiddleware[UploadsMiddlewareState]):
if new_files: if new_files:
for file in new_files: for file in new_files:
self._format_file_entry(file, lines) self._format_file_entry(file, lines)
if omitted_new_files:
lines.append(f"... ({len(omitted_new_files)} more file(s) from this message omitted from this context.)")
lines.append(f" Omitted file types: {_format_omitted_file_types(omitted_new_files)}")
lines.append(" Use `glob(pattern='**/*', path='/mnt/user-data/uploads/')` to list all uploads.")
lines.append(" Use `grep(pattern='keyword', path='/mnt/user-data/uploads/')` to search across uploads.")
lines.append("")
else: else:
lines.append("(empty)") lines.append("(empty)")
lines.append("") lines.append("")
@ -137,6 +234,12 @@ class UploadsMiddleware(AgentMiddleware[UploadsMiddlewareState]):
lines.append("") lines.append("")
for file in historical_files: for file in historical_files:
self._format_file_entry(file, lines) self._format_file_entry(file, lines)
if omitted_historical_files:
lines.append(f"... ({len(omitted_historical_files)} more historical file(s) omitted from this context.)")
lines.append(f" Omitted file types: {_format_omitted_file_types(omitted_historical_files)}")
lines.append(" Use `glob(pattern='**/*', path='/mnt/user-data/uploads/')` to list all uploads.")
lines.append(" Use `grep(pattern='keyword', path='/mnt/user-data/uploads/')` to search across uploads.")
lines.append("")
lines.append("To work with these files:") lines.append("To work with these files:")
lines.append("- Read from the file first — use the outline line numbers and `read_file` to locate relevant sections.") lines.append("- Read from the file first — use the outline line numbers and `read_file` to locate relevant sections.")
@ -227,45 +330,68 @@ class UploadsMiddleware(AgentMiddleware[UploadsMiddlewareState]):
pass # get_config() raises outside a runnable context (e.g. unit tests) pass # get_config() raises outside a runnable context (e.g. unit tests)
uploads_dir = self._paths.sandbox_uploads_dir(thread_id, user_id=get_effective_user_id()) if thread_id else None uploads_dir = self._paths.sandbox_uploads_dir(thread_id, user_id=get_effective_user_id()) if thread_id else None
query_text = get_original_user_content_text(last_message.content, last_message.additional_kwargs)
# Get newly uploaded files from the current message's additional_kwargs.files # Get newly uploaded files from the current message's additional_kwargs.files
new_files = self._files_from_kwargs(last_message, uploads_dir) or [] new_files = self._files_from_kwargs(last_message, uploads_dir) or []
context_new_files, omitted_new_files = self._select_files_for_context(new_files, query_text)
# Collect historical files from the uploads directory (all except the new ones) # Collect historical files from the uploads directory (all except the new ones)
new_filenames = {f["filename"] for f in new_files} new_filenames = {f["filename"] for f in new_files}
historical_files: list[dict] = [] historical_candidates: list[dict] = []
if uploads_dir and uploads_dir.exists(): if uploads_dir and uploads_dir.exists():
for file_path in sorted(uploads_dir.iterdir()): for file_path in sorted(uploads_dir.iterdir()):
if is_upload_staging_file(file_path.name): if is_upload_staging_file(file_path.name):
continue continue
if file_path.is_file() and file_path.name not in new_filenames: if file_path.is_file() and file_path.name not in new_filenames:
stat = file_path.stat() stat = file_path.stat()
outline, preview = _extract_outline_for_file(file_path) historical_candidates.append(
historical_files.append(
{ {
"filename": file_path.name, "filename": file_path.name,
"size": stat.st_size, "size": stat.st_size,
"path": f"/mnt/user-data/uploads/{file_path.name}", "path": f"/mnt/user-data/uploads/{file_path.name}",
"extension": file_path.suffix, "extension": file_path.suffix,
"outline": outline, "_mtime": stat.st_mtime,
"outline_preview": preview, "_host_path": file_path,
} }
) )
historical_files, omitted_historical_files = self._select_files_for_context(
historical_candidates,
query_text,
recency_key="_mtime",
)
for file in historical_files:
file_path = file.pop("_host_path")
file.pop("_mtime", None)
outline, preview = _extract_outline_for_file(file_path)
file["outline"] = outline
file["outline_preview"] = preview
# Attach outlines to new files as well # Attach outlines to new files as well
if uploads_dir: if uploads_dir:
for file in new_files: new_files_by_name = {file["filename"]: file for file in new_files}
for file in context_new_files:
phys_path = uploads_dir / file["filename"] phys_path = uploads_dir / file["filename"]
outline, preview = _extract_outline_for_file(phys_path) outline, preview = _extract_outline_for_file(phys_path)
file["outline"] = outline file["outline"] = outline
file["outline_preview"] = preview file["outline_preview"] = preview
if original_file := new_files_by_name.get(file["filename"]):
original_file["outline"] = outline
original_file["outline_preview"] = preview
if not new_files and not historical_files: if not context_new_files and not historical_files:
return None return None
logger.debug(f"New files: {[f['filename'] for f in new_files]}, historical: {[f['filename'] for f in historical_files]}") logger.debug(f"New files: {[f['filename'] for f in new_files]}, historical: {[f['filename'] for f in historical_files]}")
# Create files message and prepend to the last human message content # Create files message and prepend to the last human message content
files_message = self._create_files_message(new_files, historical_files) files_message = self._create_files_message(
context_new_files,
historical_files,
omitted_new_files=omitted_new_files,
omitted_historical_files=omitted_historical_files,
)
# Extract original content - handle both string and list formats # Extract original content - handle both string and list formats
original_content = last_message.content original_content = last_message.content

View File

@ -7,6 +7,8 @@ Covers:
additional_kwargs, historical files from uploads dir, edge-cases) additional_kwargs, historical files from uploads dir, edge-cases)
""" """
import os
import re
from pathlib import Path from pathlib import Path
from unittest.mock import MagicMock from unittest.mock import MagicMock
@ -14,9 +16,10 @@ from langchain_core.messages import AIMessage, HumanMessage
from deerflow.agents.middlewares.uploads_middleware import UploadsMiddleware from deerflow.agents.middlewares.uploads_middleware import UploadsMiddleware
from deerflow.config.paths import Paths from deerflow.config.paths import Paths
from deerflow.utils.messages import ORIGINAL_USER_CONTENT_KEY from deerflow.utils.messages import ORIGINAL_USER_CONTENT_KEY, message_content_to_text
THREAD_ID = "thread-abc123" THREAD_ID = "thread-abc123"
CONTEXT_SECTION_LIMIT = 10
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@ -49,6 +52,13 @@ def _human(content, files=None, **extra_kwargs):
return HumanMessage(content=content, additional_kwargs=additional_kwargs) return HumanMessage(content=content, additional_kwargs=additional_kwargs)
def _uploaded_files_block(content) -> str:
text = message_content_to_text(content)
match = re.search(r"<uploaded_files>[\s\S]*?</uploaded_files>", text)
assert match is not None
return match.group(0)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# _files_from_kwargs # _files_from_kwargs
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@ -351,6 +361,72 @@ class TestBeforeAgent:
} }
] ]
def test_current_message_files_are_limited_in_context(self, tmp_path):
mw = _middleware(tmp_path)
uploads_dir = _uploads_dir(tmp_path)
total_files = CONTEXT_SECTION_LIMIT + 2
files = []
for i in range(total_files):
filename = f"current_{i:02}.txt"
(uploads_dir / filename).write_text(f"new upload {i}", encoding="utf-8")
files.append({"filename": filename, "size": 12, "path": f"/mnt/user-data/uploads/{filename}"})
result = mw.before_agent(self._state(_human("compare these files", files=files)), _runtime())
assert result is not None
content = result["messages"][-1].content
assert "current_09.txt" in content
assert "current_10.txt" not in content
assert "current_11.txt" not in content
assert "2 more file(s) from this message omitted from this context" in content
assert "Omitted file types: 2 .txt" in content
assert len(result["uploaded_files"]) == total_files
def test_current_message_query_matches_are_selected_before_upload_order(self, tmp_path):
mw = _middleware(tmp_path)
uploads_dir = _uploads_dir(tmp_path)
total_files = CONTEXT_SECTION_LIMIT + 2
files = []
for i in range(total_files):
filename = f"current_{i:02}.txt"
(uploads_dir / filename).write_text(f"new upload {i}", encoding="utf-8")
files.append({"filename": filename, "size": 12, "path": f"/mnt/user-data/uploads/{filename}"})
result = mw.before_agent(self._state(_human("please inspect current_11.txt", files=files)), _runtime())
assert result is not None
content = _uploaded_files_block(result["messages"][-1].content)
assert "current_11.txt" in content
assert "Selected because: matched the current query." in content
assert "current_10.txt" not in content
assert "2 more file(s) from this message omitted from this context" in content
def test_current_message_ranking_uses_original_user_content(self, tmp_path):
mw = _middleware(tmp_path)
uploads_dir = _uploads_dir(tmp_path)
total_files = CONTEXT_SECTION_LIMIT + 2
files = []
for i in range(total_files):
filename = f"current_{i:02}.txt"
(uploads_dir / filename).write_text(f"new upload {i}", encoding="utf-8")
files.append({"filename": filename, "size": 12, "path": f"/mnt/user-data/uploads/{filename}"})
msg = _human(
"<uploaded_files>\ncurrent_11.txt\n</uploaded_files>\n\ncompare these files",
files=files,
**{ORIGINAL_USER_CONTENT_KEY: "compare these files"},
)
result = mw.before_agent(self._state(msg), _runtime())
assert result is not None
content = _uploaded_files_block(result["messages"][-1].content)
assert "current_09.txt" in content
assert "current_10.txt" not in content
assert "current_11.txt" not in content
def test_historical_files_from_uploads_dir_excluding_new(self, tmp_path): def test_historical_files_from_uploads_dir_excluding_new(self, tmp_path):
mw = _middleware(tmp_path) mw = _middleware(tmp_path)
uploads_dir = _uploads_dir(tmp_path) uploads_dir = _uploads_dir(tmp_path)
@ -383,6 +459,50 @@ class TestBeforeAgent:
assert ".env" in content assert ".env" in content
assert ".upload-active.part" not in content assert ".upload-active.part" not in content
def test_historical_files_are_limited_to_recent_context_entries(self, tmp_path):
mw = _middleware(tmp_path)
uploads_dir = _uploads_dir(tmp_path)
total_files = CONTEXT_SECTION_LIMIT + 2
for i in range(total_files):
file_path = uploads_dir / f"history_{i:02}.txt"
file_path.write_text(f"old upload {i}", encoding="utf-8")
os.utime(file_path, (i + 1, i + 1))
result = mw.before_agent(self._state(_human("what files do I have?")), _runtime())
assert result is not None
content = result["messages"][-1].content
assert "history_11.txt" in content
assert "history_02.txt" in content
assert "history_01.txt" not in content
assert "history_00.txt" not in content
assert "2 more historical file(s) omitted from this context" in content
assert "Omitted file types: 2 .txt" in content
def test_historical_query_matches_are_selected_before_recency(self, tmp_path):
mw = _middleware(tmp_path)
uploads_dir = _uploads_dir(tmp_path)
target = uploads_dir / "tax_report_2019.pdf"
target.write_text("old but relevant", encoding="utf-8")
os.utime(target, (1, 1))
for i in range(CONTEXT_SECTION_LIMIT):
file_path = uploads_dir / f"recent_{i:02}.txt"
file_path.write_text(f"recent upload {i}", encoding="utf-8")
os.utime(file_path, (100 + i, 100 + i))
result = mw.before_agent(self._state(_human("analyze tax_report_2019.pdf")), _runtime())
assert result is not None
content = _uploaded_files_block(result["messages"][-1].content)
assert "tax_report_2019.pdf" in content
assert "Selected because: matched the current query." in content
assert "recent_00.txt" not in content
assert "1 more historical file(s) omitted from this context" in content
assert "Omitted file types: 1 .txt" in content
def test_no_historical_section_when_upload_dir_is_empty(self, tmp_path): def test_no_historical_section_when_upload_dir_is_empty(self, tmp_path):
mw = _middleware(tmp_path) mw = _middleware(tmp_path)
uploads_dir = _uploads_dir(tmp_path) uploads_dir = _uploads_dir(tmp_path)