Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 7 additions & 9 deletions pycodeloop/core/agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -325,19 +324,18 @@ 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

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.",
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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 = [
{
Expand Down
9 changes: 9 additions & 0 deletions pycodeloop/core/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Comment thread
FernandoCelmer marked this conversation as resolved.
default_factory=threading.Lock, repr=False, compare=False
)
Expand All @@ -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()
Expand Down
26 changes: 26 additions & 0 deletions tests/core/test_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
[
Expand Down
Loading