diff --git a/private_gpt/components/context/models/context_stack.py b/private_gpt/components/context/models/context_stack.py
index 517bb35b..73c2b9c0 100644
--- a/private_gpt/components/context/models/context_stack.py
+++ b/private_gpt/components/context/models/context_stack.py
@@ -47,60 +47,18 @@ class ContextStack(BaseModel):
# ------------------------------------------------------------------
def to_system_prompt(self) -> list[TextBlock]:
- """Render prompt layers by priority (then insertion order).
-
- Layers stay isolated in the stack; deduplication happens only here, at
- render time. Two cases must be covered so the LLM never receives the
- same text twice:
-
- - *Stale duplicate*: a layer's text is reproduced verbatim by another
- kept layer (e.g. the chat header rendered both as a
- ``RuntimeInstructionsLayer`` and baked back into a re-ingested
- ``UserInstructionsLayer``). The later candidate is dropped because a
- kept block already contains it.
- - *Snowballed aggregate*: a re-ingested ``UserInstructionsLayer`` may
- embed the header, platform guidelines and even prior skill bodies
- that the interceptors regenerate, this iteration, as isolated
- layers. The freshly-generated isolated layers carry the latest
- state (e.g. the ``ContextPromptLayer`` rebuilt from the latest
- documents accumulated across iterations), so they must survive.
- The aggregate layer is the one discarded: its rendered text *is a
- superset of* (contains) one or more kept isolated blocks.
-
- Rendering therefore walks layers in *descending* priority order — the
- freshly-generated isolated layers (highest priority numbers) are kept
- first — and a candidate is dropped when its text contains, or is
- contained in, the text of any already-kept block. That preserves the
- latest isolated layers across iterations and discards the duplicate /
- stale aggregate layer at render, restoring the prompt to its original
- isolated shape.
- """
+ """Render prompt layers by priority (then insertion order)."""
ordered_layers = sorted(
enumerate(self.layers),
- key=lambda item: (-item[1].priority, item[0]),
+ key=lambda item: (item[1].priority, item[0]),
)
- blocks: list[TextBlock] = []
- kept: list[str] = []
- for _, layer in ordered_layers:
- rendered = layer.render()
- if not rendered or not rendered.strip():
- continue
- text = rendered.strip()
- # Drop the candidate if its text already appears inside a kept
- # block (stale duplicate) OR if it reproduces, as a snowballed
- # aggregate, any kept isolated block. The freshly-generated
- # isolated layers are kept first thanks to the descending-priority
- # order, so the aggregate (e.g. a re-ingested UserInstructions
- # layer built from a previous response) is the one discarded,
- # while the latest isolated layers — including the ContextPrompt
- # rebuilt with the latest documents — survive.
- is_stale_duplicate = any(text in candidate for candidate in kept)
- is_snowballed_aggregate = any(candidate in text for candidate in kept)
- if is_stale_duplicate or is_snowballed_aggregate:
- continue
- blocks.append(TextBlock(text=text))
- kept.append(text)
- return blocks
+ chunks = [
+ rendered
+ for _, layer in ordered_layers
+ for rendered in [layer.render()]
+ if rendered.strip()
+ ]
+ return [TextBlock(text=chunk) for chunk in chunks]
def all_tools(self) -> list[ToolSpec]:
"""Return deduplicated ToolSpec list from all TOOL_DEFINITIONS layers."""
@@ -150,30 +108,17 @@ class ContextStack(BaseModel):
# ------------------------------------------------------------------
# Immutable mutation helpers
# ------------------------------------------------------------------
- def _remove_duplicate(self, layer: AnyContextLayer) -> "list[AnyContextLayer]":
- """Return current layers with any existing (type, source) duplicate removed."""
- return [
- existing
- for existing in self.layers
- if not (existing.type == layer.type and existing.source == layer.source)
- ]
-
def insert_layer(self, layer: AnyContextLayer, index: int) -> "ContextStack":
"""Return a new stack with *layer* inserted at *index* (default 0)."""
- base = self._remove_duplicate(layer)
- index = min(index, len(base))
- return ContextStack(layers=[*base[:index], layer, *base[index:]])
+ return ContextStack(layers=[*self.layers[:index], layer, *self.layers[index:]])
def append_layer(self, layer: AnyContextLayer) -> "ContextStack":
"""Return a new stack with *layer* appended."""
- return ContextStack(layers=[*self._remove_duplicate(layer), layer])
+ return ContextStack(layers=[*self.layers, layer])
def append_layers(self, layers: list[AnyContextLayer]) -> "ContextStack":
"""Return a new stack with *layers* appended."""
- stack = self
- for layer in layers:
- stack = stack.append_layer(layer)
- return stack
+ return ContextStack(layers=[*self.layers, *layers])
def remove_layers_of_type(self, layer_type: LayerType) -> "ContextStack":
"""Return a new stack with all layers of *layer_type* removed."""
diff --git a/private_gpt/components/engines/chat/async_chat_engine.py b/private_gpt/components/engines/chat/async_chat_engine.py
index 035e69a7..37f14276 100644
--- a/private_gpt/components/engines/chat/async_chat_engine.py
+++ b/private_gpt/components/engines/chat/async_chat_engine.py
@@ -172,6 +172,7 @@ class AsyncChatCheckpoint(BaseModel):
payload: IterationCheckpointPayload = Field(
default_factory=IterationCheckpointPayload
)
+ original_input: ChatInputState | None = None
@dataclass
@@ -182,6 +183,7 @@ class _CheckpointContext:
payload: IterationCheckpointPayload = field(
default_factory=IterationCheckpointPayload
)
+ original_input: ChatInputState | None = None
@dataclass
@@ -371,6 +373,7 @@ class AsyncChatEngine:
IterationCheckpointPayload(),
hooks,
channel,
+ original_input=state.original_input,
)
async def resume(
@@ -385,6 +388,7 @@ class AsyncChatEngine:
checkpoint=checkpoint.checkpoint,
payload=checkpoint.payload,
context_stack=checkpoint.input.context_stack,
+ original_input=checkpoint.original_input,
)
handler = self._resolve_checkpoint_handler(checkpoint.checkpoint)
state = await handler(
@@ -411,6 +415,7 @@ class AsyncChatEngine:
new_payload,
hooks,
channel,
+ original_input=state.original_input,
)
if state.output.pending_external_tool_calls:
close_state = close_state.model_copy(deep=True)
@@ -425,6 +430,7 @@ class AsyncChatEngine:
new_payload,
hooks,
channel,
+ original_input=state.original_input,
)
async def run(
@@ -539,10 +545,12 @@ class AsyncChatEngine:
payload: IterationCheckpointPayload,
hooks: ExecutionHooks | None,
channel: EventChannel,
+ original_input: ChatInputState | None = None,
) -> ChatState:
context = self._build_checkpoint_context(
checkpoint=_IterationCheckpoint.BEFORE_ITERATION,
payload=payload,
+ original_input=original_input,
)
return await self._execute_before_iteration_checkpoint(
request, iteration, next_block_count, channel, hooks, context
@@ -555,11 +563,13 @@ class AsyncChatEngine:
payload: IterationCheckpointPayload,
hooks: ExecutionHooks | None,
channel: EventChannel,
+ original_input: ChatInputState | None = None,
) -> ChatState:
context = self._build_checkpoint_context(
checkpoint=_IterationCheckpoint.CLOSE,
stop_reason=stop_reason,
payload=payload,
+ original_input=original_input,
)
return await self._execute_close_checkpoint(
request, 0, 0, channel, hooks, context
@@ -573,15 +583,31 @@ class AsyncChatEngine:
payload: IterationCheckpointPayload,
hooks: ExecutionHooks | None,
channel: EventChannel,
+ original_input: ChatInputState | None = None,
) -> ChatState:
while True:
if iteration >= self._max_iterations:
return await self._execute_close(
- request, StopReasonEnum.MAX_TOKENS.value, payload, hooks, channel
+ request,
+ StopReasonEnum.MAX_TOKENS.value,
+ payload,
+ hooks,
+ channel,
+ original_input=original_input,
)
state = await self._execute_before_iteration(
- request, iteration, next_block_count, payload, hooks, channel
+ request,
+ iteration,
+ next_block_count,
+ payload,
+ hooks,
+ channel,
+ original_input=original_input,
)
+ # Keep the first-iteration snapshot for the whole loop. Later
+ # checkpoints rebuild run state from the materialized request, so
+ # re-snapshotting here would poison restore/system-prompt rebuilds.
+ original_input = state.original_input or original_input
new_payload = IterationCheckpointPayload(
model_id=state.runtime.model_id,
total_input_tokens=state.runtime.total_input_tokens,
@@ -598,6 +624,7 @@ class AsyncChatEngine:
new_payload,
hooks,
channel,
+ original_input=original_input,
)
if state.output.pending_external_tool_calls:
close_state = close_state.model_copy(deep=True)
@@ -621,12 +648,14 @@ class AsyncChatEngine:
stop_reason: str | None = None,
payload: IterationCheckpointPayload | None = None,
context_stack: ContextStack | None = None,
+ original_input: ChatInputState | None = None,
) -> _CheckpointContext:
return _CheckpointContext(
checkpoint=checkpoint,
stop_reason=stop_reason,
context_stack=context_stack,
payload=payload or IterationCheckpointPayload(),
+ original_input=original_input,
)
def _resolve_checkpoint_handler(
@@ -669,7 +698,10 @@ class AsyncChatEngine:
checkpoint_context: _CheckpointContext,
) -> ChatState:
run = self._initialize_run(
- request, context_stack=checkpoint_context.context_stack, hooks=hooks
+ request,
+ context_stack=checkpoint_context.context_stack,
+ hooks=hooks,
+ original_input=checkpoint_context.original_input,
)
run.state.runtime.iteration = iteration
run.state.runtime.next_block_count = next_block_count
@@ -692,7 +724,10 @@ class AsyncChatEngine:
) -> ChatState:
del iteration, next_block_count
run = self._initialize_run(
- request, context_stack=checkpoint_context.context_stack, hooks=hooks
+ request,
+ context_stack=checkpoint_context.context_stack,
+ hooks=hooks,
+ original_input=checkpoint_context.original_input,
)
channel.emit(RawMessageStartEvent.from_defaults())
await channel.flush()
@@ -725,6 +760,7 @@ class AsyncChatEngine:
context_stack=checkpoint_context.context_stack,
hooks=hooks,
checkpoint_payload=checkpoint_context.payload,
+ original_input=checkpoint_context.original_input,
)
async def _continue_tools_checkpoint(
@@ -736,11 +772,15 @@ class AsyncChatEngine:
context_stack: ContextStack | None = None,
hooks: ExecutionHooks | None = None,
checkpoint_payload: IterationCheckpointPayload | None = None,
+ original_input: ChatInputState | None = None,
) -> ChatState:
"""Continue from the tools checkpoint without re-running the LLM call."""
try:
run = self._initialize_run(
- request, context_stack=context_stack, hooks=hooks
+ request,
+ context_stack=context_stack,
+ hooks=hooks,
+ original_input=original_input,
)
run.state.runtime.iteration = iteration
run.state.runtime.next_block_count = next_block_count
@@ -843,7 +883,10 @@ class AsyncChatEngine:
) -> ChatState:
del iteration, next_block_count
run = self._initialize_run(
- request, context_stack=checkpoint_context.context_stack, hooks=hooks
+ request,
+ context_stack=checkpoint_context.context_stack,
+ hooks=hooks,
+ original_input=checkpoint_context.original_input,
)
self._apply_payload_usage(run, checkpoint_context.payload)
stop_reason = checkpoint_context.stop_reason
@@ -872,8 +915,14 @@ class AsyncChatEngine:
request: ChatRequest,
context_stack: ContextStack | None = None,
hooks: ExecutionHooks | None = None,
+ original_input: ChatInputState | None = None,
) -> _Run:
- return self._initialize_run(request, context_stack=context_stack, hooks=hooks)
+ return self._initialize_run(
+ request,
+ context_stack=context_stack,
+ hooks=hooks,
+ original_input=original_input,
+ )
async def _pipe_events_through_interceptors(
self,
@@ -1547,6 +1596,7 @@ class AsyncChatEngine:
request: ChatRequest,
context_stack: ContextStack | None = None,
hooks: ExecutionHooks | None = None,
+ original_input: ChatInputState | None = None,
) -> _Run:
llm = self._llm_component.get_llm(request.system.model)
if not isinstance(llm, FunctionCallingLLM):
@@ -1581,7 +1631,14 @@ class AsyncChatEngine:
output=ChatOutputState(),
timeline=[],
)
- state.original_input = state.input.model_copy(deep=True)
+ # original_input is a first-iteration snapshot. Rebuilding it from later
+ # checkpoint request materializations would make restore/system-prompt
+ # interceptors treat intermediate prompt state as the user original.
+ state.original_input = (
+ original_input
+ if original_input is not None
+ else state.input.model_copy(deep=True)
+ )
run = _Run(
state=state,
llm=llm,
diff --git a/private_gpt/components/engines/chat/checkpoint_store.py b/private_gpt/components/engines/chat/checkpoint_store.py
index 3c7ddddb..4f8fdbcd 100644
--- a/private_gpt/components/engines/chat/checkpoint_store.py
+++ b/private_gpt/components/engines/chat/checkpoint_store.py
@@ -32,6 +32,7 @@ class ChatCheckpoint(BaseModel):
correlation_id: str
request_data: dict[str, Any]
context_stack_data: dict[str, Any] = Field(default_factory=dict)
+ original_input_data: dict[str, Any] | None = None
stream_type: str
metadata: dict[str, Any]
iteration: int
diff --git a/private_gpt/components/engines/chat/resumable_runner.py b/private_gpt/components/engines/chat/resumable_runner.py
index 663cc81c..4982d80a 100644
--- a/private_gpt/components/engines/chat/resumable_runner.py
+++ b/private_gpt/components/engines/chat/resumable_runner.py
@@ -192,6 +192,7 @@ class ResumableChatRunner:
payload=saved.checkpoint_payload.model_copy(
update={"tool_responses": responses}
),
+ original_input=self._original_input(saved),
),
hooks=_RESUME_HOOKS,
channel=channel,
@@ -249,6 +250,7 @@ class ResumableChatRunner:
correlation_id=execution_id,
request_data=state.input.request.model_dump(mode="json"),
context_stack_data=state.input.context_stack.checkpoint_dump(),
+ original_input_data=self._dump_original_input(state.original_input),
stream_type=stream_type,
metadata=metadata,
iteration=state.runtime.iteration,
@@ -386,6 +388,27 @@ class ResumableChatRunner:
}
return request
+ @staticmethod
+ def _dump_original_input(
+ original_input: ChatInputState | None,
+ ) -> dict[str, Any] | None:
+ if original_input is None or not isinstance(original_input, ChatInputState):
+ return None
+ return original_input.model_dump(mode="json")
+
+ @staticmethod
+ def _original_input(checkpoint: ChatCheckpoint) -> ChatInputState | None:
+ if not checkpoint.original_input_data:
+ return None
+ data = dict(checkpoint.original_input_data)
+ request_data = data.get("request")
+ if isinstance(request_data, dict):
+ data["request"] = ResumableChatRunner._request(request_data)
+ context_stack_data = data.get("context_stack")
+ if isinstance(context_stack_data, dict):
+ data["context_stack"] = ContextStack.model_validate(context_stack_data)
+ return ChatInputState.model_validate(data)
+
@staticmethod
def _context_stack(
checkpoint: ChatCheckpoint, request_data: dict[str, Any]
diff --git a/private_gpt/components/engines/chat/utils/request_builder.py b/private_gpt/components/engines/chat/utils/request_builder.py
index c10cf2c2..db86a307 100644
--- a/private_gpt/components/engines/chat/utils/request_builder.py
+++ b/private_gpt/components/engines/chat/utils/request_builder.py
@@ -10,6 +10,7 @@ from private_gpt.components.context.models.context_layer import (
UserInstructionsLayer,
)
from private_gpt.components.context.models.context_stack import ContextStack
+from private_gpt.components.context.models.layer_type import LayerType
def build_initial_context_stack(
@@ -23,11 +24,13 @@ def build_initial_context_stack(
# in the context stack if they are present in the request.
if request.system.prompt:
+ stack = stack.remove_layers_of_type(LayerType.USER_INSTRUCTIONS)
stack = stack.append_layer(
UserInstructionsLayer(text=request.system.prompt, source=source)
)
if request.tool_config.tools:
+ stack = stack.remove_layers_of_type(LayerType.TOOL_DEFINITIONS)
stack = stack.append_layer(
ToolDefinitionsLayer(
tools=list(request.tool_config.tools),
@@ -36,6 +39,7 @@ def build_initial_context_stack(
)
if request.context.documents:
+ stack = stack.remove_layers_of_type(LayerType.DOCUMENT)
for document in request.context.documents:
stack = stack.append_layer(
DocumentLayer(document=document, source=source)
diff --git a/tests/components/context/test_context_stack_render.py b/tests/components/context/test_context_stack_render.py
deleted file mode 100644
index 77121b04..00000000
--- a/tests/components/context/test_context_stack_render.py
+++ /dev/null
@@ -1,137 +0,0 @@
-"""Render-time deduplication of ``ContextStack.to_system_prompt``.
-
-Layers stay isolated in the stack; deduplication happens only at render. The
-contract enforced here:
-
-- The same text never reaches the LLM twice (stale duplicate drop).
-- A re-ingested ``UserInstructionsLayer`` that aggregates a previous response
- (header + guidelines) is discarded once the freshly-generated isolated
- layers that the interceptors rebuilt for this iteration already reproduce
- its parts (snowballed-aggregate drop).
-- The freshly-generated isolated layers, including the ``ContextPromptLayer``
- rebuilt with the latest documents accumulated across iterations, must
- survive over any stale aggregate that embeds an older version of them.
-"""
-
-from private_gpt.components.context.models.context_layer import (
- ContextPromptLayer,
- RuntimeInstructionsLayer,
- SkillBodyLayer,
- ToolInstructionsLayer,
- UserInstructionsLayer,
-)
-from private_gpt.components.context.models.context_stack import ContextStack
-
-
-def _render(stack: ContextStack) -> list[str]:
- return [b.text for b in stack.to_system_prompt() if b.text]
-
-
-class TestRenderTimeDeduplication:
- def test_identical_layers_collapse_to_one_block(self) -> None:
- text = "You are Zylon, an AI assistant.\nCurrent date: 2026-07-28."
- stack = ContextStack(
- layers=[
- UserInstructionsLayer(text=text, source="request"),
- RuntimeInstructionsLayer(text=text, source="platform_header"),
- ]
- )
-
- rendered = _render(stack)
-
- assert rendered == [text]
-
- def test_stale_duplicate_lower_priority_is_dropped(self) -> None:
- """Runtime header is kept; aggregate UserInstructions is dropped."""
- header = "You are Zylon, an AI assistant."
- guideline = "\nWrite clearly.\n"
- bloated = f"{header}\n\n{guideline}"
- stack = ContextStack(
- layers=[
- UserInstructionsLayer(text=bloated, source="request"),
- RuntimeInstructionsLayer(text=header, source="platform_header"),
- ToolInstructionsLayer(
- tool_name="response_formatting",
- instructions=guideline,
- source="platform:tool_instructions",
- ),
- ]
- )
-
- rendered = _render(stack)
- assert rendered.count(header) == 1
- assert rendered.count(guideline) == 1
- # The bloated aggregate must NOT survive: only the isolated layers do
- assert bloated not in rendered
- assert header in rendered
- assert guideline in rendered
-
- def test_latest_context_prompt_survives_stale_aggregate(self) -> None:
- """A re-ingested UserInstructions embedding an *older* rendered
- context prompt is discarded; the fresh ContextPromptLayer (rebuilt
- from the latest documents accumulated across iterations) survives.
- """
- stale_ctx = "\nold content\n"
- fresh_ctx = "\nlatest content\n"
- # Aggregate layer reproduces the *stale* version of the context.
- bloated = f"You are Zylon.\n\n{stale_ctx}\n\n..."
- stack = ContextStack(
- layers=[
- UserInstructionsLayer(text=bloated, source="request"),
- RuntimeInstructionsLayer(
- text="You are Zylon.", source="platform_header"
- ),
- ContextPromptLayer(text=fresh_ctx, source="system_prompt"),
- ]
- )
-
- rendered = _render(stack)
- assert fresh_ctx in rendered, "Latest context prompt must survive"
- assert stale_ctx not in rendered, "Stale context prompt must be dropped"
- assert bloated not in rendered, "Snowballed aggregate must be dropped"
-
- def test_distinct_isolated_layers_kept(self) -> None:
- """No false positives: layers with non-overlapping content are kept."""
- header = "You are Zylon."
- guideline = "\nWrite clearly.\n"
- # SkillBodyLayer wraps the instructions in
- skill_body = '\nbody\n'
- stack = ContextStack(
- layers=[
- RuntimeInstructionsLayer(text=header, source="platform_header"),
- ToolInstructionsLayer(
- tool_name="response_formatting",
- instructions=guideline,
- source="platform:tool_instructions",
- ),
- SkillBodyLayer(
- skill_id="x",
- name="x",
- version="1",
- instructions="body",
- source="skill:x",
- ),
- ]
- )
-
- rendered = _render(stack)
-
- assert set(rendered) == {header, guideline, skill_body}
-
- def test_render_is_idempotent_across_iterations(self) -> None:
- """Simulate two iterations where the same starter stack is rendered
- twice — output must not grow with repeated calls.
- """
- header = "You are Zylon."
- stack = ContextStack(
- layers=[
- UserInstructionsLayer(text=f"{header}\n", source="request"),
- RuntimeInstructionsLayer(text=header, source="platform_header"),
- ]
- )
-
- first = _render(stack)
- second = _render(ContextStack(layers=list(stack.layers)))
-
- assert first == second
- assert first.count(header) == 1
diff --git a/tests/engines/test_async_chat_engine.py b/tests/engines/test_async_chat_engine.py
index 02adb4ff..57a1dba4 100644
--- a/tests/engines/test_async_chat_engine.py
+++ b/tests/engines/test_async_chat_engine.py
@@ -329,6 +329,7 @@ async def _run_async_engine(
has_input_usage=state.runtime.has_input_usage,
has_output_usage=state.runtime.has_output_usage,
),
+ original_input=state.original_input,
),
channel=channel2,
)
@@ -869,6 +870,7 @@ async def test_extract_citation_interceptor_converts_bracket_refs_on_resume(
has_input_usage=state.runtime.has_input_usage,
has_output_usage=state.runtime.has_output_usage,
),
+ original_input=state.original_input,
),
channel=channel2,
)
@@ -894,3 +896,110 @@ async def test_extract_citation_interceptor_converts_bracket_refs_on_resume(
assert " XML tag in output, got: {full_text!r}"
)
+
+
+@pytest.mark.asyncio
+async def test_initialize_run_reuses_provided_original_input() -> None:
+ """_initialize_run must not resnapshot original_input on later checkpoints."""
+ from llama_index.core.base.llms.types import TextBlock
+ from llama_index.core.llms.function_calling import FunctionCallingLLM
+
+ from private_gpt.components.context.models.context_layer import (
+ UserInstructionsLayer,
+ )
+ from private_gpt.components.context.models.context_stack import ContextStack
+ from private_gpt.components.context.models.layer_type import LayerType
+
+ class _FakeFunctionLLM(FunctionCallingLLM):
+ @property
+ def metadata(self):
+ return MagicMock(is_function_calling_model=True, context_window=8192)
+
+ def _prepare_chat_with_tools(self, *a, **k):
+ return {}
+
+ async def achat(self, *a, **k):
+ raise NotImplementedError
+
+ def chat(self, *a, **k):
+ raise NotImplementedError
+
+ def stream_chat(self, *a, **k):
+ raise NotImplementedError
+
+ async def astream_chat(self, *a, **k):
+ raise NotImplementedError
+
+ def complete(self, *a, **k):
+ raise NotImplementedError
+
+ async def acomplete(self, *a, **k):
+ raise NotImplementedError
+
+ def stream_complete(self, *a, **k):
+ raise NotImplementedError
+
+ async def astream_complete(self, *a, **k):
+ raise NotImplementedError
+
+ def chat_with_tools(self, *a, **k):
+ raise NotImplementedError
+
+ async def achat_with_tools(self, *a, **k):
+ raise NotImplementedError
+
+ def stream_chat_with_tools(self, *a, **k):
+ raise NotImplementedError
+
+ async def astream_chat_with_tools(self, *a, **k):
+ raise NotImplementedError
+
+ def get_tool_calls_from_response(self, *a, **k):
+ return []
+
+ first_request = ResolvedChatRequest(
+ messages=[ChatMessage(role=MessageRole.USER, content="hello")],
+ system=ResolvedSystemConfig(
+ model="default",
+ prompt=[TextBlock(text="USER PROMPT")],
+ ),
+ )
+ later_request = first_request.model_copy(deep=True)
+ later_request.system.prompt = [TextBlock(text="FULL RENDERED PROMPT")]
+ later_request.messages = [
+ *later_request.messages,
+ ChatMessage(role=MessageRole.ASSISTANT, content="tool-turn"),
+ ]
+
+ llm_component = MagicMock(spec=LLMComponent)
+ llm_component.get_llm.return_value = _FakeFunctionLLM()
+ engine = AsyncChatEngine(
+ llm_component=llm_component,
+ chat_scheduler=MagicMock(),
+ )
+
+ first_run = engine.initialize_run(first_request)
+ original = first_run.state.original_input
+ assert original is not None
+ first_layers = original.context_stack.layers_of_type(LayerType.USER_INSTRUCTIONS)
+ assert first_layers
+ assert first_layers[0].text == [TextBlock(text="USER PROMPT")]
+
+ second_run = engine.initialize_run(
+ later_request,
+ context_stack=ContextStack(
+ layers=[
+ UserInstructionsLayer(
+ text=[TextBlock(text="FULL RENDERED PROMPT")],
+ source="request",
+ )
+ ]
+ ),
+ original_input=original,
+ )
+ assert second_run.state.original_input is original
+ second_layers = second_run.state.original_input.context_stack.layers_of_type(
+ LayerType.USER_INSTRUCTIONS
+ )
+ assert second_layers
+ assert second_layers[0].text == [TextBlock(text="USER PROMPT")]
diff --git a/tests/server/chat/interceptors/test_system_prompt_interceptor.py b/tests/server/chat/interceptors/test_system_prompt_interceptor.py
deleted file mode 100644
index 256e7daf..00000000
--- a/tests/server/chat/interceptors/test_system_prompt_interceptor.py
+++ /dev/null
@@ -1,300 +0,0 @@
-"""Tests for SystemPromptRequestInterceptor layer deduplication.
-
-Verifies that running the interceptor N times (simulating tool-call loops or
-the recalculate branch) never accumulates duplicate layers in the context
-stack or duplicates text in the rendered system prompt.
-"""
-
-from unittest.mock import MagicMock
-
-import pytest
-from llama_index.core.base.llms.types import ChatMessage, MessageRole, TextBlock
-
-from private_gpt.components.chat.models.chat_config_models import (
- ResolvedChatRequest,
- ResolvedSystemConfig,
-)
-from private_gpt.components.context.models.context_stack import ContextStack
-from private_gpt.components.context.models.layer_type import LayerType
-from private_gpt.components.engines.chat.models.chat_interceptor_context import (
- ChatInterceptorContext,
-)
-from private_gpt.components.engines.chat.models.chat_phase import InterceptorPhase
-from private_gpt.components.engines.chat.models.chat_state import (
- ChatInputState,
- ChatOutputState,
- ChatRuntimeState,
- ChatState,
-)
-from private_gpt.components.engines.chat.utils.request_builder import (
- build_initial_context_stack,
-)
-from private_gpt.server.chat.interceptors.system_prompt_interceptor import (
- SystemPromptRequestInterceptor,
-)
-from tests.fixtures.mock_function_llm import get_mock_function_calling_llm
-
-_SYSTEM_PROMPT = "You are Zylon, an AI assistant.\nCurrent date: 2026-07-27."
-
-
-def _make_request(
- system_prompt: str | list[TextBlock] | None = _SYSTEM_PROMPT,
-) -> ResolvedChatRequest:
- return ResolvedChatRequest(
- messages=[ChatMessage(role=MessageRole.USER, content="hello")],
- system=ResolvedSystemConfig(prompt=system_prompt),
- )
-
-
-def _make_context(
- request: ResolvedChatRequest,
- context_stack: ContextStack | None = None,
- phase: InterceptorPhase = InterceptorPhase.BEFORE_ITERATION,
-) -> ChatInterceptorContext:
- stack = (
- context_stack
- if context_stack is not None
- else build_initial_context_stack(request)
- )
- state = ChatState(
- input=ChatInputState(
- request=request,
- context_stack=stack,
- ),
- runtime=ChatRuntimeState(),
- output=ChatOutputState(),
- timeline=[],
- )
- return ChatInterceptorContext(
- state=state,
- llm=get_mock_function_calling_llm(["ok"]),
- phase=phase,
- emit_fn=lambda _: None,
- )
-
-
-def _make_interceptor(
- add_context_to_system_prompt: bool = False,
-) -> SystemPromptRequestInterceptor:
- """Build a SystemPromptRequestInterceptor with a minimal PromptBuilderService."""
- prompt_builder = MagicMock()
- prompt_template = MagicMock()
- prompt_template.format.return_value = _SYSTEM_PROMPT
- prompt_builder.create_chat_header_prompt.return_value = prompt_template
- return SystemPromptRequestInterceptor(
- prompt_builder_service=prompt_builder,
- add_context_to_system_prompt=add_context_to_system_prompt,
- )
-
-
-class TestSystemPromptInterceptorIdempotency:
- """The interceptor must be idempotent across repeated calls."""
-
- @pytest.mark.asyncio
- async def test_single_run_produces_one_platform_header_layer(self) -> None:
- interceptor = _make_interceptor()
- request = _make_request()
- context = _make_context(request)
-
- await interceptor.intercept(context)
-
- platform_layers = [
- layer
- for layer in context.state.input.context_stack.layers
- if layer.source == "platform_header"
- ]
- assert len(platform_layers) == 1, "Expected exactly 1 platform_header layer"
-
- @pytest.mark.asyncio
- async def test_running_n_times_does_not_accumulate_layers(self) -> None:
- """Simulates multiple BEFORE_ITERATION passes (tool-call loop)."""
- interceptor = _make_interceptor()
- request = _make_request()
- context = _make_context(request)
-
- for _ in range(5):
- await interceptor.intercept(context)
-
- platform_layers = [
- layer
- for layer in context.state.input.context_stack.layers
- if layer.source == "platform_header"
- ]
- assert len(platform_layers) == 1, (
- f"After 5 iterations got {len(platform_layers)} platform_header layers — "
- "interceptor is accumulating duplicates!"
- )
-
- @pytest.mark.asyncio
- async def test_running_n_times_no_user_instruction_duplication(self) -> None:
- """Multiple BEFORE_ITERATION passes: user instructions appear once."""
- interceptor = _make_interceptor()
- request = _make_request()
- context = _make_context(request)
-
- for _ in range(5):
- await interceptor.intercept(context)
-
- user_layers = [
- layer
- for layer in context.state.input.context_stack.layers
- if layer.type == LayerType.USER_INSTRUCTIONS
- ]
- sources = [layer.source for layer in user_layers]
- assert sources.count("request") <= 1, (
- f"'request' USER_INSTRUCTIONS layer duplicated: {sources}"
- )
- assert sources.count("platform_header") <= 1, (
- f"'platform_header' USER_INSTRUCTIONS layer duplicated: {sources}"
- )
-
- @pytest.mark.asyncio
- async def test_system_prompt_text_not_duplicated_after_n_iterations(self) -> None:
- """The rendered system prompt text must not repeat after N runs."""
- interceptor = _make_interceptor()
- request = _make_request()
- context = _make_context(request)
-
- for _ in range(5):
- await interceptor.intercept(context)
-
- prompt = context.state.input.request.system.prompt
- # Normalise to list of text strings
- if isinstance(prompt, str):
- texts = [prompt]
- elif isinstance(prompt, list):
- texts = [b.text for b in prompt if isinstance(b, TextBlock) and b.text]
- else:
- texts = []
-
- full_text = "\n".join(texts)
- occurrences = full_text.count(_SYSTEM_PROMPT)
- assert (
- occurrences <= 2
- ), ( # at most 2: once in user layer, once in platform_header
- f"System prompt text appears {occurrences} times after 5 iterations. "
- "Likely a duplication bug!"
- )
-
- @pytest.mark.asyncio
- async def test_checkpoint_saves_original_not_mutated_prompt(self) -> None:
- """Verify checkpoint round-trip with original request avoids duplication.
-
- The ResumableChatRunner now saves ``state.original_input.request``.
- This test simulates the resume path where the original (clean) request
- is used together with ``build_initial_context_stack``.
- """
- interceptor = _make_interceptor()
- original_request = _make_request()
-
- # --- Simulate first request execution ---
- context = _make_context(original_request)
- await interceptor.intercept(context)
-
- # Verify _render_system_prompt_text returns a single TextBlock
- mutated_prompt = context.state.input.request.system.prompt
- if isinstance(mutated_prompt, list):
- assert len(mutated_prompt) == 1, (
- "_render_system_prompt_text should return exactly 1 TextBlock"
- )
-
- # --- Simulate resume: build fresh stack from ORIGINAL (clean) request ---
- restored_stack = build_initial_context_stack(original_request)
- restored_context = _make_context(original_request, context_stack=restored_stack)
- await interceptor.intercept(restored_context)
-
- prompt = restored_context.state.input.request.system.prompt
- if isinstance(prompt, str):
- full_text = prompt
- elif isinstance(prompt, list):
- full_text = "\n".join(
- b.text for b in prompt if isinstance(b, TextBlock) and b.text
- )
- else:
- full_text = ""
-
- occurrences = full_text.count(_SYSTEM_PROMPT)
- assert occurrences <= 2, (
- f"After resume the system prompt appears {occurrences} times."
- )
-
- @pytest.mark.asyncio
- async def test_round_tripped_prompt_does_not_duplicate_rendered_blocks(
- self,
- ) -> None:
- """A request whose ``system.prompt`` already carries the rendered
- stack (client echoes back a previous response) must not produce a
- system message whose content blocks repeat the header or guidelines
- that the interceptors regenerate as isolated layers.
-
- Layers stay isolated in the stack; at render time duplicate rendered
- text is discarded so the snowball never reaches the LLM.
- """
- header = "You are Zylon, an AI assistant.\nCurrent date: 2026-07-28."
- guideline = "\nWrite clearly.\n"
- # Client re-sends the fully rendered prompt (header + guideline baked in)
- bloated_prompt = f"{header}\n\n{guideline}"
-
- interceptor = _make_interceptor()
- # Header template renders the same header the client already embedded
- interceptor._prompt_builder_service.create_chat_header_prompt.return_value.format.return_value = header
- request = _make_request(system_prompt=bloated_prompt)
- context = _make_context(request)
- await interceptor.intercept(context)
-
- prompt = context.state.input.request.system.prompt
- if isinstance(prompt, str):
- texts = [prompt]
- elif isinstance(prompt, list):
- texts = [b.text for b in prompt if isinstance(b, TextBlock) and b.text]
- else:
- texts = []
-
- full = "\n".join(texts)
- # Header must appear at most once across all rendered blocks
- assert full.count(header) <= 1, (
- f"Header rendered {full.count(header)} times after round-trip: {texts!r}"
- )
- # Guideline must appear at most once
- assert full.count(guideline) <= 1, (
- f"Guideline rendered {full.count(guideline)} times: {texts!r}"
- )
- # No rendered block may fully duplicate another kept block
- stripped = [t.strip() for t in texts if t.strip()]
- for i, block in enumerate(stripped):
- for j, other in enumerate(stripped):
- if i != j and block and block in other:
- raise AssertionError(
- f"Rendered block #{i} is contained in block #{j} — "
- "duplicate content reached the rendered prompt."
- )
-
- @pytest.mark.asyncio
- async def test_fallback_build_with_mutated_prompt_is_safe(self) -> None:
- """Defensive: even if build_initial is called on a mutated request,
- the system prompt should not explode (snowball test)."""
- interceptor = _make_interceptor()
- request = _make_request()
- context = _make_context(request)
- await interceptor.intercept(context)
-
- # Simulate mutated request being re-ingested
- mutated_request = context.state.input.request
- restored_stack = build_initial_context_stack(mutated_request)
- restored_context = _make_context(mutated_request, context_stack=restored_stack)
- await interceptor.intercept(restored_context)
-
- prompt = restored_context.state.input.request.system.prompt
- if isinstance(prompt, str):
- full_text = prompt
- elif isinstance(prompt, list):
- full_text = "\n".join(
- b.text for b in prompt if isinstance(b, TextBlock) and b.text
- )
- else:
- full_text = ""
-
- occurrences = full_text.count(_SYSTEM_PROMPT)
- assert occurrences <= 3, (
- f"Even on fallback path, the system prompt explodes to {occurrences} occurrences."
- )