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
97 changes: 73 additions & 24 deletions pycodeloop/providers/generic.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,10 +5,12 @@
import json
Comment thread
FernandoCelmer marked this conversation as resolved.
import os
Comment thread
FernandoCelmer marked this conversation as resolved.
import re
import threading
import urllib.error
import urllib.request
import uuid
from collections.abc import Callable
from dataclasses import dataclass
from pathlib import Path

from pycodeloop.abc.provider import Provider, ProviderResponse, ToolCall, Usage
Expand Down Expand Up @@ -101,6 +103,26 @@ def _fallback_call_id() -> str:
return f"fallback-{uuid.uuid4().hex[:8]}"


@dataclass(frozen=True)
class _ConnectionSnapshot:
"""A consistent, point-in-time copy of everything `reload()` can
mutate — read once under `GenericProvider._lock` at the start of
`complete()`/`_stream()` so a concurrent `reload()` (e.g. `ask()`
running on another thread mid-`run()`) can't hand a request a
half-old/half-new mix of url/model/headers/parser."""

url: str
model: str
headers: dict[str, str]
auth_header: str
auth_prefix: str
api_key: str | None
timeout: float
request_builder: RequestBuilder
response_parser: ResponseParser
uses_default_parser: bool


class GenericProvider(Provider):
"""Any JSON chat-completions HTTP API via the stdlib, no vendor SDK.
Defaults to the OpenAI request/response shape; override
Expand Down Expand Up @@ -179,6 +201,7 @@ def __init__(
self.repetition_repeats = repetition_repeats
self._uses_default_parser = response_parser is None
self._config_path: Path | None = None
self._lock = threading.Lock()

@classmethod
def from_json(cls, path: str | Path) -> GenericProvider:
Expand Down Expand Up @@ -232,16 +255,17 @@ def reload(self) -> None:
return

fresh = self._build_from_json(self._config_path)
self.url = fresh.url
self.model = fresh.model
self.api_key = fresh.api_key
self.headers = fresh.headers
self.auth_header = fresh.auth_header
self.auth_prefix = fresh.auth_prefix
self.request_builder = fresh.request_builder
self.response_parser = fresh.response_parser
self.timeout = fresh.timeout
self._uses_default_parser = fresh._uses_default_parser
with self._lock:
self.url = fresh.url
self.model = fresh.model
self.api_key = fresh.api_key
self.headers = fresh.headers
self.auth_header = fresh.auth_header
self.auth_prefix = fresh.auth_prefix
self.request_builder = fresh.request_builder
self.response_parser = fresh.response_parser
self.timeout = fresh.timeout
self._uses_default_parser = fresh._uses_default_parser

@staticmethod
def _default_request(
Expand All @@ -256,19 +280,39 @@ def _default_request(
"tools": openai_tool_schema(tools) if tools else None,
}

def _headers(self) -> dict[str, str]:
headers = {"Content-Type": "application/json", **self.headers}
if self.api_key and self.auth_header not in headers:
headers[self.auth_header] = f"{self.auth_prefix}{self.api_key}"
def _snapshot_locked(self) -> _ConnectionSnapshot:
"""Caller must hold `self._lock`."""
return _ConnectionSnapshot(
url=self.url,
model=self.model,
headers=dict(self.headers),
auth_header=self.auth_header,
auth_prefix=self.auth_prefix,
api_key=self.api_key,
timeout=self.timeout,
request_builder=self.request_builder,
response_parser=self.response_parser,
uses_default_parser=self._uses_default_parser,
)

def _headers(self, config: _ConnectionSnapshot) -> dict[str, str]:
headers = {"Content-Type": "application/json", **config.headers}
if config.api_key and config.auth_header not in headers:
headers[config.auth_header] = (
f"{config.auth_prefix}{config.api_key}"
)
return headers

def _open(self, body: dict):
def _open(self, body: dict, config: _ConnectionSnapshot):
data = json.dumps(body).encode()
request = urllib.request.Request(
self.url, data=data, headers=self._headers(), method="POST"
config.url,
data=data,
headers=self._headers(config),
method="POST",
)
try:
return urllib.request.urlopen(request, timeout=self.timeout)
return urllib.request.urlopen(request, timeout=config.timeout)
except urllib.error.HTTPError as exc:
detail = exc.read().decode(errors="replace")
raise urllib.error.HTTPError(
Expand All @@ -286,24 +330,28 @@ def complete(
tools: list[dict],
on_delta: Callable[[str], None] | None = None,
) -> ProviderResponse:
body = self.request_builder(system_prompt, messages, tools, self.model)
with self._lock:
config = self._snapshot_locked()
body = config.request_builder(
system_prompt, messages, tools, config.model
)
known_tools = {tool["name"] for tool in tools}

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

with self._open(body) as response:
with self._open(body, config) as response:
raw = response.read()

try:
data = json.loads(raw)
except json.JSONDecodeError as exc:
snippet = raw.decode(errors="replace")[:500]
raise ValueError(
f"{self.url} returned malformed/truncated JSON ({exc}): {snippet!r}"
f"{config.url} returned malformed/truncated JSON ({exc}): {snippet!r}"
) from None

result = self.response_parser(data)
result = config.response_parser(data)

if on_delta is not None and result.text:
on_delta(result.text)
Expand All @@ -315,14 +363,15 @@ def _stream(
body: dict,
on_delta: Callable[[str], None],
known_tools: set[str],
config: _ConnectionSnapshot,
) -> ProviderResponse:
body = {**body, "stream": True}
text = ""
pending: dict[int, dict] = {}
stop_reason = "stop"
usage = Usage()

with self._open(body) as response:
with self._open(body, config) as response:
for raw_line in response:
line = raw_line.decode().strip()
if not line or not line.startswith("data: "):
Expand Down
54 changes: 54 additions & 0 deletions tests/providers/test_generic.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import io
import json
import tempfile
import threading
import unittest
from pathlib import Path
from unittest import mock
Expand Down Expand Up @@ -510,5 +511,58 @@ def test_get_provider_model_kwarg_overrides_json_config(self):
self.assertEqual(provider.model, "from-cli")


class TestReloadThreadSafety(GenericProviderTestCase):
def test_concurrent_complete_never_sees_a_mixed_config(self):
path = self._write_config({"url": "http://fake/A", "model": "model-A"})
provider = GenericProvider.from_json(path)

seen: list[tuple[str, str]] = []
seen_lock = threading.Lock()

def fake_urlopen(request, timeout=None):
body = json.loads(request.data)
with seen_lock:
seen.append((request.full_url, body["model"]))
payload = json.dumps(
{
"choices": [
{"message": {"content": "ok"}, "finish_reason": "stop"}
],
"usage": {},
}
).encode()
return _FakeResponse(payload)

def flip_config():
for i in range(50):
tag = "A" if i % 2 == 0 else "B"
path.write_text(
json.dumps(
{"url": f"http://fake/{tag}", "model": f"model-{tag}"}
)
)
provider.reload()

def call_complete():
for _ in range(50):
provider.complete("sys", [], [])

with mock.patch(
"pycodeloop.providers.generic.urllib.request.urlopen",
side_effect=fake_urlopen,
):
reloader = threading.Thread(target=flip_config)
caller = threading.Thread(target=call_complete)
reloader.start()
caller.start()
reloader.join()
caller.join()

self.assertTrue(seen)
for url, model in seen:
tag = url.rsplit("/", 1)[-1]
self.assertEqual(model, f"model-{tag}")


if __name__ == "__main__":
unittest.main()
Loading