diff --git a/pycodeloop/core/agent.py b/pycodeloop/core/agent.py index dd860cb..bf5e443 100644 --- a/pycodeloop/core/agent.py +++ b/pycodeloop/core/agent.py @@ -111,7 +111,6 @@ def __init__( self.on_turn_end = on_turn_end self.on_trace_event = on_trace_event self.usage = Usage() - self._last_context_tokens = 0 def _trace(self, event_type: str, **fields) -> None: if self.on_trace_event: @@ -325,9 +324,8 @@ def _compact(self, session: Session) -> None: provider itself, replacing the older history with one condensed message — keeps the conversation going instead of hitting the model's context limit.""" - turn_starts = [ - i for i, m in enumerate(session.messages) if m.role == "user" - ] + history = session.history() + turn_starts = [i for i, m in enumerate(history) if m.role == "user"] if len(turn_starts) <= _COMPACT_KEEP_RECENT_TURNS: return @@ -335,9 +333,9 @@ def _compact(self, session: Session) -> None: if self.on_compact_start: self.on_compact_start() - before_count = len(session.messages) + before_count = len(history) cutoff = turn_starts[-_COMPACT_KEEP_RECENT_TURNS] - older, recent = session.messages[:cutoff], session.messages[cutoff:] + older, recent = history[:cutoff], history[cutoff:] summary = self._complete( system_prompt="Summarize conversations concisely for context compaction.", @@ -397,7 +395,7 @@ def run( context_window = getattr(self.provider, "context_window", None) if context_window is None: context_window = context_window_for(self.provider.model) - if self.auto_compact and self._last_context_tokens >= ( + if self.auto_compact and session.get_last_context_tokens() >= ( context_window * self.compact_threshold ): self._compact(session) @@ -429,9 +427,9 @@ def run( output_tokens=response.usage.output_tokens, ) - self._last_context_tokens = response.usage.input_tokens + session.update_last_context_tokens(response.usage.input_tokens) if self.on_context: - self.on_context(self._last_context_tokens, context_window) + self.on_context(response.usage.input_tokens, context_window) tool_calls = [ { diff --git a/pycodeloop/core/session.py b/pycodeloop/core/session.py index 9e29227..a2e5ca0 100644 --- a/pycodeloop/core/session.py +++ b/pycodeloop/core/session.py @@ -16,6 +16,7 @@ class Session: messages: list[Message] = field(default_factory=list) cwd: str = "." dirty: bool = field(default=False, repr=False, compare=False) + last_context_tokens: int = field(default=0, repr=False, compare=False) _lock: threading.Lock = field( default_factory=threading.Lock, repr=False, compare=False ) @@ -42,6 +43,14 @@ def add_tool_result(self, tool_call_id: str, content: str) -> None: ) ) + def get_last_context_tokens(self) -> int: + with self._lock: + return self.last_context_tokens + + def update_last_context_tokens(self, value: int) -> None: + with self._lock: + self.last_context_tokens = value + def history(self) -> list[Message]: with self._lock: self._repair_dangling_tool_calls() diff --git a/tests/core/test_agent.py b/tests/core/test_agent.py index e916d53..7d5ead3 100644 --- a/tests/core/test_agent.py +++ b/tests/core/test_agent.py @@ -451,6 +451,32 @@ def test_compacts_when_context_usage_crosses_threshold(self): self.assertIn("second", session.messages[0].content) self.assertEqual(len(session.messages), 4) + def test_context_usage_does_not_leak_across_sessions(self): + provider = FakeProvider( + [ + ProviderResponse(text="b1", usage=Usage(input_tokens=500)), + ProviderResponse(text="b2", usage=Usage(input_tokens=500)), + ProviderResponse(text="a1", usage=Usage(input_tokens=190_000)), + ProviderResponse(text="b3", usage=Usage(input_tokens=500)), + ] + ) + provider.model = "claude-sonnet-5" + events = [] + agent = Agent( + provider=provider, + on_compact_start=lambda: events.append("start"), + ) + session_a = Session(system_prompt="sys") + session_b = Session(system_prompt="sys") + + agent.run("b-first", session=session_b) + agent.run("b-second", session=session_b) + agent.run("a-first", session=session_a) + agent.run("b-third", session=session_b) + + self.assertEqual(events, []) + self.assertEqual(len(session_b.messages), 6) + def test_auto_compact_false_never_compacts(self): provider = FakeProvider( [