fix(ai): keep GCP Ollama lane on safe models
This commit is contained in:
@@ -23,6 +23,7 @@ from src.models.nvidia import (
|
||||
from src.services.nvidia_provider import (
|
||||
HIGH_RISK_TOOLS,
|
||||
NvidiaProvider,
|
||||
OllamaToolProvider,
|
||||
create_tool_definition,
|
||||
get_nvidia_provider,
|
||||
reset_nvidia_provider,
|
||||
@@ -286,6 +287,58 @@ class TestProtocolCompliance:
|
||||
assert hasattr(INvidiaProvider, "close")
|
||||
|
||||
|
||||
class TestOllamaToolProviderRouting:
|
||||
"""Ollama tool-calling must not pollute the GCP alert lane."""
|
||||
|
||||
def test_base_url_uses_hermes_resolver_lane(self, monkeypatch):
|
||||
from src.services import nvidia_provider as nvidia_provider_module
|
||||
|
||||
captured_workloads = []
|
||||
|
||||
def fake_resolve(workload):
|
||||
captured_workloads.append(workload)
|
||||
return "http://local-111:11434"
|
||||
|
||||
monkeypatch.setattr(nvidia_provider_module, "resolve_ollama_endpoint", fake_resolve)
|
||||
|
||||
provider = OllamaToolProvider()
|
||||
|
||||
assert provider._base_url() == "http://local-111:11434"
|
||||
assert captured_workloads == ["hermes"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_uses_local_hermes_lane(self, monkeypatch):
|
||||
from src.services import nvidia_provider as nvidia_provider_module
|
||||
|
||||
class _FakeResponse:
|
||||
status_code = 200
|
||||
|
||||
class _FakeClient:
|
||||
def __init__(self):
|
||||
self.checked_urls = []
|
||||
|
||||
async def get(self, url, **kwargs):
|
||||
self.checked_urls.append(url)
|
||||
return _FakeResponse()
|
||||
|
||||
monkeypatch.setattr(
|
||||
nvidia_provider_module,
|
||||
"resolve_ollama_endpoint",
|
||||
lambda workload: "http://local-111:11434",
|
||||
)
|
||||
|
||||
provider = OllamaToolProvider()
|
||||
client = _FakeClient()
|
||||
|
||||
async def _get_client():
|
||||
return client
|
||||
|
||||
monkeypatch.setattr(provider, "_get_client", _get_client)
|
||||
|
||||
assert await provider.health_check() is True
|
||||
assert client.checked_urls == ["http://local-111:11434/api/tags"]
|
||||
|
||||
|
||||
class TestEdgeCases:
|
||||
"""邊界測試案例 (P2-2)"""
|
||||
|
||||
@@ -406,7 +459,11 @@ class TestAIRouterNvidiaIntegration:
|
||||
|
||||
def test_tool_calling_route(self):
|
||||
"""測試 Tool Calling 路由"""
|
||||
from src.services.ai_router import AIProviderEnum, get_ai_router, reset_ai_router
|
||||
from src.services.ai_router import (
|
||||
AIProviderEnum,
|
||||
get_ai_router,
|
||||
reset_ai_router,
|
||||
)
|
||||
|
||||
reset_ai_router()
|
||||
router = get_ai_router()
|
||||
@@ -424,7 +481,11 @@ class TestAIRouterNvidiaIntegration:
|
||||
|
||||
def test_existing_routing_not_affected(self):
|
||||
"""測試現有路由規則不受影響"""
|
||||
from src.services.ai_router import AIProviderEnum, get_ai_router, reset_ai_router
|
||||
from src.services.ai_router import (
|
||||
AIProviderEnum,
|
||||
get_ai_router,
|
||||
reset_ai_router,
|
||||
)
|
||||
|
||||
reset_ai_router()
|
||||
router = get_ai_router()
|
||||
|
||||
@@ -10,7 +10,7 @@ from src.services.ai_providers.ollama import OllamaGcpBProvider, OllamaProvider
|
||||
|
||||
class _FakeRegistry:
|
||||
def get_model(self, provider: str, use_case: str) -> str:
|
||||
return "qwen2.5:7b-instruct"
|
||||
return "qwen3:14b"
|
||||
|
||||
def get_provider_options(self, provider: str) -> dict[str, Any]:
|
||||
return {"num_predict": 32, "temperature": 0.1, "top_p": 0.9}
|
||||
@@ -33,10 +33,12 @@ class _FakeResponse:
|
||||
class _FakeClient:
|
||||
def __init__(self) -> None:
|
||||
self.posted_urls: list[str] = []
|
||||
self.posted_payloads: list[dict[str, Any]] = []
|
||||
self.checked_urls: list[str] = []
|
||||
|
||||
async def post(self, url: str, **kwargs: Any) -> _FakeResponse:
|
||||
self.posted_urls.append(url)
|
||||
self.posted_payloads.append(kwargs.get("json", {}))
|
||||
return _FakeResponse()
|
||||
|
||||
async def get(self, url: str, **kwargs: Any) -> _FakeResponse:
|
||||
@@ -53,6 +55,7 @@ async def test_ollama_gcp_b_analyze_uses_secondary_url(monkeypatch: pytest.Monke
|
||||
"OLLAMA_SECONDARY_URL",
|
||||
"http://secondary:11436",
|
||||
)
|
||||
monkeypatch.setattr(ollama_module.settings, "ALERT_OLLAMA_MODEL", "gemma3:4b")
|
||||
|
||||
client = _FakeClient()
|
||||
provider = OllamaGcpBProvider()
|
||||
@@ -67,6 +70,57 @@ async def test_ollama_gcp_b_analyze_uses_secondary_url(monkeypatch: pytest.Monke
|
||||
assert result.success is True
|
||||
assert result.provider == "ollama_gcp_b"
|
||||
assert client.posted_urls == ["http://secondary:11436/api/generate"]
|
||||
assert client.posted_payloads[0]["model"] == "gemma3:4b"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ollama_gcp_a_coerces_heavy_diagnose_model_to_alert_model(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(ollama_module, "get_model_registry", lambda: _FakeRegistry())
|
||||
monkeypatch.setattr(ollama_module.settings, "OLLAMA_URL", "http://primary:11435")
|
||||
monkeypatch.setattr(ollama_module.settings, "OLLAMA_SECONDARY_URL", "http://secondary:11436")
|
||||
monkeypatch.setattr(ollama_module.settings, "ALERT_OLLAMA_MODEL", "gemma3:4b")
|
||||
|
||||
client = _FakeClient()
|
||||
provider = OllamaProvider()
|
||||
|
||||
async def _get_client() -> _FakeClient:
|
||||
return client
|
||||
|
||||
monkeypatch.setattr(provider, "_get_client", _get_client)
|
||||
|
||||
result = await provider.analyze("diagnose", context={"task_type": "diagnose"})
|
||||
|
||||
assert result.success is True
|
||||
assert client.posted_urls == ["http://primary:11435/api/generate"]
|
||||
assert client.posted_payloads[0]["model"] == "gemma3:4b"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ollama_gcp_a_can_explicitly_allow_heavy_model(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(ollama_module, "get_model_registry", lambda: _FakeRegistry())
|
||||
monkeypatch.setattr(ollama_module.settings, "OLLAMA_URL", "http://primary:11435")
|
||||
monkeypatch.setattr(ollama_module.settings, "OLLAMA_SECONDARY_URL", "http://secondary:11436")
|
||||
monkeypatch.setattr(ollama_module.settings, "ALERT_OLLAMA_MODEL", "gemma3:4b")
|
||||
|
||||
client = _FakeClient()
|
||||
provider = OllamaProvider()
|
||||
|
||||
async def _get_client() -> _FakeClient:
|
||||
return client
|
||||
|
||||
monkeypatch.setattr(provider, "_get_client", _get_client)
|
||||
|
||||
result = await provider.analyze(
|
||||
"deep diagnose",
|
||||
context={"task_type": "diagnose", "allow_gcp_heavy_model": True},
|
||||
)
|
||||
|
||||
assert result.success is True
|
||||
assert client.posted_payloads[0]["model"] == "qwen3:14b"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
Reference in New Issue
Block a user