mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-11 03:12:55 +00:00
Merge pull request #2841 from arc53-machine/default-chunks-6
Default retrieval to 6 chunks instead of 2
This commit is contained in:
36 files changed
+223
-50
No files matched your search
@@ -65,7 +65,7 @@ Common request body fields:
|
||||
| `prompt_id` | `string` | No | `/api/answer`, `/stream` | Ignored when `api_key` already defines prompt. |
|
||||
| `active_docs` | `string` or `string[]` | No | `/api/answer`, `/stream` | Overrides active docs when not using key-owned source config. |
|
||||
| `retriever` | `string` | No | `/api/answer`, `/stream` | Retriever type (for example `classic`). |
|
||||
| `chunks` | `number` | No | `/api/answer`, `/stream` | Retrieval chunk count, default `2`. |
|
||||
| `chunks` | `number` | No | `/api/answer`, `/stream` | Retrieval chunk count, default `6`: a total for the request, split across its sources. |
|
||||
| `isNoneDoc` | `boolean` | No | `/api/answer`, `/stream` | Skip document retrieval. |
|
||||
| `agent_id` | `string` | No | `/api/answer`, `/stream` | Alternative to `api_key` when using authenticated user context. |
|
||||
|
||||
|
||||
@@ -33,7 +33,7 @@ When you create or configure an agent, you'll work with these key components:
|
||||
|
||||
**Source:**
|
||||
* **Select source:** The knowledge base for the agent. You can select from previously uploaded documents or data sources. This is what the agent will "know."
|
||||
* **Chunks per query:** A numerical value determining how many relevant text chunks from the selected source are sent to the LLM with each query. This helps manage context length and relevance.
|
||||
* **Chunks:** How many relevant text chunks are sent to the LLM with each query, in total across the agent's sources (default 6). The agent form doesn't show it; set `chunks` when creating or updating the agent through the API, or in an imported agent YAML file (see [Agents API](/Agents/api)).
|
||||
|
||||
**Prompt:**
|
||||
The main set of instructions or system [prompt](/Guides/Customising-prompts) that defines the agent's persona, objectives, constraints, and how it should behave or respond.
|
||||
|
||||
@@ -79,7 +79,7 @@ Retrieval decides which chunks are pulled in to answer a question. These setting
|
||||
"retrieval": {
|
||||
"retriever": "classic",
|
||||
"exposure": "prefetch",
|
||||
"chunks": 2,
|
||||
"chunks": 6,
|
||||
"score_threshold": null,
|
||||
"rephrase_query": true,
|
||||
"prescreen": null
|
||||
@@ -91,7 +91,7 @@ Retrieval decides which chunks are pulled in to answer a question. These setting
|
||||
| --- | --- | --- |
|
||||
| `retriever` | `classic` | Retrieval strategy: `classic`, `hybrid`, or `graphrag`. |
|
||||
| `exposure` | `prefetch` | How retrieved context reaches the model: `prefetch` or `agentic_tool` (see below). |
|
||||
| `chunks` | `2` | Final number of chunks (top-k) returned to the answer. Range 1–500. Set here, it **overrides** whatever a request asks for. |
|
||||
| `chunks` | `6` | Final number of chunks (top-k) returned to the answer. Range 1–500. Set here, it **overrides** whatever a request asks for. |
|
||||
| `score_threshold` | `null` | Minimum similarity score. Honored by pgvector and MongoDB Atlas; FAISS, Qdrant, Milvus and the `hybrid` retriever ignore it — the config API returns a `warnings` entry when you set it on one of those. |
|
||||
| `rephrase_query` | `true` | Whether to run a query-rephrasing side-call before retrieval. |
|
||||
| `prescreen` | `null` | Optional LLM relevance filter (see below). `null` = off. |
|
||||
@@ -102,6 +102,14 @@ A source that has configured `retrieval.chunks` **outranks the value sent with a
|
||||
request**. The owner tuned top-k for that corpus, so a client cannot raise or
|
||||
lower it per call. Sources left at the default still let the request decide.
|
||||
|
||||
`chunks` is a total for the request, not a count per source: a request with
|
||||
`chunks: 6` over three sources takes two chunks from each. Every attached source
|
||||
still gets at least one chunk, so with more sources than `chunks` the answer
|
||||
receives one chunk per source (eight sources at `chunks: 6` return eight
|
||||
chunks). A source counts as configured when its stored
|
||||
value differs from the default, so a source saved while the default was `2`
|
||||
keeps `2` as its own setting.
|
||||
|
||||
Requests are also bounded: `chunks` is clamped to 0–500, and `0` still means
|
||||
"skip retrieval for this turn".
|
||||
|
||||
|
||||
@@ -153,7 +153,9 @@ def _run_agent_headless(
|
||||
source_active = str(src_row["id"])
|
||||
retriever_kind = src_row.get("retriever", retriever_kind)
|
||||
source = {"active_docs": source_active}
|
||||
chunks = int(agent_config.get("chunks", 2) or 2)
|
||||
# ``chunks=0`` switches retrieval off; only a missing value takes the default.
|
||||
raw_chunks = agent_config.get("chunks")
|
||||
chunks = 6 if raw_chunks in (None, "") else int(raw_chunks)
|
||||
prompt_id = agent_config.get("prompt_id", "default")
|
||||
user_api_key = agent_config.get("key")
|
||||
agent_id = _resolve_agent_id(agent_config)
|
||||
|
||||
@@ -36,7 +36,7 @@ class InternalSearchTool(Tool):
|
||||
source=self.config.get("source", {}),
|
||||
chat_history=[],
|
||||
prompt="",
|
||||
chunks=int(self.config.get("chunks", 2)),
|
||||
chunks=int(self.config.get("chunks", 6)),
|
||||
doc_token_limit=int(self.config.get("doc_token_limit", 50000)),
|
||||
model_id=self.config.get("model_id", "docsgpt-local"),
|
||||
model_user_id=self.config.get("model_user_id"),
|
||||
@@ -464,7 +464,7 @@ def add_internal_search_tool(tools_dict: Dict, retriever_config: Dict) -> None:
|
||||
def build_internal_tool_config(
|
||||
source: Dict,
|
||||
retriever_name: str = "classic",
|
||||
chunks: int = 2,
|
||||
chunks: int = 6,
|
||||
doc_token_limit: int = 50000,
|
||||
sources: Optional[List[Dict]] = None,
|
||||
model_id: str = "docsgpt-local",
|
||||
|
||||
@@ -45,7 +45,7 @@ class AgentNodeConfig(BaseModel):
|
||||
stream_to_user: bool = True
|
||||
tools: List[str] = Field(default_factory=list)
|
||||
sources: List[str] = Field(default_factory=list)
|
||||
chunks: str = "2"
|
||||
chunks: str = "6"
|
||||
retriever: str = ""
|
||||
model_id: Optional[str] = None
|
||||
json_schema: Optional[Dict[str, Any]] = None
|
||||
|
||||
@@ -429,7 +429,7 @@ class WorkflowEngine:
|
||||
else {}
|
||||
),
|
||||
"retriever_name": node_config.retriever or "classic",
|
||||
"chunks": int(node_config.chunks) if node_config.chunks else 2,
|
||||
"chunks": int(node_config.chunks) if node_config.chunks else 6,
|
||||
"model_id": node_model_id,
|
||||
"llm_name": node_llm_name,
|
||||
"api_key": node_api_key,
|
||||
@@ -1388,7 +1388,7 @@ class WorkflowEngine:
|
||||
source={"active_docs": self._authorized_node_sources(node_config.sources)},
|
||||
chat_history=[],
|
||||
prompt="",
|
||||
chunks=int(node_config.chunks) if node_config.chunks else 2,
|
||||
chunks=int(node_config.chunks) if node_config.chunks else 6,
|
||||
decoded_token=self.agent.decoded_token,
|
||||
)
|
||||
docs = retriever.search(query)
|
||||
|
||||
@@ -46,7 +46,7 @@ class AnswerResource(Resource, BaseAnswerResource):
|
||||
required=False, default="default", description="Prompt ID"
|
||||
),
|
||||
"chunks": fields.Integer(
|
||||
required=False, default=2, description="Number of chunks"
|
||||
required=False, default=6, description="Number of chunks"
|
||||
),
|
||||
"retriever": fields.String(required=False, description="Retriever type"),
|
||||
"api_key": fields.String(required=False, description="API key"),
|
||||
|
||||
@@ -47,7 +47,7 @@ class StreamResource(Resource, BaseAnswerResource):
|
||||
required=False, default="default", description="Prompt ID"
|
||||
),
|
||||
"chunks": fields.Integer(
|
||||
required=False, default=2, description="Number of chunks"
|
||||
required=False, default=6, description="Number of chunks"
|
||||
),
|
||||
"retriever": fields.String(required=False, description="Retriever type"),
|
||||
"api_key": fields.String(required=False, description="API key"),
|
||||
|
||||
@@ -698,7 +698,7 @@ class StreamProcessor:
|
||||
"retriever": src_retriever or "classic",
|
||||
"chunks": (
|
||||
src_chunks if src_chunks is not None
|
||||
else data.get("chunks", "2")
|
||||
else data.get("chunks", "6")
|
||||
),
|
||||
# Per-source behaviour contract (lenient read).
|
||||
"retrieval": SourceConfig.parse(
|
||||
@@ -729,7 +729,7 @@ class StreamProcessor:
|
||||
"retriever": src_retriever or "classic",
|
||||
"chunks": (
|
||||
src_chunks if src_chunks is not None
|
||||
else data.get("chunks", "2")
|
||||
else data.get("chunks", "6")
|
||||
),
|
||||
"retrieval": SourceConfig.parse(
|
||||
source_doc.get("config")
|
||||
@@ -1056,7 +1056,7 @@ class StreamProcessor:
|
||||
)
|
||||
|
||||
retriever_name = "classic"
|
||||
chunks = 2
|
||||
chunks = 6
|
||||
|
||||
if self._agent_data is not None:
|
||||
# Agent-bound: agent wins, body's retriever/chunks are dropped.
|
||||
@@ -1068,7 +1068,7 @@ class StreamProcessor:
|
||||
except (ValueError, TypeError):
|
||||
logger.warning(
|
||||
f"Invalid agent chunks value: {self._agent_data['chunks']}, "
|
||||
"using default value 2"
|
||||
"using default value 6"
|
||||
)
|
||||
else:
|
||||
if "retriever" in self.data:
|
||||
@@ -1079,7 +1079,7 @@ class StreamProcessor:
|
||||
except (ValueError, TypeError):
|
||||
logger.warning(
|
||||
f"Invalid request chunks value: {self.data['chunks']}, "
|
||||
"using default value 2"
|
||||
"using default value 6"
|
||||
)
|
||||
# A source that configured its own retrieval knobs outranks the
|
||||
# request body: the owner tuned top-k for that corpus, a client
|
||||
@@ -1893,7 +1893,7 @@ class StreamProcessor:
|
||||
"retriever_name": self.retriever_config.get(
|
||||
"retriever_name", "classic"
|
||||
),
|
||||
"chunks": self.retriever_config.get("chunks", 2),
|
||||
"chunks": self.retriever_config.get("chunks", 6),
|
||||
"doc_token_limit": self.retriever_config.get(
|
||||
"doc_token_limit", 50000
|
||||
),
|
||||
|
||||
@@ -1552,9 +1552,9 @@ def apply_import(conn, user: str, doc: dict, resolution: Optional[dict] = None)
|
||||
slug = _unique_slug(agents_repo, user, metadata.get("slug") or spec.get("name"), exclude_id=exclude_id)
|
||||
|
||||
try:
|
||||
chunks_value = int(spec["chunks"]) if spec.get("chunks") is not None else 2
|
||||
chunks_value = int(spec["chunks"]) if spec.get("chunks") is not None else 6
|
||||
except (TypeError, ValueError):
|
||||
chunks_value = 2
|
||||
chunks_value = 6
|
||||
|
||||
# YAML-authoritative fields — written even when the resolved value is None,
|
||||
# so a re-import can CLEAR models / json_schema / prompt-to-default. On
|
||||
|
||||
@@ -259,7 +259,7 @@ def _format_agent_output(
|
||||
),
|
||||
"source": source_value,
|
||||
"sources": sources_list,
|
||||
"chunks": str(agent["chunks"]) if agent.get("chunks") is not None else "2",
|
||||
"chunks": str(agent["chunks"]) if agent.get("chunks") is not None else "6",
|
||||
"retriever": agent.get("retriever", "") or "",
|
||||
"prompt_id": str(agent["prompt_id"]) if agent.get("prompt_id") else "",
|
||||
"tools": agent.get("tools", []) or [],
|
||||
@@ -743,7 +743,7 @@ class CreateAgent(Resource):
|
||||
# For classic agents: default chunks/retriever if nothing else supplied.
|
||||
if agent_type != "workflow":
|
||||
if build_data.get("chunks") in (None, ""):
|
||||
build_data["chunks"] = 2
|
||||
build_data["chunks"] = 6
|
||||
if (
|
||||
not source_id_resolved
|
||||
and not extra_source_ids
|
||||
@@ -989,7 +989,7 @@ class UpdateAgent(Resource):
|
||||
elif field == "chunks":
|
||||
chunks_value = data.get("chunks")
|
||||
if chunks_value in ("", None):
|
||||
update_fields["chunks"] = 2
|
||||
update_fields["chunks"] = 6
|
||||
else:
|
||||
try:
|
||||
chunks_int = int(chunks_value)
|
||||
|
||||
@@ -74,7 +74,7 @@ def _ephemeral_agent_for_agentless(
|
||||
"user_id": user_id,
|
||||
"agent_type": "classic",
|
||||
"retriever": "classic",
|
||||
"chunks": 2,
|
||||
"chunks": 6,
|
||||
"prompt_id": "default",
|
||||
"source_id": None,
|
||||
"default_model_id": schedule.get("model_id") or "",
|
||||
|
||||
@@ -205,7 +205,7 @@ class ShareConversation(Resource):
|
||||
|
||||
if is_promptable:
|
||||
prompt_id_raw = data.get("prompt_id", "default")
|
||||
chunks_raw = data.get("chunks", "2")
|
||||
chunks_raw = data.get("chunks", "6")
|
||||
try:
|
||||
chunks_int = int(chunks_raw) if chunks_raw not in (None, "") else None
|
||||
except (TypeError, ValueError):
|
||||
|
||||
@@ -32,7 +32,7 @@ class ClassicRAG(BaseRetriever):
|
||||
source,
|
||||
chat_history=None,
|
||||
prompt="",
|
||||
chunks=2,
|
||||
chunks=6,
|
||||
doc_token_limit=50000,
|
||||
model_id="docsgpt-local",
|
||||
user_api_key=None,
|
||||
@@ -54,9 +54,9 @@ class ClassicRAG(BaseRetriever):
|
||||
self.chunks = int(chunks)
|
||||
except ValueError:
|
||||
logger.warning(
|
||||
f"Invalid chunks value '{chunks}', using default value 2"
|
||||
f"Invalid chunks value '{chunks}', using default value 6"
|
||||
)
|
||||
self.chunks = 2
|
||||
self.chunks = 6
|
||||
else:
|
||||
self.chunks = chunks
|
||||
user_id = decoded_token.get("sub") if decoded_token else "default"
|
||||
@@ -405,7 +405,7 @@ class ClassicRAG(BaseRetriever):
|
||||
|
||||
# ``chunks_per_source`` has a floor of 1 so no attached source is
|
||||
# starved, which means N sources always yield at least N documents —
|
||||
# ``chunks=2`` across 4 sources returned 4, though ``chunks`` is
|
||||
# ``chunks=6`` across 8 sources returned 8, though ``chunks`` is
|
||||
# documented as a top-k. Bound the overshoot to exactly that floor so
|
||||
# attaching more sources can no longer inflate the result without limit.
|
||||
# Ceiling on ``self.chunks`` (the actual fetch target), not
|
||||
|
||||
@@ -57,7 +57,7 @@ class Dispatcher(BaseRetriever):
|
||||
source,
|
||||
chat_history=None,
|
||||
prompt="",
|
||||
chunks=2,
|
||||
chunks=6,
|
||||
doc_token_limit=50000,
|
||||
model_id="docsgpt-local",
|
||||
user_api_key=None,
|
||||
|
||||
@@ -195,7 +195,7 @@ class GraphRAGRetriever(BaseRetriever):
|
||||
source,
|
||||
chat_history=None,
|
||||
prompt="",
|
||||
chunks=2,
|
||||
chunks=6,
|
||||
doc_token_limit=50000,
|
||||
model_id="docsgpt-local",
|
||||
user_api_key=None,
|
||||
|
||||
@@ -118,7 +118,7 @@ class RetrievalConfig(BaseModel):
|
||||
|
||||
retriever: str = "classic" # RetrieverCreator key
|
||||
exposure: str = "prefetch" # prefetch | agentic_tool (D11)
|
||||
chunks: int = 2 # final top-k
|
||||
chunks: int = 6 # final top-k
|
||||
score_threshold: Optional[float] = None # pgvector/mongo honor it; others ignore
|
||||
rephrase_query: bool = True # toggle ClassicRAG._rephrase_query side-call
|
||||
reranker: Optional[dict] = None # reserved: future cross-encoder/LLM reorder
|
||||
|
||||
@@ -135,7 +135,7 @@ export default function NewAgent({ mode }: { mode: 'new' | 'edit' | 'draft' }) {
|
||||
image: '',
|
||||
source: '',
|
||||
sources: [],
|
||||
chunks: '2',
|
||||
chunks: '6',
|
||||
retriever: 'classic',
|
||||
prompt_id: 'default',
|
||||
tools: [],
|
||||
@@ -708,7 +708,7 @@ export default function NewAgent({ mode }: { mode: 'new' | 'edit' | 'draft' }) {
|
||||
agent_type: data.agent_type || 'classic',
|
||||
prompt_id: data.prompt_id || 'default',
|
||||
retriever: agentSourceIds.length === 0 ? 'classic' : '',
|
||||
chunks: data.chunks || '2',
|
||||
chunks: data.chunks || '6',
|
||||
tools: data.tools || [],
|
||||
...serializeAgentSources(agentSourceIds, sourceDocs),
|
||||
models: agentModels,
|
||||
|
||||
@@ -329,7 +329,7 @@ function createEmptyWorkflowAgent(): Agent {
|
||||
description: '',
|
||||
image: '',
|
||||
source: '',
|
||||
chunks: '2',
|
||||
chunks: '6',
|
||||
retriever: '',
|
||||
prompt_id: '',
|
||||
tools: [],
|
||||
|
||||
@@ -43,7 +43,7 @@ export type SourceGraphRetrievalConfig = {
|
||||
export type SourceRetrievalConfig = {
|
||||
retriever?: string; // default 'classic' (only option for now)
|
||||
exposure?: RetrievalExposure; // default 'prefetch'
|
||||
chunks?: number; // top-k, default 2
|
||||
chunks?: number; // top-k, default 6
|
||||
score_threshold?: number | null; // default null
|
||||
rephrase_query?: boolean; // default true
|
||||
prescreen?: SourcePrescreenConfig | null; // null = off
|
||||
|
||||
@@ -58,7 +58,7 @@ const initialState: Preference = {
|
||||
{ name: 'creative', id: 'creative', type: 'public' },
|
||||
{ name: 'strict', id: 'strict', type: 'public' },
|
||||
],
|
||||
chunks: '2',
|
||||
chunks: '6',
|
||||
selectedDocs: [],
|
||||
sourceDocs: null,
|
||||
conversations: {
|
||||
|
||||
@@ -16,6 +16,13 @@ const clone = (v: RetrievalOptionsValue): RetrievalOptionsValue =>
|
||||
JSON.parse(JSON.stringify(v));
|
||||
|
||||
describe('configToOptions (lenient read)', () => {
|
||||
it('defaults a source to 6 chunks, the backend default', () => {
|
||||
// Equal to the backend's RetrievalConfig default, so a new upload is not
|
||||
// read as a per-source override.
|
||||
expect(DEFAULT_RETRIEVAL_OPTIONS.retrieval.chunks).toBe(6);
|
||||
expect(configToOptions(undefined).retrieval.chunks).toBe(6);
|
||||
});
|
||||
|
||||
it('returns all defaults for an absent config', () => {
|
||||
expect(configToOptions(undefined)).toEqual(DEFAULT_RETRIEVAL_OPTIONS);
|
||||
});
|
||||
|
||||
@@ -84,7 +84,7 @@ export const DEFAULT_RETRIEVAL_OPTIONS: RetrievalOptionsValue = {
|
||||
retrieval: {
|
||||
retriever: 'classic',
|
||||
exposure: 'prefetch',
|
||||
chunks: 2,
|
||||
chunks: 6,
|
||||
score_threshold: null,
|
||||
rephrase_query: true,
|
||||
prescreen: {
|
||||
|
||||
@@ -38,7 +38,7 @@ const preloadedState: { preference: Preference } = {
|
||||
{ name: 'creative', id: 'creative', type: 'public' },
|
||||
{ name: 'strict', id: 'strict', type: 'public' },
|
||||
],
|
||||
chunks: JSON.parse(chunks ?? '2').toString(),
|
||||
chunks: JSON.parse(chunks ?? '6').toString(),
|
||||
selectedDocs: getStoredRecentDocs(),
|
||||
conversations: {
|
||||
data: null,
|
||||
|
||||
@@ -0,0 +1,68 @@
|
||||
"""``run_agent_headless`` passes the agent's ``chunks`` to the retriever.
|
||||
|
||||
``chunks=0`` switches retrieval off, but ``int(... or 2)`` read it as unset,
|
||||
so a scheduled run of an agent with retrieval off retrieved anyway.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _retriever_chunks(agent_config, monkeypatch):
|
||||
"""The ``chunks`` a headless run hands to the retriever."""
|
||||
from docsgpt.agents import headless_runner as hr
|
||||
|
||||
agent = MagicMock(name="agent")
|
||||
agent.gen.return_value = iter([{"answer": "ok"}])
|
||||
agent.llm.token_usage = {"prompt_tokens": 1, "generated_tokens": 1}
|
||||
|
||||
retriever = MagicMock(name="retriever")
|
||||
retriever.search.return_value = []
|
||||
created = {}
|
||||
|
||||
def create_retriever(cls, *args, **kwargs):
|
||||
created.update(kwargs)
|
||||
return retriever
|
||||
|
||||
tool_executor = MagicMock(name="tool_executor")
|
||||
tool_executor.headless_denials = []
|
||||
|
||||
monkeypatch.setattr(hr, "get_prompt", lambda _pid: "system prompt")
|
||||
monkeypatch.setattr(
|
||||
hr.RetrieverCreator, "create_retriever", classmethod(create_retriever),
|
||||
)
|
||||
monkeypatch.setattr(hr, "ToolExecutor", lambda *a, **kw: tool_executor)
|
||||
monkeypatch.setattr(
|
||||
hr.AgentCreator, "create_agent",
|
||||
classmethod(lambda cls, *a, **kw: agent),
|
||||
)
|
||||
|
||||
config = {"user_id": "u1", "id": "agent-1", "default_model_id": "m", **agent_config}
|
||||
with patch("docsgpt.core.model_utils.validate_model_id", return_value=True), \
|
||||
patch("docsgpt.core.model_utils.get_default_model_id", return_value="m"), \
|
||||
patch(
|
||||
"docsgpt.core.model_utils.get_provider_from_model_id",
|
||||
return_value="openai",
|
||||
), \
|
||||
patch("docsgpt.core.model_utils.get_api_key_for_provider", return_value="k"), \
|
||||
patch("docsgpt.utils.calculate_doc_token_budget", return_value=1000):
|
||||
hr.run_agent_headless(config, "do the thing")
|
||||
return created["chunks"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestHeadlessRunnerChunks:
|
||||
def test_unset_chunks_uses_the_default(self, monkeypatch):
|
||||
assert _retriever_chunks({}, monkeypatch) == 6
|
||||
|
||||
def test_null_chunks_uses_the_default(self, monkeypatch):
|
||||
assert _retriever_chunks({"chunks": None}, monkeypatch) == 6
|
||||
|
||||
def test_zero_chunks_keeps_retrieval_off(self, monkeypatch):
|
||||
assert _retriever_chunks({"chunks": 0}, monkeypatch) == 0
|
||||
|
||||
def test_explicit_chunks_is_kept(self, monkeypatch):
|
||||
assert _retriever_chunks({"chunks": 4}, monkeypatch) == 4
|
||||
@@ -108,7 +108,7 @@ class TestAgentNodeConfig:
|
||||
assert c.stream_to_user is True
|
||||
assert c.tools == []
|
||||
assert c.sources == []
|
||||
assert c.chunks == "2"
|
||||
assert c.chunks == "6"
|
||||
assert c.retriever == ""
|
||||
assert c.model_id is None
|
||||
assert c.json_schema is None
|
||||
|
||||
@@ -383,7 +383,7 @@ class TestBuildHelpers:
|
||||
def test_build_config_defaults(self):
|
||||
config = build_internal_tool_config(source={"active_docs": ["abc"]})
|
||||
assert config["retriever_name"] == "classic"
|
||||
assert config["chunks"] == 2
|
||||
assert config["chunks"] == 6
|
||||
assert config["doc_token_limit"] == 50000
|
||||
|
||||
def test_internal_tool_id(self):
|
||||
|
||||
@@ -473,7 +473,7 @@ class TestConfigureRetriever:
|
||||
sp = StreamProcessor({}, {"sub": "u"})
|
||||
sp._configure_retriever()
|
||||
assert sp.retriever_config["retriever_name"] == "classic"
|
||||
assert sp.retriever_config["chunks"] == 2
|
||||
assert sp.retriever_config["chunks"] == 6
|
||||
|
||||
def test_agent_overrides(self):
|
||||
from docsgpt.api.answer.services.stream_processor import (
|
||||
@@ -519,7 +519,7 @@ class TestConfigureRetriever:
|
||||
sp._agent_data = {}
|
||||
sp._configure_retriever()
|
||||
assert sp.retriever_config["retriever_name"] == "classic"
|
||||
assert sp.retriever_config["chunks"] == 2
|
||||
assert sp.retriever_config["chunks"] == 6
|
||||
|
||||
def test_invalid_agent_chunks_falls_back(self):
|
||||
from docsgpt.api.answer.services.stream_processor import (
|
||||
@@ -528,7 +528,7 @@ class TestConfigureRetriever:
|
||||
sp = StreamProcessor({}, {"sub": "u"})
|
||||
sp._agent_data = {"chunks": "not-a-number"}
|
||||
sp._configure_retriever()
|
||||
assert sp.retriever_config["chunks"] == 2
|
||||
assert sp.retriever_config["chunks"] == 6
|
||||
|
||||
def test_invalid_request_chunks_falls_back(self):
|
||||
from docsgpt.api.answer.services.stream_processor import (
|
||||
@@ -536,7 +536,7 @@ class TestConfigureRetriever:
|
||||
)
|
||||
sp = StreamProcessor({"chunks": "abc"}, {"sub": "u"})
|
||||
sp._configure_retriever()
|
||||
assert sp.retriever_config["chunks"] == 2
|
||||
assert sp.retriever_config["chunks"] == 6
|
||||
|
||||
def test_isnonedoc_without_api_key_sets_chunks_to_0(self):
|
||||
from docsgpt.api.answer.services.stream_processor import (
|
||||
|
||||
@@ -126,6 +126,15 @@ class TestChunksPrecedence:
|
||||
sp._configure_retriever()
|
||||
return sp.retriever_config["chunks"]
|
||||
|
||||
def test_default_applies_when_nothing_sets_chunks(self):
|
||||
assert self._sp() == 6
|
||||
|
||||
def test_source_stored_at_the_old_default_counts_as_configured(self):
|
||||
"""Accepted when the default moved from 2 to 6: sources created while
|
||||
the UI always sent the full config store ``chunks=2`` explicitly, and
|
||||
stay a per-source setting of 2 rather than being migrated."""
|
||||
assert self._sp(request_chunks="7", source_chunks=2) == 2
|
||||
|
||||
def test_request_applies_when_source_is_unconfigured(self):
|
||||
assert self._sp(request_chunks="7") == 7
|
||||
|
||||
@@ -139,7 +148,7 @@ class TestChunksPrecedence:
|
||||
("-5", 0),
|
||||
("100000", 500),
|
||||
("501", 500),
|
||||
("abc", 2),
|
||||
("abc", 6),
|
||||
],
|
||||
)
|
||||
def test_request_chunks_is_clamped(self, sent, expected):
|
||||
|
||||
@@ -1202,7 +1202,7 @@ class TestConfigureAgent:
|
||||
assert sp.retriever_config["chunks"] == 5
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_configure_agent_invalid_chunks_defaults_to_2(self):
|
||||
def test_configure_agent_invalid_chunks_uses_the_default(self):
|
||||
sp = self._make_sp()
|
||||
sp._resolve_agent_id = MagicMock(return_value="agent_id_1")
|
||||
sp._get_agent_key = MagicMock(return_value=("agent_key", False, None))
|
||||
@@ -1220,7 +1220,7 @@ class TestConfigureAgent:
|
||||
sp._configure_agent()
|
||||
sp.model_id = "test-model"
|
||||
sp._configure_retriever()
|
||||
assert sp.retriever_config["chunks"] == 2
|
||||
assert sp.retriever_config["chunks"] == 6
|
||||
|
||||
|
||||
# ---- Additional coverage: _load_conversation_history ----
|
||||
|
||||
@@ -176,6 +176,27 @@ def test_import_by_slug_idempotent(pg_conn):
|
||||
assert len(AgentsRepository(pg_conn).list_for_user(user)) == 1
|
||||
|
||||
|
||||
def test_import_without_chunks_uses_the_default(pg_conn):
|
||||
user = "u_default_chunks"
|
||||
result = apply_import(pg_conn, user, _doc(name="Plain Bot", _slug="plain-bot"))
|
||||
agent = AgentsRepository(pg_conn).get(result["agent_id"], user)
|
||||
assert agent["chunks"] == 6
|
||||
|
||||
|
||||
def test_import_with_invalid_chunks_uses_the_default(pg_conn):
|
||||
user = "u_invalid_chunks"
|
||||
result = apply_import(pg_conn, user, _doc(name="Odd Bot", _slug="odd-bot", chunks="lots"))
|
||||
agent = AgentsRepository(pg_conn).get(result["agent_id"], user)
|
||||
assert agent["chunks"] == 6
|
||||
|
||||
|
||||
def test_import_keeps_explicit_chunks(pg_conn):
|
||||
user = "u_explicit_chunks"
|
||||
result = apply_import(pg_conn, user, _doc(name="Off Bot", _slug="off-bot", chunks=0))
|
||||
agent = AgentsRepository(pg_conn).get(result["agent_id"], user)
|
||||
assert agent["chunks"] == 0
|
||||
|
||||
|
||||
def test_import_missing_source_drafts_and_warns(pg_conn):
|
||||
user = "u_missing"
|
||||
doc = _doc(
|
||||
|
||||
@@ -426,6 +426,31 @@ class TestCreateAgent:
|
||||
agents = AgentsRepository(pg_conn).list_for_user(user)
|
||||
assert any(a["name"] == "My Draft Agent" for a in agents)
|
||||
|
||||
def test_agent_created_without_chunks_stores_the_default(self, app, pg_conn):
|
||||
# With two sources, 2 left each one a single chunk, so the answer hung
|
||||
# on which chunk won; 6 is the default now.
|
||||
from docsgpt.api.user.agents.routes import CreateAgent
|
||||
from docsgpt.storage.db.repositories.agents import AgentsRepository
|
||||
|
||||
user = "u-create-default-chunks"
|
||||
|
||||
with _patch_db(pg_conn), app.test_request_context(
|
||||
"/api/create_agent",
|
||||
method="POST",
|
||||
json={
|
||||
"name": "Handbook Bot",
|
||||
"description": "d",
|
||||
"agent_type": "classic",
|
||||
"status": "draft",
|
||||
},
|
||||
):
|
||||
from flask import request
|
||||
request.decoded_token = {"sub": user}
|
||||
response = CreateAgent().post()
|
||||
assert response.status_code == 201
|
||||
(agent,) = AgentsRepository(pg_conn).list_for_user(user)
|
||||
assert agent["chunks"] == 6
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# UpdateAgent — big method with many validation branches
|
||||
@@ -623,6 +648,23 @@ class TestUpdateAgent:
|
||||
field in msg and user in msg for msg in warnings
|
||||
), f"no WARN naming field={field!r} and user={user!r}; got {warnings!r}"
|
||||
|
||||
def test_blank_chunks_resets_to_the_default(self, app, pg_conn):
|
||||
from docsgpt.api.user.agents.routes import UpdateAgent
|
||||
from docsgpt.storage.db.repositories.agents import AgentsRepository
|
||||
|
||||
user = "u-upd-blank-chunks"
|
||||
agent = _seed_agent(pg_conn, user=user)
|
||||
with _patch_db(pg_conn), app.test_request_context(
|
||||
f"/api/update_agent/{agent['id']}",
|
||||
method="PUT",
|
||||
json={"name": "n", "description": "d", "status": "draft", "chunks": ""},
|
||||
):
|
||||
from flask import request
|
||||
request.decoded_token = {"sub": user}
|
||||
response = UpdateAgent().put(str(agent["id"]))
|
||||
assert response.status_code == 200
|
||||
assert AgentsRepository(pg_conn).get(str(agent["id"]), user)["chunks"] == 6
|
||||
|
||||
def test_invalid_chunks_returns_400(self, app, pg_conn):
|
||||
from docsgpt.api.user.agents.routes import UpdateAgent
|
||||
|
||||
|
||||
@@ -40,7 +40,7 @@ class TestParseLenient:
|
||||
r = RetrievalConfig()
|
||||
assert r.retriever == "classic"
|
||||
assert r.exposure == "prefetch"
|
||||
assert r.chunks == 2
|
||||
assert r.chunks == 6
|
||||
assert r.score_threshold is None
|
||||
assert r.rephrase_query is True
|
||||
assert r.reranker is None
|
||||
|
||||
@@ -74,6 +74,22 @@ class TestDispatcherGrouping:
|
||||
assert "b" not in retrievals
|
||||
|
||||
|
||||
def test_source_stored_at_the_old_default_is_an_override(self, _patch_llm_creator):
|
||||
# Accepted when the default moved from 2 to 6: sources saved with the
|
||||
# full config store chunks=2 and keep it as their own setting, while a
|
||||
# new upload stores 6, the default, and takes the global path.
|
||||
sources = [
|
||||
{"id": "old", "retrieval": RetrievalConfig(chunks=2)},
|
||||
{"id": "new", "retrieval": RetrievalConfig(chunks=6)},
|
||||
]
|
||||
d = Dispatcher(source={"question": "q", "active_docs": ["old", "new"]}, sources=sources)
|
||||
retrievals = d._groups[0]["retrievals"]
|
||||
assert retrievals["old"].chunks == 2
|
||||
assert "new" not in retrievals
|
||||
|
||||
def test_default_budget_is_six(self, _patch_llm_creator):
|
||||
assert Dispatcher(source={"question": "q", "active_docs": ["a"]}).chunks == 6
|
||||
|
||||
def test_graph_options_count_as_an_override(self, _patch_llm_creator):
|
||||
"""A graph source that changes only its graph options still needs its
|
||||
config carried over: those options live on the per-source retrieval the
|
||||
|
||||
@@ -151,7 +151,7 @@ class TestClassicRAGInit:
|
||||
|
||||
def test_chunks_invalid_string_defaults(self, _patch_llm_creator):
|
||||
rag = _make_rag(chunks="abc")
|
||||
assert rag.chunks == 2
|
||||
assert rag.chunks == 6
|
||||
|
||||
def test_decoded_token_none(self, _patch_llm_creator):
|
||||
rag = _make_rag(decoded_token=None)
|
||||
|
||||
Reference in new issue
Block a user