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
2 changes: 2 additions & 0 deletions council/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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},
Expand Down
101 changes: 90 additions & 11 deletions council/ollama.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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("/")


Expand All @@ -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:
Expand All @@ -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
Expand All @@ -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)
155 changes: 155 additions & 0 deletions tests/test_ollama.py
Original file line number Diff line number Diff line change
@@ -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
Loading