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
7 changes: 0 additions & 7 deletions pycodeloop/abc/provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,6 @@

from __future__ import annotations

import threading
from abc import ABC, abstractmethod
from collections.abc import Callable
from dataclasses import dataclass, field
Expand Down Expand Up @@ -66,15 +65,9 @@ def complete(
messages: list,
tools: list[dict],
on_delta: Callable[[str], None] | None = None,
cancel_event: threading.Event | None = None,
) -> ProviderResponse:
"""Send conversation + tool schema, return the model's response.

When `on_delta` is given, call it with each text chunk as it
arrives instead of only returning the assembled text at the end.

When `cancel_event` is given and set while streaming, the
implementation should stop reading as soon as practical and
return with `stop_reason="cancelled"` instead of waiting for the
provider to finish on its own.
"""
21 changes: 0 additions & 21 deletions pycodeloop/core/agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -412,22 +412,9 @@ def run(
messages=session.history(),
tools=tools,
on_delta=self.on_text_delta,
cancel_event=cancel_event,
)
elapsed = time.perf_counter() - started_at

if response.stop_reason == "cancelled":
self.usage = self.usage + response.usage
if self.on_usage:
self.on_usage(response.usage, self.usage, elapsed)
if response.text.strip() or response.tool_calls:
session.add_assistant(response.text)
self._notify_message()
if self.on_turn_end:
self.on_turn_end()
self._trace("run_end", reason="cancelled")
return "Cancelled by user."

