fix: honor hyde_enabled toggle and record ping failures before client creation (Task 6 review)
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
This commit is contained in:
@@ -169,6 +169,13 @@ def _build_model_status(role: str) -> dict[str, Any]:
|
|||||||
# Config always wins: report "disabled" even if the reranker was
|
# Config always wins: report "disabled" even if the reranker was
|
||||||
# enabled and called successfully earlier in this process's life.
|
# enabled and called successfully earlier in this process's life.
|
||||||
status = "disabled"
|
status = "disabled"
|
||||||
|
elif role == "hyde_llm":
|
||||||
|
enabled = settings.hyde_enabled
|
||||||
|
if not enabled:
|
||||||
|
# Same "config always wins" override as the reranker branch above:
|
||||||
|
# report "disabled" even if HyDE ran successfully before being
|
||||||
|
# turned off in settings during this process's life.
|
||||||
|
status = "disabled"
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"role": role,
|
"role": role,
|
||||||
@@ -199,7 +206,15 @@ async def get_model_statuses():
|
|||||||
async def _ping_main_or_hyde(role: str) -> None:
|
async def _ping_main_or_hyde(role: str) -> None:
|
||||||
"""Send one minimal chat completion to the LLM configured for `role`."""
|
"""Send one minimal chat completion to the LLM configured for `role`."""
|
||||||
provider, model = _resolve_role_provider_model(role)
|
provider, model = _resolve_role_provider_model(role)
|
||||||
client = get_llm_client(provider=provider, model=model)
|
try:
|
||||||
|
client = get_llm_client(provider=provider, model=model)
|
||||||
|
except Exception as exc: # noqa: BLE001 - record, then re-raise so gather() still isolates this ping
|
||||||
|
# get_llm_client() can fail before any TrackedLLMClient exists to
|
||||||
|
# record the outcome itself (e.g. missing API key, unsupported
|
||||||
|
# provider string), so record the failure here directly, otherwise it
|
||||||
|
# would be invisible on the /status/models page afterward.
|
||||||
|
get_model_usage_tracker().record(provider=provider, model=model, success=False, error=str(exc))
|
||||||
|
raise
|
||||||
await asyncio.to_thread(client.chat, [{"role": "user", "content": "ping"}], max_tokens=1)
|
await asyncio.to_thread(client.chat, [{"role": "user", "content": "ping"}], max_tokens=1)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -58,6 +58,23 @@ def test_get_models_reflects_recorded_usage(client):
|
|||||||
assert main_row["status"] == "ok"
|
assert main_row["status"] == "ok"
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_models_hyde_llm_disabled_forces_disabled_status(client):
|
||||||
|
"""settings.hyde_enabled=False must force hyde_llm to enabled=False/status='disabled',
|
||||||
|
mirroring the reranker override, even if HyDE previously ran successfully."""
|
||||||
|
from app.config.settings import settings
|
||||||
|
get_model_usage_tracker().record(
|
||||||
|
provider=settings.hyde_llm_provider or settings.llm_provider,
|
||||||
|
model=settings.hyde_llm_model or settings.llm_model,
|
||||||
|
success=True,
|
||||||
|
)
|
||||||
|
with patch.object(settings, "hyde_enabled", False):
|
||||||
|
resp = client.get("/api/v1/status/models")
|
||||||
|
assert resp.status_code == 200
|
||||||
|
hyde_row = next(m for m in resp.json()["models"] if m["role"] == "hyde_llm")
|
||||||
|
assert hyde_row["enabled"] is False
|
||||||
|
assert hyde_row["status"] == "disabled"
|
||||||
|
|
||||||
|
|
||||||
def test_ping_models_calls_each_enabled_model_once(client):
|
def test_ping_models_calls_each_enabled_model_once(client):
|
||||||
"""POST /status/models/ping must invoke chat()/embed_query() and return fresh statuses."""
|
"""POST /status/models/ping must invoke chat()/embed_query() and return fresh statuses."""
|
||||||
mock_llm_response = LLMResponse(content="pong", model="test-model", usage={"total_tokens": 1})
|
mock_llm_response = LLMResponse(content="pong", model="test-model", usage={"total_tokens": 1})
|
||||||
@@ -90,3 +107,27 @@ def test_ping_models_survives_one_model_failing(client):
|
|||||||
|
|
||||||
assert resp.status_code == 200
|
assert resp.status_code == 200
|
||||||
mock_embedding.embed_query.assert_called_once()
|
mock_embedding.embed_query.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
def test_ping_records_get_llm_client_failure_instead_of_dropping_it(client):
|
||||||
|
"""A get_llm_client() failure (raised before any TrackedLLMClient exists) must still
|
||||||
|
be recorded into the tracker, so it is visible afterwards via _build_model_status()
|
||||||
|
instead of being silently discarded by asyncio.gather(return_exceptions=True)."""
|
||||||
|
mock_embedding = MagicMock()
|
||||||
|
mock_embedding.embed_query.return_value = [0.1]
|
||||||
|
|
||||||
|
with patch("app.api.routes.status.get_llm_client", side_effect=RuntimeError("missing api key")), \
|
||||||
|
patch("app.api.routes.status.get_embedding_provider", return_value=mock_embedding), \
|
||||||
|
patch("app.api.routes.status.get_reranker", return_value=None):
|
||||||
|
resp = client.post("/api/v1/status/models/ping")
|
||||||
|
|
||||||
|
assert resp.status_code == 200
|
||||||
|
body = resp.json()
|
||||||
|
main_row = next(m for m in body["models"] if m["role"] == "main_llm")
|
||||||
|
hyde_row = next(m for m in body["models"] if m["role"] == "hyde_llm")
|
||||||
|
assert main_row["status"] == "error"
|
||||||
|
assert main_row["call_count_error"] == 1
|
||||||
|
assert main_row["last_error"] == "missing api key"
|
||||||
|
assert hyde_row["status"] == "error"
|
||||||
|
assert hyde_row["call_count_error"] == 1
|
||||||
|
assert hyde_row["last_error"] == "missing api key"
|
||||||
|
|||||||
Reference in New Issue
Block a user