From 1411c634fd295d848531198e7ad4b72b05c4f4a7 Mon Sep 17 00:00:00 2001 From: Ahmed Janabi Date: Sun, 30 Aug 2026 20:23:09 +0000 Subject: [PATCH] feat: add OpenAI-compatible inference backend (llama.cpp-server / LM Studio) Adds PW_INFERENCE_BACKEND=ollama|openai + PW_INFERENCE_API_BASE to council/ollama.py, the single shared HTTP client used by worker, researcher, judge, batch, and rerank. On the default (ollama) backend behavior is unchanged. On openai, base()/generate()/smallest_chat_model() speak the OpenAI-compatible /v1/chat/completions + /v1/models shape that llama.cpp-server and LM Studio expose, giving every council/ollama.py caller a new backend option with zero changes to those callers. Other direct-to-Ollama call sites (council/local.py, council/library.py, council/doctor.py, council/net/agent.py, council/net/baseline.py, council/operator.py) are intentionally out of scope and remain Ollama-only, per the task boundary. Adds tests/test_ollama.py (module had zero dedicated tests) covering base() resolution on both backends, generate()'s request shape on both backends, resolve_timeout() precedence, and KNOWN registry entries. Co-Authored-By: Claude Sonnet 5 --- council/config.py | 2 + council/ollama.py | 101 +++++++++++++++++++++++++--- tests/test_ollama.py | 155 +++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 247 insertions(+), 11 deletions(-) create mode 100644 tests/test_ollama.py diff --git a/council/config.py b/council/config.py index 3cc491c..38c1e51 100644 --- a/council/config.py +++ b/council/config.py @@ -43,6 +43,8 @@ "PW_OLLAMA_BASE": {"help": "Ollama endpoint (default http://localhost:11434)", "secret": False}, "PW_OLLAMA_KEEP_ALIVE": {"help": "how long Ollama keeps a model warm (default 30m)", "secret": False}, "PW_OLLAMA_TIMEOUT": {"help": "per-generation timeout seconds (default 300)", "secret": False}, + "PW_INFERENCE_BACKEND": {"help": "ollama | openai — which inference API shape to speak (default ollama)", "secret": False}, + "PW_INFERENCE_API_BASE": {"help": "OpenAI-compatible endpoint (llama.cpp-server/LM Studio) when PW_INFERENCE_BACKEND=openai", "secret": False}, "PW_MODEL_CAP_GB": {"help": "skip models larger than this many GB (0 = no cap)", "secret": False}, "PW_PAGE_EVIDENCE": {"help": "fetch full source pages for drafting, not just snippets (default 1)", "secret": False}, "PW_COUNTRY": {"help": "your location, tags the analyst vantage point", "secret": False}, diff --git a/council/ollama.py b/council/ollama.py index 7293f31..369a7e5 100644 --- a/council/ollama.py +++ b/council/ollama.py @@ -10,6 +10,20 @@ model *detection* honored ``PW_OLLAMA_BASE`` — so pointing it at a GPU box listed that box's models and then failed every generation against localhost (D48). Routing them all through `base()` / `generate()` fixes that in one place and removes four near-identical HTTP clients. + +Inference backend: ``PW_INFERENCE_BACKEND`` selects which HTTP API shape this module speaks — +``ollama`` (default) for Ollama's native ``/api/generate`` + ``/api/tags``, or ``openai`` for the +OpenAI-compatible ``/v1/chat/completions`` (+ ``/v1/models``) shape that both llama.cpp-server and +LM Studio expose out of the box as of 2026. On ``openai`` the endpoint comes from +``PW_INFERENCE_API_BASE`` (no hardcoded default — there is no sane guess for an arbitrary +llama.cpp-server/LM Studio port the way ``localhost:11434`` is for Ollama). This is purely a +local-inference-runtime choice; it does not touch D1/D4/D18. Known/intentional limitation: only +callers that go through *this module's* `generate()` / `base()` / `smallest_chat_model()` get the +new backend for free (`worker`, `researcher`, `judge`, `batch`, `rerank`). Other call sites that +talk to Ollama directly elsewhere in the codebase (e.g. `council/local.py`, `council/library.py`, +`council/doctor.py`, `council/net/agent.py`, `council/net/baseline.py`, `council/operator.py`) are +NOT routed through this dispatch and remain Ollama-only — that's an intentional scope boundary for +now, not an oversight. """ from __future__ import annotations @@ -21,9 +35,22 @@ DEFAULT_BASE = "http://localhost:11434" +def backend() -> str: + """Which inference API shape to speak: ``ollama`` (default) or ``openai``, from + ``PW_INFERENCE_BACKEND``. Normalized (stripped, lowercased) so ``" OpenAI "`` still matches.""" + return (os.environ.get("PW_INFERENCE_BACKEND") or "ollama").strip().lower() + + def base() -> str: - """The Ollama endpoint every call should use, resolved from ``PW_OLLAMA_BASE`` at call/instance - time (so a process can target a remote host without editing code). Trailing slash stripped.""" + """The inference endpoint every call should use. On the default ``ollama`` backend, resolved + from ``PW_OLLAMA_BASE`` (or `DEFAULT_BASE`) at call/instance time (so a process can target a + remote host without editing code). On the ``openai`` backend, resolved from + ``PW_INFERENCE_API_BASE`` — there is no hardcoded default for it (no sane guess for an + arbitrary llama.cpp-server/LM Studio port); if unset, an empty string is returned and the + subsequent `requests` call fails naturally with a clear connection error, matching this + module's existing style of not pre-validating env vars. Trailing slash stripped.""" + if backend() == "openai": + return (os.environ.get("PW_INFERENCE_API_BASE") or "").rstrip("/") return (os.environ.get("PW_OLLAMA_BASE") or DEFAULT_BASE).rstrip("/") @@ -35,12 +62,22 @@ def keep_alive() -> str: def smallest_chat_model(base_url: str | None = None) -> str: """The smallest installed non-embedding model — cheap for auxiliary calls like listwise reranking (``council.rerank``). ``PW_SMALL_MODEL`` overrides; ``''`` if none reachable. - Resolved at call time against ``base_url`` (or ``base()``) so a remote host is honored.""" + Resolved at call time against ``base_url`` (or ``base()``) so a remote host is honored. + + On the ``openai`` backend this queries ``/v1/models`` (supported by both llama.cpp-server and + LM Studio) and picks the first entry in ``data["data"]`` — that endpoint carries no per-model + size, so there is no "smallest" to sort by; any lookup failure (unreachable host, no models, + malformed response) falls through to the same ``except Exception: return ""`` used by the + Ollama path.""" override = os.environ.get("PW_SMALL_MODEL", "") if override: return override + resolved = (base_url or base()).rstrip("/") try: - r = requests.get(f"{(base_url or base()).rstrip('/')}/api/tags", timeout=10) + if backend() == "openai": + r = requests.get(f"{resolved}/v1/models", timeout=10) + return r.json()["data"][0]["id"] + r = requests.get(f"{resolved}/api/tags", timeout=10) models = [m for m in r.json().get("models", []) if "embed" not in m["name"].lower()] return sorted(models, key=lambda m: m.get("size", 0))[0]["name"] except Exception: @@ -61,16 +98,14 @@ def resolve_timeout(env_primary: str | None, default: float) -> float: return float(default) -def generate(prompt: str, *, model: str, base_url: str | None = None, - temperature: float = 0.0, num_predict: int | None = None, - timeout: float | None = None, timeout_env: str | None = None, - timeout_default: float = 300.0) -> tuple[str, int]: +def _generate_ollama(prompt: str, *, model: str, base_url: str | None, + temperature: float, num_predict: int | None, + timeout: float | None, timeout_env: str | None, + timeout_default: float) -> tuple[str, int]: """POST a non-streaming ``/api/generate`` and return ``(text, tokens)``. `text` is the model's response, stripped; `tokens` is Ollama's ``eval_count`` (0 if absent — - callers that want a word-count fallback apply it themselves). `base_url` defaults to `base()`. - Timeout precedence: explicit `timeout` → `timeout_env`/`PW_OLLAMA_TIMEOUT` → `timeout_default`. - Raises ``requests.HTTPError`` on a non-2xx response (callers already handle that).""" + callers that want a word-count fallback apply it themselves).""" options: dict = {"temperature": temperature} if num_predict is not None: options["num_predict"] = num_predict @@ -83,3 +118,47 @@ def generate(prompt: str, *, model: str, base_url: str | None = None, r.raise_for_status() d = r.json() return (d.get("response") or "").strip(), int(d.get("eval_count") or 0) + + +def _generate_openai(prompt: str, *, model: str, base_url: str | None, + temperature: float, num_predict: int | None, + timeout: float | None, timeout_env: str | None, + timeout_default: float) -> tuple[str, int]: + """POST a non-streaming ``/v1/chat/completions`` (OpenAI-compatible: llama.cpp-server, LM + Studio) and return ``(text, tokens)``. + + `text` is ``choices[0].message.content``, stripped; `tokens` is + ``usage.completion_tokens`` (0 if absent — mirrors the Ollama path's "0 if absent" + error-tolerance contract exactly).""" + body: dict = {"model": model, "messages": [{"role": "user", "content": prompt}], + "temperature": temperature} + if num_predict is not None: + body["max_tokens"] = num_predict + r = requests.post( + f"{(base_url or base()).rstrip('/')}/v1/chat/completions", + json=body, + timeout=timeout if timeout is not None else resolve_timeout(timeout_env, timeout_default), + ) + r.raise_for_status() + d = r.json() + text = (d.get("choices") or [{}])[0].get("message", {}).get("content") or "" + tokens = d.get("usage", {}).get("completion_tokens", 0) + return text.strip(), int(tokens or 0) + + +def generate(prompt: str, *, model: str, base_url: str | None = None, + temperature: float = 0.0, num_predict: int | None = None, + timeout: float | None = None, timeout_env: str | None = None, + timeout_default: float = 300.0) -> tuple[str, int]: + """Generate a completion and return ``(text, tokens)``, dispatching to the configured + `backend()` — Ollama's native ``/api/generate`` (default) or an OpenAI-compatible + ``/v1/chat/completions`` (``PW_INFERENCE_BACKEND=openai``). + + `text` is the model's response, stripped; `tokens` is 0 if the backend didn't report a count + (callers that want a word-count fallback apply it themselves). `base_url` defaults to `base()`. + Timeout precedence: explicit `timeout` → `timeout_env`/`PW_OLLAMA_TIMEOUT` → `timeout_default`. + Raises ``requests.HTTPError`` on a non-2xx response (callers already handle that).""" + impl = _generate_openai if backend() == "openai" else _generate_ollama + return impl(prompt, model=model, base_url=base_url, temperature=temperature, + num_predict=num_predict, timeout=timeout, timeout_env=timeout_env, + timeout_default=timeout_default) diff --git a/tests/test_ollama.py b/tests/test_ollama.py new file mode 100644 index 0000000..c48ef63 --- /dev/null +++ b/tests/test_ollama.py @@ -0,0 +1,155 @@ +"""tests/test_ollama.py — council/ollama.py, the single shared HTTP client every generation call +(worker/researcher/judge/batch/rerank) goes through. + +Covers: `base()` resolution for both the default Ollama backend and the OpenAI-compatible +backend (PW_INFERENCE_BACKEND=openai), `generate()`'s request shape on each backend (mocking +`requests.post` and asserting on call args, following tests/test_doctor.py's fake-response +convention), `resolve_timeout()` precedence, and that the new config knobs are registered in +`council.config.KNOWN`. +""" +from __future__ import annotations + +import council.config as C +import council.ollama as O + + +class _FakeResp: + """Minimal stand-in for `requests.Response` (mirrors tests/test_doctor.py's _FakeGenResp).""" + + def __init__(self, payload): + self._p = payload + + def raise_for_status(self): + pass + + def json(self): + return self._p + + +# ---- base() ---------------------------------------------------------------------------------- + +def test_base_default_is_localhost_ollama(monkeypatch): + monkeypatch.delenv("PW_OLLAMA_BASE", raising=False) + monkeypatch.delenv("PW_INFERENCE_BACKEND", raising=False) + assert O.base() == "http://localhost:11434" + + +def test_base_honors_pw_ollama_base_override(monkeypatch): + monkeypatch.delenv("PW_INFERENCE_BACKEND", raising=False) + monkeypatch.setenv("PW_OLLAMA_BASE", "http://gpu-box:11434/") + assert O.base() == "http://gpu-box:11434" # trailing slash stripped + + +def test_base_openai_backend_resolves_from_inference_api_base(monkeypatch): + monkeypatch.setenv("PW_INFERENCE_BACKEND", "openai") + monkeypatch.setenv("PW_INFERENCE_API_BASE", "http://localhost:8080/") + assert O.base() == "http://localhost:8080" # trailing slash stripped + + +def test_base_openai_backend_with_no_api_base_set_is_empty_not_a_crash(monkeypatch): + monkeypatch.setenv("PW_INFERENCE_BACKEND", "openai") + monkeypatch.delenv("PW_INFERENCE_API_BASE", raising=False) + assert O.base() == "" # no hardcoded default for openai + + +def test_backend_default_and_normalization(monkeypatch): + monkeypatch.delenv("PW_INFERENCE_BACKEND", raising=False) + assert O.backend() == "ollama" + monkeypatch.setenv("PW_INFERENCE_BACKEND", " OpenAI ") + assert O.backend() == "openai" + + +# ---- generate() -------------------------------------------------------------------------------- + +def test_generate_default_backend_posts_to_api_generate(monkeypatch): + monkeypatch.delenv("PW_INFERENCE_BACKEND", raising=False) + calls = [] + + def fake_post(url, json=None, timeout=None): + calls.append((url, json, timeout)) + return _FakeResp({"response": " hello ", "eval_count": 7}) + + monkeypatch.setattr(O.requests, "post", fake_post) + text, tokens = O.generate("hi", model="m1", base_url="http://x:11434", num_predict=64) + + assert (text, tokens) == ("hello", 7) + assert len(calls) == 1 + url, body, _timeout = calls[0] + assert url == "http://x:11434/api/generate" + assert body["model"] == "m1" + assert body["prompt"] == "hi" + assert body["stream"] is False + assert body["options"]["num_predict"] == 64 + assert "keep_alive" in body + + +def test_generate_openai_backend_posts_to_chat_completions(monkeypatch): + monkeypatch.setenv("PW_INFERENCE_BACKEND", "openai") + calls = [] + + def fake_post(url, json=None, timeout=None): + calls.append((url, json, timeout)) + return _FakeResp({ + "choices": [{"message": {"content": " world "}}], + "usage": {"completion_tokens": 12}, + }) + + monkeypatch.setattr(O.requests, "post", fake_post) + text, tokens = O.generate("hi", model="m2", base_url="http://localhost:8080", + temperature=0.5, num_predict=32) + + assert (text, tokens) == ("world", 12) + assert len(calls) == 1 + url, body, _timeout = calls[0] + assert url == "http://localhost:8080/v1/chat/completions" + assert body["model"] == "m2" + assert body["messages"] == [{"role": "user", "content": "hi"}] + assert body["temperature"] == 0.5 + assert body["max_tokens"] == 32 + + +def test_generate_openai_backend_missing_fields_fall_back_to_zero_and_empty(monkeypatch): + monkeypatch.setenv("PW_INFERENCE_BACKEND", "openai") + monkeypatch.setattr(O.requests, "post", + lambda url, json=None, timeout=None: _FakeResp({"choices": [{"message": {}}]})) + text, tokens = O.generate("hi", model="m3", base_url="http://x") + assert (text, tokens) == ("", 0) # mirrors the Ollama path's 0-if-absent contract + + +# ---- resolve_timeout() ------------------------------------------------------------------------- + +def test_resolve_timeout_prefers_env_primary(monkeypatch): + monkeypatch.setenv("PW_JUDGE_TIMEOUT", "12") + monkeypatch.setenv("PW_OLLAMA_TIMEOUT", "99") + assert O.resolve_timeout("PW_JUDGE_TIMEOUT", 300.0) == 12.0 + + +def test_resolve_timeout_falls_back_to_pw_ollama_timeout(monkeypatch): + monkeypatch.delenv("PW_JUDGE_TIMEOUT", raising=False) + monkeypatch.setenv("PW_OLLAMA_TIMEOUT", "45") + assert O.resolve_timeout("PW_JUDGE_TIMEOUT", 300.0) == 45.0 + + +def test_resolve_timeout_falls_back_to_default_when_unset(monkeypatch): + monkeypatch.delenv("PW_JUDGE_TIMEOUT", raising=False) + monkeypatch.delenv("PW_OLLAMA_TIMEOUT", raising=False) + assert O.resolve_timeout("PW_JUDGE_TIMEOUT", 300.0) == 300.0 + + +def test_resolve_timeout_falls_back_to_default_on_invalid_value(monkeypatch): + # an empty/invalid value falls through instead of raising + monkeypatch.delenv("PW_JUDGE_TIMEOUT", raising=False) + monkeypatch.setenv("PW_OLLAMA_TIMEOUT", "not-a-number") + assert O.resolve_timeout(None, 300.0) == 300.0 + + +def test_resolve_timeout_with_no_env_primary_key(monkeypatch): + monkeypatch.setenv("PW_OLLAMA_TIMEOUT", "7") + assert O.resolve_timeout(None, 300.0) == 7.0 + + +# ---- config registration ------------------------------------------------------------------------ + +def test_config_known_includes_inference_backend_keys(): + assert "PW_INFERENCE_BACKEND" in C.KNOWN + assert "PW_INFERENCE_API_BASE" in C.KNOWN