empty_retries = 0
while (
not response.text.strip()
Expand Down Expand Up @@ -459,17 +446,9 @@ def run(
messages=session.history(),
tools=tools,
on_delta=self.on_text_delta,
cancel_event=cancel_event,
)
elapsed = time.perf_counter() - started_at

if response.stop_reason == "cancelled":
self.usage = self.usage + response.usage
if self.on_usage:
self.on_usage(response.usage, self.usage, elapsed)
self._trace("run_end", reason="cancelled")
return "Cancelled by user."

if not response.text.strip() and not response.tool_calls:
self.usage = self.usage + response.usage
if self.on_usage:
Expand Down
10 changes: 1 addition & 9 deletions pycodeloop/providers/generic.py
Original file line number Diff line number Diff line change
Expand Up @@ -350,7 +350,6 @@ def complete(
messages: list[Message],
tools: list[dict],
on_delta: Callable[[str], None] | None = None,
cancel_event: threading.Event | None = None,
) -> ProviderResponse:
with self._lock:
config = self._snapshot_locked()
Expand All @@ -360,9 +359,7 @@ def complete(
known_tools = {tool["name"] for tool in tools}

if on_delta is not None and config.supports_openai_sse:
return self._stream(
body, on_delta, known_tools, config, cancel_event
)
return self._stream(body, on_delta, known_tools, config)

with self._open(body, config) as response:
raw = response.read()
Expand All @@ -388,7 +385,6 @@ def _stream(
on_delta: Callable[[str], None],
known_tools: set[str],
config: _ConnectionSnapshot,
cancel_event: threading.Event | None = None,
) -> ProviderResponse:
body = {**body, "stream": True}
if config.include_usage_in_stream:
Expand All @@ -405,10 +401,6 @@ def _stream(

with self._open(body, config) as response:
for raw_line in response:
if cancel_event is not None and cancel_event.is_set():
stop_reason = "cancelled"
saw_terminal_marker = True
break
line = raw_line.decode().strip()
if not line or not line.startswith("data: "):
continue
Expand Down
28 changes: 2 additions & 26 deletions tests/core/test_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@ def __init__(self, scripted: list[ProviderResponse]) -> None:
self._scripted = list(scripted)

def complete(
self, system_prompt, messages, tools, on_delta=None, cancel_event=None
self, system_prompt, messages, tools, on_delta=None
) -> ProviderResponse:
return self._scripted.pop(0)

Expand Down Expand Up @@ -631,30 +631,6 @@ def test_cancel_event_set_before_run_returns_immediately(self):
self.assertEqual(result, "Cancelled by user.")
self.assertEqual(len(provider._scripted), 1)

def test_cancelled_provider_response_stops_the_run_immediately(self):
"""Regression: cancel_event was never threaded into
provider.complete(), so cancelling mid-stream had no effect
until the current turn finished on its own. Once the provider
reports stop_reason='cancelled', the run must stop right away
instead of retrying or continuing to the next turn."""
provider = FakeProvider(
[
ProviderResponse(
text="partial answer", stop_reason="cancelled"
),
ProviderResponse(text="should never be reached"),
]
)
cancel_event = threading.Event()
agent = Agent(provider=provider, tools=[EchoTool()])
session = Session(system_prompt="sys")

result = agent.run("go", session=session, cancel_event=cancel_event)

self.assertEqual(result, "Cancelled by user.")
self.assertEqual(len(provider._scripted), 1)
self.assertEqual(session.messages[-1].content, "partial answer")


class FlakyProvider(Provider):
"""Raises a retryable error `fail_times` times, then succeeds."""
Expand All @@ -668,7 +644,7 @@ def __init__(self, fail_times: int, status_code: int = 429) -> None:
self.calls = 0

def complete(
self, system_prompt, messages, tools, on_delta=None, cancel_event=None
self, system_prompt, messages, tools, on_delta=None
) -> ProviderResponse:
self.calls += 1
if self.calls <= self.fail_times:
Expand Down
2 changes: 1 addition & 1 deletion tests/core/test_codeloop.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,7 @@ def __init__(self, scripted: list[ProviderResponse]) -> None:
self._scripted = list(scripted)

def complete(
self, system_prompt, messages, tools, on_delta=None, cancel_event=None
self, system_prompt, messages, tools, on_delta=None
) -> ProviderResponse:
return self._scripted.pop(0)

Expand Down
35 changes: 0 additions & 35 deletions tests/providers/test_generic.py
Original file line number Diff line number Diff line change
Expand Up @@ -338,41 +338,6 @@ def test_streaming_flags_a_connection_dropped_mid_response(self):
self.assertEqual(result.text, "cut off mid")
self.assertEqual(result.stop_reason, "connection_lost")

def test_streaming_stops_promptly_when_cancel_event_is_set(self):
"""Regression: cancel_event was accepted nowhere in the streaming
read loop, so pressing Esc/Cancel mid-response did nothing until
the provider finished the turn on its own."""
path = self._write_config(
{"url": "http://fake/v1/chat/completions", "model": "my-model"}
)
provider = GenericProvider.from_json(path)

chunks = [
{"choices": [{"delta": {"content": "first"}}]},
{"choices": [{"delta": {"content": "second"}}]},
{"choices": [{"delta": {}, "finish_reason": "stop"}]},
]
sse_body = (
"".join(f"data: {json.dumps(c)}\n" for c in chunks)
+ "data: [DONE]\n"
).encode()

cancel_event = threading.Event()

def on_delta(_chunk):
cancel_event.set()

with mock.patch(
"pycodeloop.providers.generic.urllib.request.urlopen",
return_value=_FakeResponse(sse_body),
):
result = provider.complete(
"sys", [], [], on_delta=on_delta, cancel_event=cancel_event
)

self.assertEqual(result.text, "first")
self.assertEqual(result.stop_reason, "cancelled")

def test_streaming_requests_usage_and_captures_it_from_final_chunk(self):
"""Regression: streaming previously sent `stream: True` with no
`stream_options.include_usage`, so OpenAI-compatible servers that
Expand Down
2 changes: 1 addition & 1 deletion tests/tools/test_delegate.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@ def __init__(self, scripted: list[ProviderResponse]) -> None:
self.requests: list[list] = []

def complete(
self, system_prompt, messages, tools, on_delta=None, cancel_event=None
self, system_prompt, messages, tools, on_delta=None
) -> ProviderResponse:
self.requests.append(tools)
return self._scripted.pop(0)
Expand Down
Loading