diff --git a/src/physcom/llm/providers/ollama.py b/src/physcom/llm/providers/ollama.py new file mode 100644 index 0000000..8391f54 --- /dev/null +++ b/src/physcom/llm/providers/ollama.py @@ -0,0 +1,73 @@ +"""Ollama LLM provider — local models via the Ollama HTTP API.""" + +from __future__ import annotations + +import json +import re +import urllib.error +import urllib.request + +from physcom.llm.base import LLMProvider +from physcom.llm.prompts import PHYSICS_ESTIMATION_PROMPT, PLAUSIBILITY_REVIEW_PROMPT + + +class OllamaLLMProvider(LLMProvider): + """LLM provider backed by a local Ollama server (see `ollama serve`).""" + + def __init__(self, model: str = "qwen2.5:7b", host: str = "http://localhost:11434") -> None: + self._model = model + self._host = host.rstrip("/") + + def estimate_physics( + self, combination_description: str, metrics: list[str] + ) -> dict[str, float]: + prompt = PHYSICS_ESTIMATION_PROMPT.format( + description=combination_description, + metrics=", ".join(metrics), + ) + text = self._generate(prompt, json_mode=True) + return self._parse_json(text, metrics) + + def review_plausibility( + self, combination_description: str, scores: dict[str, float] + ) -> tuple[str, bool]: + scores_str = "\n".join(f"- {k}: {v:.3f}" for k, v in scores.items()) + prompt = PLAUSIBILITY_REVIEW_PROMPT.format( + description=combination_description, + scores=scores_str, + ) + text = self._generate(prompt, json_mode=False).strip() + return (text, self._parse_verdict(text)) + + def _generate(self, prompt: str, json_mode: bool) -> str: + payload = {"model": self._model, "prompt": prompt, "stream": False} + if json_mode: + payload["format"] = "json" + req = urllib.request.Request( + f"{self._host}/api/generate", + data=json.dumps(payload).encode("utf-8"), + headers={"Content-Type": "application/json"}, + ) + try: + with urllib.request.urlopen(req, timeout=120) as resp: + return json.loads(resp.read())["response"] + except urllib.error.URLError as exc: + raise ConnectionError( + f"Could not reach Ollama at {self._host} (is `ollama serve` running?)" + ) from exc + + def _parse_verdict(self, text: str) -> bool: + """Extract VERDICT: PLAUSIBLE/IMPLAUSIBLE from response; default to True.""" + m = re.search(r"VERDICT:\s*(PLAUSIBLE|IMPLAUSIBLE)", text, re.IGNORECASE) + if m: + return m.group(1).upper() == "PLAUSIBLE" + return True + + def _parse_json(self, text: str, metrics: list[str]) -> dict[str, float]: + """Strip markdown fences and parse JSON; fall back to 0.5 per metric on error.""" + text = re.sub(r"```(?:json)?\s*", "", text).strip().rstrip("`").strip() + try: + data = json.loads(text) + return {k: float(v) for k, v in data.items() if k in metrics} + except (json.JSONDecodeError, ValueError, TypeError): + return {m: 0.5 for m in metrics} diff --git a/src/physcom/llm/registry.py b/src/physcom/llm/registry.py index a508860..0f4a719 100644 --- a/src/physcom/llm/registry.py +++ b/src/physcom/llm/registry.py @@ -10,9 +10,11 @@ from physcom.llm.base import LLMProvider def build_llm_provider() -> LLMProvider | None: """Return an LLMProvider based on env vars, or None if not configured. - LLM_PROVIDER — provider name ('gemini'; more can be added) + LLM_PROVIDER — provider name ('gemini', 'ollama'; more can be added) GEMINI_API_KEY — required when LLM_PROVIDER=gemini GEMINI_MODEL — optional Gemini model name (default: gemini-2.0-flash) + OLLAMA_MODEL — optional Ollama model name (default: qwen2.5:7b) + OLLAMA_HOST — optional Ollama server URL (default: http://localhost:11434) """ provider = os.environ.get("LLM_PROVIDER", "").lower().strip() @@ -27,4 +29,10 @@ def build_llm_provider() -> LLMProvider | None: from physcom.llm.providers.gemini import GeminiLLMProvider return GeminiLLMProvider(api_key=api_key, model=model) - raise ValueError(f"Unknown LLM_PROVIDER: {provider!r}. Supported: gemini") + if provider == "ollama": + model = os.environ.get("OLLAMA_MODEL", "qwen2.5:7b") + host = os.environ.get("OLLAMA_HOST", "http://localhost:11434") + from physcom.llm.providers.ollama import OllamaLLMProvider + return OllamaLLMProvider(model=model, host=host) + + raise ValueError(f"Unknown LLM_PROVIDER: {provider!r}. Supported: gemini, ollama") diff --git a/tests/test_ollama_provider.py b/tests/test_ollama_provider.py new file mode 100644 index 0000000..24410a3 --- /dev/null +++ b/tests/test_ollama_provider.py @@ -0,0 +1,43 @@ +"""Tests for the Ollama provider's parsing logic and registry wiring.""" + +from __future__ import annotations + +import pytest + +from physcom.llm.providers.ollama import OllamaLLMProvider + + +@pytest.fixture +def provider(): + return OllamaLLMProvider() + + +def test_parse_json_strips_fences(provider): + text = '```json\n{"power_density": 500.0, "safety": 0.7}\n```' + result = provider._parse_json(text, ["power_density", "safety"]) + assert result == {"power_density": 500.0, "safety": 0.7} + + +def test_parse_json_falls_back_on_invalid(provider): + result = provider._parse_json("not json", ["power_density", "safety"]) + assert result == {"power_density": 0.5, "safety": 0.5} + + +def test_parse_verdict_plausible(provider): + assert provider._parse_verdict("blah blah\nVERDICT: PLAUSIBLE") is True + + +def test_parse_verdict_implausible(provider): + assert provider._parse_verdict("blah blah\nVERDICT: IMPLAUSIBLE") is False + + +def test_registry_builds_ollama_provider(monkeypatch): + from physcom.llm.registry import build_llm_provider + + monkeypatch.setenv("LLM_PROVIDER", "ollama") + monkeypatch.setenv("OLLAMA_MODEL", "phi4:14b") + monkeypatch.setenv("OLLAMA_HOST", "http://example:1234") + provider = build_llm_provider() + assert isinstance(provider, OllamaLLMProvider) + assert provider._model == "phi4:14b" + assert provider._host == "http://example:1234"