diff --git a/pycodeloop/providers/generic.py b/pycodeloop/providers/generic.py index 6ee9471..721c00c 100644 --- a/pycodeloop/providers/generic.py +++ b/pycodeloop/providers/generic.py @@ -5,10 +5,12 @@ import json import os 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 @@ -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 @@ -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: @@ -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( @@ -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( @@ -286,13 +330,17 @@ 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: @@ -300,10 +348,10 @@ def complete( 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) @@ -315,6 +363,7 @@ def _stream( body: dict, on_delta: Callable[[str], None], known_tools: set[str], + config: _ConnectionSnapshot, ) -> ProviderResponse: body = {**body, "stream": True} text = "" @@ -322,7 +371,7 @@ def _stream( 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: "): diff --git a/tests/providers/test_generic.py b/tests/providers/test_generic.py index e342328..cf310e8 100644 --- a/tests/providers/test_generic.py +++ b/tests/providers/test_generic.py @@ -3,6 +3,7 @@ import io import json import tempfile +import threading import unittest from pathlib import Path from unittest import mock @@ -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()