local LLM support
This commit is contained in:
73
src/physcom/llm/providers/ollama.py
Normal file
73
src/physcom/llm/providers/ollama.py
Normal file
@@ -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}
|
||||||
@@ -10,9 +10,11 @@ from physcom.llm.base import LLMProvider
|
|||||||
def build_llm_provider() -> LLMProvider | None:
|
def build_llm_provider() -> LLMProvider | None:
|
||||||
"""Return an LLMProvider based on env vars, or None if not configured.
|
"""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_API_KEY — required when LLM_PROVIDER=gemini
|
||||||
GEMINI_MODEL — optional Gemini model name (default: gemini-2.0-flash)
|
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()
|
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
|
from physcom.llm.providers.gemini import GeminiLLMProvider
|
||||||
return GeminiLLMProvider(api_key=api_key, model=model)
|
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")
|
||||||
|
|||||||
43
tests/test_ollama_provider.py
Normal file
43
tests/test_ollama_provider.py
Normal file
@@ -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"
|
||||||
Reference in New Issue
Block a user