Newer
Older
navi-1 / tests / unit / workers / test_compression_worker.py
"""Tests for the post-turn CompressionWorker."""

from __future__ import annotations

import pytest

from navi.core.events import ContextCompressed
from navi.core.session import InMemorySessionStore, Session
from navi.workers.base import WorkerContext
from navi.workers.compressor import CompressionWorker
from navi.llm.base import Message
from tests.conftest_factory import FakeLLMBackend


def _ctx(session_id: str, llm, store) -> WorkerContext:
    return WorkerContext(
        session_id=session_id,
        context_tokens=60_000,  # over the 0.90 threshold -> worker attempts compression
        max_context_tokens=65536,
        llm=llm,
        model="test",
        temperature=0.3,
        session_store=store,
    )


@pytest.mark.asyncio
async def test_worker_compresses_single_long_turn():
    """Regression: a single long autonomous turn (one user message + many tool
    iterations = one turn) used to no-op because the worker called
    compress_context without keep_recent_messages. Now it mirrors the midturn
    path and compresses via the intra-turn fallback."""
    worker = CompressionWorker()
    backend = FakeLLMBackend(responses=["summary of the work so far"])
    store = InMemorySessionStore()
    session = Session(profile_id="test")
    store._sessions[session.id] = session
    session.context.append(Message(role="user", content="do the task"))
    for i in range(20):
        session.context.append(Message(role="assistant", content=f"step {i} " * 40))

    result = await worker.run(session, _ctx(session.id, backend, store))
    assert any(isinstance(ev, ContextCompressed) for ev in result.events)
    compressed = next(ev for ev in result.events if isinstance(ev, ContextCompressed))
    assert compressed.messages_after < compressed.messages_before
    assert compressed.summary == "summary of the work so far"


@pytest.mark.asyncio
async def test_worker_no_op_when_context_small():
    """Below the token threshold the worker does nothing."""
    worker = CompressionWorker()
    backend = FakeLLMBackend(responses=["irrelevant"])
    store = InMemorySessionStore()
    session = Session(profile_id="test")
    store._sessions[session.id] = session
    session.context.append(Message(role="user", content="hi"))
    session.context.append(Message(role="assistant", content="hello"))

    ctx = WorkerContext(
        session_id=session.id,
        context_tokens=10,  # well below threshold
        max_context_tokens=65536,
        llm=backend,
        model="test",
        temperature=0.3,
        session_store=store,
    )
    result = await worker.run(session, ctx)
    assert result.events == []


@pytest.mark.asyncio
async def test_worker_hard_truncates_when_summarizer_always_fails():
    """The worker delegates to compress_and_save_session, so a summarizer LLM
    that always fails still ends with the hard-truncate fallback — the context
    shrinks and ContextCompressed is emitted. The old worker called
    compress_context directly and silently no-op'ed on any LLM failure,
    leaving the session over the threshold until the next turn's gates."""
    import navi.core.compressor as compressor_module

    worker = CompressionWorker()
    backend = FakeLLMBackend(responses=["unused"])
    store = InMemorySessionStore()
    session = Session(profile_id="test")
    store._sessions[session.id] = session
    # 4 turns, each assistant ~20k tokens -> hard-truncate (0.5 * 65536)
    # keeps only the newest turn.
    big = "x" * 60_000
    for i in range(4):
        session.context.append(Message(role="user", content=str(i)))
        session.context.append(Message(role="assistant", content=big))

    async def _always_fail(*args, **kwargs):
        raise RuntimeError("summarizer down")

    original = compressor_module.compress_context
    compressor_module.compress_context = _always_fail
    try:
        result = await worker.run(session, _ctx(session.id, backend, store))
    finally:
        compressor_module.compress_context = original

    compressed = [ev for ev in result.events if isinstance(ev, ContextCompressed)]
    assert compressed, "worker must not silently no-op when the summarizer fails"
    assert compressed[0].messages_after < compressed[0].messages_before
    assert "truncated" in compressed[0].summary.lower()
    assert len(session.context) == compressed[0].messages_after


@pytest.mark.asyncio
async def test_worker_marks_dropped_messages_not_in_context():
    """Delegated path keeps the message marking: dropped messages are flagged
    is_context=False so a reload does not resurrect them into the context."""
    worker = CompressionWorker()
    backend = FakeLLMBackend(responses=["worker summary"])
    store = InMemorySessionStore()
    session = Session(profile_id="test")
    store._sessions[session.id] = session
    session.context.append(Message(role="user", content="1"))
    session.context.append(Message(role="assistant", content="a1"))
    session.context.append(Message(role="user", content="2"))
    session.context.append(Message(role="assistant", content="a2"))
    for i in range(12):
        session.context.append(Message(role="user", content=f"task {i}"))
        session.context.append(Message(role="assistant", content=f"answer {i}"))
    session.messages = list(session.context)

    result = await worker.run(session, _ctx(session.id, backend, store))
    assert any(isinstance(ev, ContextCompressed) for ev in result.events)
    kept_ids = {id(m) for m in session.context}
    dropped = [m for m in session.messages if id(m) not in kept_ids and m.role != "system"]
    assert dropped, "old turns must be marked is_context=False"
    assert all(m.is_context is False for m in dropped)


def test_worker_keep_recent_messages_mirrors_midturn():
    """The worker passes keep_recent_messages=max(12, context_keep_recent*2),
    matching the midturn auto-compress path (source of truth for the intra-turn
    fallback). Guards against regressing back to keep_recent_messages=None."""
    import inspect

    src = inspect.getsource(CompressionWorker.run)
    assert "keep_recent_messages" in src
    assert "max(12" in src