diff --git a/app.py b/app.py index 463f7e1..063bdd5 100644 --- a/app.py +++ b/app.py @@ -289,7 +289,10 @@ def _openai_client_for(provider: str, base_url: str | None): from openai import OpenAI final_base_url = resolve_provider_base_url(provider, base_url) - api_key = os.environ.get("OPENAI_API_KEY", "dev-local") + if provider and provider.lower() == "lmstudio": + api_key = os.environ.get("OPENAI_API_KEY") or "lmstudio-local" + else: + api_key = os.environ.get("OPENAI_API_KEY", "dev-local") return OpenAI(api_key=api_key, base_url=final_base_url) diff --git a/tests/contracts/test_control_plane_endpoints.py b/tests/contracts/test_control_plane_endpoints.py index d991b33..a83bfcc 100644 --- a/tests/contracts/test_control_plane_endpoints.py +++ b/tests/contracts/test_control_plane_endpoints.py @@ -1,6 +1,7 @@ import pytest +from types import SimpleNamespace -from app import app +from app import app, _openai_client_for pytestmark = pytest.mark.contracts @@ -36,3 +37,64 @@ def test_ai_test_provider_mock(): r = client.post("/api/ai/test_provider", json={"provider": "mock", "model": "mock-model"}) assert r.status_code == 200 assert r.get_json()["ok"] is True + + +def test_ai_test_provider_lmstudio_allows_empty_env_api_key(monkeypatch): + captured = {} + + class _FakeModels: + @staticmethod + def list(): + return SimpleNamespace(data=[SimpleNamespace(id="local-model")]) + + class _FakeCompletions: + @staticmethod + def create(**_kwargs): + return SimpleNamespace( + choices=[SimpleNamespace(message=SimpleNamespace(content="ok"))] + ) + + class _FakeChat: + completions = _FakeCompletions() + + class _FakeOpenAI: + def __init__(self, api_key, base_url): + captured["api_key"] = api_key + captured["base_url"] = base_url + self.models = _FakeModels() + self.chat = _FakeChat() + + monkeypatch.setenv("OPENAI_API_KEY", "") + monkeypatch.setattr("openai.OpenAI", _FakeOpenAI) + + with app.test_client() as client: + r = client.post( + "/api/ai/test_provider", + json={ + "provider": "lmstudio", + "base_url": "http://localhost:1234/v1", + "model": "local-model", + }, + ) + assert r.status_code == 200 + assert r.get_json()["ok"] is True + assert captured["api_key"] == "lmstudio-local" + + +@pytest.mark.parametrize("provider", ["openai", "ollama", "openrouter"]) +def test_openai_client_for_non_lmstudio_keeps_existing_env_behavior( + monkeypatch, provider +): + captured = {} + + class _FakeOpenAI: + def __init__(self, api_key, base_url): + captured["api_key"] = api_key + captured["base_url"] = base_url + + monkeypatch.setenv("OPENAI_API_KEY", "") + monkeypatch.setattr("openai.OpenAI", _FakeOpenAI) + + _openai_client_for(provider, None) + + assert captured["api_key"] == ""