Merge branch 'pr-5126-head' into stage/brick-5096

# Conflicts:
#	api/session_ops.py
This commit is contained in:
nesquena-hermes
2026-06-28 18:14:03 +00:00
4 changed files with 113 additions and 18 deletions
+5 -18
View File
@@ -12227,24 +12227,9 @@ def handle_post(handler, parsed) -> bool:
if keep < 0:
return bad(handler, "keep_count must be non-negative")
with _get_session_agent_lock(body["session_id"]):
old_msg_count = len(s.messages or [])
old_ctx_count = len(getattr(s, 'context_messages', None) or [])
s.messages = s.messages[:keep]
# Truncate context_messages in sync with messages so the agent's
# model-facing context doesn't retain rows the user removed via
# Edit / Regenerate. Without this, context_messages still contains
# the full pre-truncation history and the agent sees "deleted"
# turns on the next turn (#2914).
if isinstance(getattr(s, 'context_messages', None), list):
s.context_messages = s.context_messages[:keep]
try:
from api.session_ops import _truncation_watermark_for
s.truncation_watermark = _truncation_watermark_for(s.messages)
# Persist the original truncate cutoff.
s.truncation_boundary = s.truncation_watermark
except Exception:
s.truncation_watermark = 0.0
s.truncation_boundary = 0.0
from api.session_ops import truncate_session_at_keep
old_msg_count, old_ctx_count = truncate_session_at_keep(s, keep)
s.save()
logger.info(
"truncate %s: messages %d%d, context_messages %d%d, watermark=%.2f",
@@ -12252,6 +12237,8 @@ def handle_post(handler, parsed) -> bool:
old_ctx_count, len(getattr(s, 'context_messages', None) or []),
s.truncation_watermark or 0,
)
from api.config import _evict_session_agent
_evict_session_agent(body["session_id"])
return j(
handler, {"ok": True, "session": s.compact() | {"messages": s.messages}}
)
+17
View File
@@ -90,6 +90,23 @@ def truncate_context_for_display_keep(
return prefix + suffix[:keep]
def truncate_session_at_keep(session, keep: int) -> tuple[int, int]:
"""Truncate display + context; set watermark/boundary. Returns old counts."""
full_messages = list(session.messages or [])
old_msg_count = len(full_messages)
old_ctx_count = len(getattr(session, 'context_messages', None) or [])
session.messages = full_messages[:keep]
if isinstance(getattr(session, 'context_messages', None), list):
session.context_messages = truncate_context_for_display_keep(
session.context_messages,
full_messages,
keep,
)
session.truncation_watermark = _truncation_watermark_for(session.messages)
session.truncation_boundary = session.truncation_watermark
return old_msg_count, old_ctx_count
def retry_last(session_id: str) -> dict[str, Any]:
"""Truncate the session to before the last user message, return its text.
@@ -139,6 +139,72 @@ def test_truncate_endpoint_also_truncates_context_messages(monkeypatch, tmp_path
assert loaded.truncation_watermark == 2.0
def test_truncate_endpoint_compaction_leading_context_row(monkeypatch, tmp_path):
"""Context longer than display (leading compaction row): naive [:keep] would
leave a stale tail; truncate must align suffix to display prefix (#5096 / C).
"""
import json
from io import BytesIO
from types import SimpleNamespace
import api.models as models
import api.routes as routes
from api.models import Session
session_dir = tmp_path / "sessions"
session_dir.mkdir(parents=True)
monkeypatch.setattr(models, "SESSION_DIR", session_dir)
monkeypatch.setattr(models, "SESSION_INDEX_FILE", session_dir / "_index.json")
models.SESSIONS.clear()
monkeypatch.setattr(
"api.config._evict_session_agent",
lambda _sid: None,
)
display = [
_msg("user", "u1", 1.0, "u1"),
_msg("assistant", "a1", 2.0, "a1"),
_msg("user", "REMOVE", 3.0, "u2"),
]
context = [
_msg("user", "compaction-only", 0.5, "cref"),
_msg("user", "u1", 1.0, "cu1"),
_msg("assistant", "a1", 2.0, "ca1"),
_msg("user", "REMOVE", 3.0, "cu2"),
]
session = Session(
session_id="issue5096truncatectx",
messages=display,
context_messages=context,
)
session.save()
body = {"session_id": "issue5096truncatectx", "keep_count": 2}
body_bytes = json.dumps(body).encode()
monkeypatch.setattr(routes, "_check_csrf", lambda handler: True)
captured_response = {}
def fake_j(handler, payload, status=200, extra_headers=None):
captured_response["payload"] = payload
monkeypatch.setattr(routes, "j", fake_j)
handler = SimpleNamespace(
headers={"Content-Length": str(len(body_bytes))},
rfile=BytesIO(body_bytes),
)
routes.handle_post(handler, SimpleNamespace(path="/api/session/truncate"))
assert captured_response["payload"].get("ok") is True
loaded = Session.load("issue5096truncatectx")
assert loaded is not None
assert [m["content"] for m in loaded.messages] == ["u1", "a1"]
assert len(loaded.context_messages) == 3
assert loaded.context_messages[0]["content"] == "compaction-only"
assert "REMOVE" not in [m["content"] for m in loaded.context_messages]
def test_truncate_without_context_messages_truncation_leaks_to_agent(monkeypatch, tmp_path):
"""Prove the bug: if context_messages is NOT truncated, agent sees old rows."""
import api.models as models
+25
View File
@@ -0,0 +1,25 @@
"""truncate_session_at_keep aligns context when lengths differ (#5096 C)."""
from api.models import Session
from api.session_ops import truncate_session_at_keep
def test_truncate_session_at_keep_compaction_prefix():
s = Session(
session_id="t1",
messages=[
{"role": "user", "content": "u1", "timestamp": 1},
{"role": "assistant", "content": "a1", "timestamp": 2},
{"role": "user", "content": "gone", "timestamp": 3},
],
context_messages=[
{"role": "user", "content": "cref", "timestamp": 0.5},
{"role": "user", "content": "u1", "timestamp": 1},
{"role": "assistant", "content": "a1", "timestamp": 2},
{"role": "user", "content": "gone", "timestamp": 3},
],
)
truncate_session_at_keep(s, 2)
assert len(s.messages) == 2
assert len(s.context_messages) == 3
assert "gone" not in [m["content"] for m in s.context_messages]