diff --git a/README.md b/README.md index cdcd304..9c02b7b 100644 --- a/README.md +++ b/README.md @@ -89,6 +89,10 @@ connected through SSH local forwarding. on push and pull request, including `python scripts/check_deployment_safety.py --skip-listen`. + The MCP server also calls `classify_base_url` on start and refuses a + fail-status `OLLAMA_BASE_URL` (wildcard, tunnel, or public IP). Hostname + and private LAN URLs still only warn. + 6. Deployment safety checks (no GPU required; inspects this host only): ```bash diff --git a/spec.md b/spec.md index b563e86..fe9dea7 100644 --- a/spec.md +++ b/spec.md @@ -587,7 +587,8 @@ PYTHONPATH=src python3 scripts/check_deployment_safety.py The checker looks at *this* host only: `OLLAMA_BASE_URL`, model tags, whether `.env` is ignored, placeholder IPs in git, and whether port 11434 is listening on a wildcard. It does **not** scan other machines and it cannot prove weights -are clean. +are clean. The stdio MCP server also calls `classify_base_url` at start and +exits if that check fails; hostname and private LAN URLs still only warn. ### 12.2 Open-weight models diff --git a/src/local_coding_slm/server.py b/src/local_coding_slm/server.py index 7fde678..e8ffddd 100644 --- a/src/local_coding_slm/server.py +++ b/src/local_coding_slm/server.py @@ -14,6 +14,7 @@ sys.path.insert(0, str(_SRC)) from local_coding_slm.ollama_client import ( # noqa: E402 + DEFAULT_BASE_URL, OllamaError, chat, format_user_task, @@ -21,6 +22,7 @@ ) from local_coding_slm.payload import inspect_payload, refusal_message # noqa: E402 from local_coding_slm.prompts import SYSTEM_PROMPTS # noqa: E402 +from local_coding_slm.safety import classify_base_url # noqa: E402 def _load_dotenv() -> None: @@ -153,7 +155,22 @@ def local_review( return _run_tool("local_review", task, files, language, style, model, max_tokens) +def enforce_runtime_base_url(url: str | None = None) -> None: + """Refuse to start when OLLAMA_BASE_URL fails classify_base_url. + + Warn status still starts (hostname and private LAN URLs only warn). + """ + raw = os.environ.get("OLLAMA_BASE_URL", DEFAULT_BASE_URL) if url is None else url + result = classify_base_url(raw) + if result.status == "fail": + print(f"ERROR: {result.message}", file=sys.stderr) + raise SystemExit(1) + if result.status == "warn": + print(f"WARNING: {result.message}", file=sys.stderr) + + def main() -> None: + enforce_runtime_base_url() mcp.run(transport="stdio") diff --git a/tests/test_server_runtime_safety.py b/tests/test_server_runtime_safety.py new file mode 100644 index 0000000..fd0f89d --- /dev/null +++ b/tests/test_server_runtime_safety.py @@ -0,0 +1,102 @@ +"""Server start must call classify_base_url. Leftover #9 CI step stays.""" + +from __future__ import annotations + +import os +import unittest +from pathlib import Path +from unittest.mock import patch + +from local_coding_slm.safety import CheckResult, classify_base_url +from local_coding_slm.server import enforce_runtime_base_url, main + +ROOT = Path(__file__).resolve().parents[1] +SERVER = ROOT / "src" / "local_coding_slm" / "server.py" +WORKFLOW = ROOT / ".github" / "workflows" / "tests.yml" +CI_COMMAND = "python scripts/check_deployment_safety.py --skip-listen" + + +class ServerRuntimeSafetyTests(unittest.TestCase): + def test_server_start_calls_classify_base_url(self) -> None: + passed = CheckResult("base_url", "pass", "OLLAMA_BASE_URL is loopback") + with ( + patch("local_coding_slm.server.mcp.run") as run, + patch( + "local_coding_slm.server.classify_base_url", + return_value=passed, + ) as classify, + patch.dict( + os.environ, + {"OLLAMA_BASE_URL": "http://127.0.0.1:11434"}, + clear=False, + ), + ): + main() + classify.assert_called() + self.assertEqual(classify.call_args.args[0], "http://127.0.0.1:11434") + run.assert_called_once_with(transport="stdio") + + def test_server_start_refuses_fail_base_url(self) -> None: + self.assertEqual(classify_base_url("http://8.8.8.8:11434").status, "fail") + with ( + patch("local_coding_slm.server.mcp.run") as run, + patch.dict( + os.environ, + {"OLLAMA_BASE_URL": "http://8.8.8.8:11434"}, + clear=False, + ), + ): + with self.assertRaises(SystemExit) as raised: + main() + self.assertEqual(raised.exception.code, 1) + run.assert_not_called() + + def test_server_start_allows_loopback(self) -> None: + with ( + patch("local_coding_slm.server.mcp.run") as run, + patch.dict( + os.environ, + {"OLLAMA_BASE_URL": "http://127.0.0.1:11434"}, + clear=False, + ), + ): + main() + run.assert_called_once_with(transport="stdio") + + def test_server_start_hostname_warn_still_starts(self) -> None: + # Leftover #11 stays warn; runtime start must not fail-close hostnames. + url = "https://my-ollama.evil.com" + self.assertEqual(classify_base_url(url).status, "warn") + with ( + patch("local_coding_slm.server.mcp.run") as run, + patch.dict(os.environ, {"OLLAMA_BASE_URL": url}, clear=False), + ): + main() + run.assert_called_once_with(transport="stdio") + + def test_server_imports_classify_base_url(self) -> None: + text = SERVER.read_text(encoding="utf-8") + self.assertIn("from local_coding_slm.safety import classify_base_url", text) + self.assertIn("classify_base_url", text) + + def test_leftover_9_ci_deployment_safety_step_stays(self) -> None: + text = WORKFLOW.read_text(encoding="utf-8") + self.assertIn("- name: Deployment safety", text) + self.assertIn(CI_COMMAND, text) + + def test_enforce_runtime_base_url_uses_default_when_unset(self) -> None: + env = {key: value for key, value in os.environ.items() if key != "OLLAMA_BASE_URL"} + with ( + patch.dict(os.environ, env, clear=True), + patch( + "local_coding_slm.server.classify_base_url", + return_value=CheckResult("base_url", "pass", "loopback"), + ) as classify, + ): + enforce_runtime_base_url() + classify.assert_called_once() + self.assertIn("127.0.0.1", classify.call_args.args[0]) + + +if __name__ == "__main__": + unittest.main()