Files
GarminHealthLab/backend/tests/test_ai.py
ericwyuan c83340742c [阶段4.1] AI 健康建议 - 多模型可切换 + 大上下文 + 失败兜底
services/ai.py:
- 模型目录(catalog)按短 id 索引,业务代码不感知厂商
  gemini-flash (Google, 1M 上下文)
  llama-70b / qwen-72b / deepseek-r1 (NVIDIA NIM, 128k)
  仅注册纯文本模型,不含视觉模型
- 两个 provider: GeminiProvider、OpenAICompatProvider
  (后者兼容 NVIDIA NIM / Ollama / vLLM)
- 大上下文: 每日指标序列化为 CSV 而非 JSON,同样的数据 token 数约为
  1/4,一整年历史仍远小于最小的 128k 窗口;按 AI_DAY_BUDGET 截断
- 兜底链: 首选模型超时/报错/返回无法解析的文本时自动降级到下一个,
  meta.fallbackFrom 记录降级路径
- 响应解析容忍 markdown 代码块包裹和 JSON 前的多余句子

services/analysis.py:
- get_ai_recommendations(): 所有模型都失败时回落到规则引擎,
  端点始终 200,meta.source 区分 ai / rules

routes/analysis.py:
- GET /api/analysis/models 列出模型及各自是否已配置密钥
- GET /api/analysis/ai-recommendations?model=&days=

tests/test_ai.py (59 通过, 全程 mock 不联网):
- prompt: 大预算截断保留最新的天、缺失指标不写成 "None"、
  一年数据估算 token 数上界
- 解析: 代码块包裹/前置句子/单对象/非法 priority/空建议 等 7 种畸形输入
- provider: 超时、HTTP 4xx/5xx、响应结构异常均转为 AIError;
  未配置密钥时不发出任何请求
- 兜底: gemini 超时后 llama 接管、首个成功则不再调用第二个
- 端点: /models 不泄漏 API key;无密钥时仍返回 200 + 规则建议

密钥一律从环境变量读取,.env.example 只留空占位符。

全量: 161 passed, 1 skipped

Co-Authored-By: Claude Haiku 4.5 <noreply@anthropic.com>
2026-08-23 12:38:19 +08:00

461 lines
18 KiB
Python

"""
Unit tests for the multi-provider LLM layer.
Every HTTP call is mocked — the suite never touches the network and never
needs a real API key.
"""
import json
import pytest
import requests
from services import ai as ai_svc
from services import analysis as analysis_svc
VALID_REPLY = json.dumps(
[
{"category": "睡眠", "recommendation": "固定就寝时间,目标 7-8 小时。",
"priority": "high", "basedOn": ["sleep_duration"]},
{"category": "运动", "recommendation": "每天增加 20 分钟快走。",
"priority": "medium", "basedOn": ["steps"]},
],
ensure_ascii=False,
)
SUMMARY = [
{"date": "2026-08-20", "steps": 6500, "heartRate": 70,
"heartRateVariability": 45, "stress": 55, "caloriesBurned": 260,
"sleep": {"duration": 6, "quality": 80}},
{"date": "2026-08-21", "steps": 9000, "heartRate": 62,
"heartRateVariability": 46, "stress": 40, "caloriesBurned": 360,
"sleep": {"duration": 8, "quality": 79}},
]
class FakeResponse:
def __init__(self, status_code=200, payload=None, text=""):
self.status_code = status_code
self._payload = payload
self.text = text or json.dumps(payload or {})
def json(self):
if self._payload is None:
raise ValueError("no json")
return self._payload
def gemini_payload(text):
return {"candidates": [{"content": {"parts": [{"text": text}]}}]}
def openai_payload(text):
return {"choices": [{"message": {"content": text}}]}
@pytest.fixture
def keys(monkeypatch):
"""Pretend both vendors are configured."""
monkeypatch.setenv("GEMINI_API_KEY", "test-gemini-key")
monkeypatch.setenv("NVIDIA_API_KEY", "test-nvidia-key")
return True
@pytest.fixture
def no_keys(monkeypatch):
monkeypatch.delenv("GEMINI_API_KEY", raising=False)
monkeypatch.delenv("NVIDIA_API_KEY", raising=False)
return True
# --- prompt construction ----------------------------------------------------
class TestBuildPrompt:
def test_includes_every_day_as_a_csv_row(self):
prompt = ai_svc.build_prompt(SUMMARY)
assert "2026-08-20" in prompt and "2026-08-21" in prompt
assert "共 2 天" in prompt
def test_uses_csv_not_json(self):
"""CSV keeps a year of history affordable; JSON would not."""
prompt = ai_svc.build_prompt(SUMMARY)
assert "6500,45" in prompt.replace(" ", "") or "6500" in prompt
assert '"steps":' not in prompt
def test_missing_metrics_become_empty_cells_not_the_word_none(self):
prompt = ai_svc.build_prompt([{"date": "2026-08-20", "steps": None}])
assert "None" not in prompt
def test_sleep_is_flattened_into_columns(self):
prompt = ai_svc.build_prompt(SUMMARY)
assert "sleep_h,sleep_q" in prompt
def test_day_budget_trims_to_the_most_recent_days(self):
many = [{"date": f"2026-01-{d:02d}", "steps": d} for d in range(1, 32)]
prompt = ai_svc.build_prompt(many, day_budget=5)
assert "共 5 天" in prompt
assert "2026-01-31" in prompt, "must keep the newest days"
assert "2026-01-01" not in prompt, "must drop the oldest days"
def test_activities_included_when_supplied(self):
prompt = ai_svc.build_prompt(
SUMMARY, [{"activity_type": "running", "distance": 5.0}]
)
assert "running" in prompt
def test_activities_capped(self):
acts = [{"activity_type": f"run{i}"} for i in range(500)]
prompt = ai_svc.build_prompt(SUMMARY, acts)
assert "共 200 条" in prompt
def test_prompt_forbids_fabricating_numbers(self):
assert "不要编造" in ai_svc.build_prompt(SUMMARY)
def test_prompt_disclaims_medical_advice(self):
assert "不是医生" in ai_svc.build_prompt(SUMMARY)
def test_a_year_of_data_stays_compact(self):
year = [
{"date": f"2026-{m:02d}-{d:02d}", "steps": 8000, "heartRate": 60,
"sleep": {"duration": 7, "quality": 80}}
for m in range(1, 13) for d in range(1, 29)
]
prompt = ai_svc.build_prompt(year)
# ~4 chars/token: a year must stay far under even the smallest window.
assert len(prompt) / 4 < 50_000
# --- response parsing -------------------------------------------------------
class TestParseRecommendations:
def test_plain_json_array(self):
recs = ai_svc.parse_recommendations(VALID_REPLY)
assert len(recs) == 2
assert recs[0]["category"] == "睡眠"
def test_markdown_fenced_json(self):
recs = ai_svc.parse_recommendations(f"```json\n{VALID_REPLY}\n```")
assert len(recs) == 2
def test_json_with_a_preamble_sentence(self):
recs = ai_svc.parse_recommendations(f"好的,分析结果如下:\n{VALID_REPLY}")
assert len(recs) == 2
def test_single_object_is_wrapped(self):
recs = ai_svc.parse_recommendations(
'{"category":"睡眠","recommendation":"早点睡","priority":"high"}'
)
assert len(recs) == 1
def test_results_are_sorted_by_priority(self):
reply = json.dumps([
{"category": "a", "recommendation": "low one", "priority": "low"},
{"category": "b", "recommendation": "high one", "priority": "high"},
{"category": "c", "recommendation": "medium one", "priority": "medium"},
])
assert [r["priority"] for r in ai_svc.parse_recommendations(reply)] == [
"high", "medium", "low"
]
def test_invalid_priority_defaults_to_medium(self):
reply = json.dumps([
{"category": "a", "recommendation": "x", "priority": "URGENT!!"}
])
assert ai_svc.parse_recommendations(reply)[0]["priority"] == "medium"
def test_entries_without_recommendation_text_are_dropped(self):
reply = json.dumps([
{"category": "a", "recommendation": ""},
{"category": "b", "recommendation": "keep me"},
])
recs = ai_svc.parse_recommendations(reply)
assert len(recs) == 1 and recs[0]["recommendation"] == "keep me"
def test_non_list_based_on_is_normalised(self):
reply = json.dumps([
{"category": "a", "recommendation": "x", "basedOn": "steps"}
])
assert ai_svc.parse_recommendations(reply)[0]["basedOn"] == []
def test_results_are_tagged_as_ai_generated(self):
assert all(r["source"] == "ai" for r in ai_svc.parse_recommendations(VALID_REPLY))
@pytest.mark.parametrize(
"reply", ["", " ", "抱歉,我无法回答。", "[", "null", "[]", "[1,2,3]"]
)
def test_unusable_replies_raise_aierror(self, reply):
with pytest.raises(ai_svc.AIError):
ai_svc.parse_recommendations(reply)
# --- providers --------------------------------------------------------------
class TestGeminiProvider:
def test_successful_call(self, keys, monkeypatch):
captured = {}
def fake_post(url, **kwargs):
captured["url"] = url
captured["headers"] = kwargs.get("headers", {})
captured["json"] = kwargs.get("json")
return FakeResponse(200, gemini_payload("hello"))
monkeypatch.setattr(requests, "post", fake_post)
out = ai_svc.CATALOG["gemini-flash"].generate("prompt text")
assert out == "hello"
assert "gemini-flash-latest:generateContent" in captured["url"]
assert captured["headers"]["X-goog-api-key"] == "test-gemini-key"
assert captured["json"]["contents"][0]["parts"][0]["text"] == "prompt text"
def test_http_error_becomes_aierror(self, keys, monkeypatch):
monkeypatch.setattr(
requests, "post", lambda *a, **k: FakeResponse(429, text="rate limited")
)
with pytest.raises(ai_svc.AIError, match="429"):
ai_svc.CATALOG["gemini-flash"].generate("p")
def test_timeout_becomes_aierror(self, keys, monkeypatch):
def boom(*a, **k):
raise requests.Timeout("timed out")
monkeypatch.setattr(requests, "post", boom)
with pytest.raises(ai_svc.AIError, match="请求失败"):
ai_svc.CATALOG["gemini-flash"].generate("p")
def test_unexpected_shape_becomes_aierror(self, keys, monkeypatch):
monkeypatch.setattr(
requests, "post", lambda *a, **k: FakeResponse(200, {"unexpected": True})
)
with pytest.raises(ai_svc.AIError, match="响应格式异常"):
ai_svc.CATALOG["gemini-flash"].generate("p")
def test_missing_key_raises_before_any_request(self, no_keys, monkeypatch):
def boom(*a, **k):
raise AssertionError("must not issue a request without a key")
monkeypatch.setattr(requests, "post", boom)
with pytest.raises(ai_svc.AIError, match="GEMINI_API_KEY"):
ai_svc.CATALOG["gemini-flash"].generate("p")
class TestOpenAICompatProvider:
def test_successful_call(self, keys, monkeypatch):
captured = {}
def fake_post(url, **kwargs):
captured["url"] = url
captured["headers"] = kwargs.get("headers", {})
captured["json"] = kwargs.get("json")
return FakeResponse(200, openai_payload("hi"))
monkeypatch.setattr(requests, "post", fake_post)
out = ai_svc.CATALOG["llama-70b"].generate("prompt text")
assert out == "hi"
assert captured["url"].endswith("/chat/completions")
assert captured["headers"]["Authorization"] == "Bearer test-nvidia-key"
assert captured["json"]["model"] == "meta/llama-3.3-70b-instruct"
def test_http_error_becomes_aierror(self, keys, monkeypatch):
monkeypatch.setattr(
requests, "post", lambda *a, **k: FakeResponse(500, text="boom")
)
with pytest.raises(ai_svc.AIError, match="500"):
ai_svc.CATALOG["llama-70b"].generate("p")
# --- catalog & chain --------------------------------------------------------
class TestCatalog:
def test_all_models_listed(self, keys):
assert {m["id"] for m in ai_svc.list_models()} == {
"gemini-flash", "llama-70b", "qwen-72b", "deepseek-r1"
}
def test_configured_flag_tracks_the_environment(self, no_keys, monkeypatch):
assert all(not m["configured"] for m in ai_svc.list_models())
monkeypatch.setenv("GEMINI_API_KEY", "k")
by_id = {m["id"]: m for m in ai_svc.list_models()}
assert by_id["gemini-flash"]["configured"] is True
assert by_id["llama-70b"]["configured"] is False
def test_every_model_declares_a_large_window(self):
assert all(m["contextWindow"] >= 128_000 for m in ai_svc.list_models())
def test_no_vision_models_registered(self):
assert not any("vision" in m["model"] for m in ai_svc.list_models())
class TestResolveChain:
def test_preferred_model_goes_first(self, keys):
assert ai_svc.resolve_chain("qwen-72b")[0] == "qwen-72b"
def test_chain_has_no_duplicates(self, keys):
chain = ai_svc.resolve_chain("gemini-flash")
assert len(chain) == len(set(chain))
def test_unconfigured_models_are_skipped(self, monkeypatch):
monkeypatch.setenv("GEMINI_API_KEY", "k")
monkeypatch.delenv("NVIDIA_API_KEY", raising=False)
assert ai_svc.resolve_chain() == ["gemini-flash"]
def test_unknown_model_raises(self, keys):
with pytest.raises(ai_svc.AIError, match="未知模型"):
ai_svc.resolve_chain("gpt-nonexistent")
def test_no_credentials_raises_with_actionable_message(self, no_keys):
with pytest.raises(ai_svc.AIError, match="GEMINI_API_KEY"):
ai_svc.resolve_chain()
# --- generate + fallback ----------------------------------------------------
class TestGenerate:
def test_returns_recommendations_and_meta(self, keys, monkeypatch):
monkeypatch.setattr(
requests, "post", lambda *a, **k: FakeResponse(200, gemini_payload(VALID_REPLY))
)
recs, meta = ai_svc.generate(SUMMARY)
assert len(recs) == 2
assert meta["model"] == "gemini-flash"
assert meta["days"] == 2
assert meta["fallbackFrom"] == []
def test_falls_back_to_the_next_model(self, keys, monkeypatch):
calls = []
def fake_post(url, **kwargs):
calls.append(url)
if "generativelanguage" in url:
raise requests.Timeout("gemini down")
return FakeResponse(200, openai_payload(VALID_REPLY))
monkeypatch.setattr(requests, "post", fake_post)
recs, meta = ai_svc.generate(SUMMARY)
assert len(recs) == 2
assert meta["model"] == "llama-70b"
assert meta["fallbackFrom"] == ["gemini-flash"]
assert len(calls) == 2
def test_falls_back_when_a_model_returns_unparseable_text(self, keys, monkeypatch):
def fake_post(url, **kwargs):
if "generativelanguage" in url:
return FakeResponse(200, gemini_payload("抱歉,我帮不了你。"))
return FakeResponse(200, openai_payload(VALID_REPLY))
monkeypatch.setattr(requests, "post", fake_post)
_, meta = ai_svc.generate(SUMMARY)
assert meta["model"] == "llama-70b"
def test_raises_when_every_model_fails(self, keys, monkeypatch):
def boom(*a, **k):
raise requests.Timeout("all down")
monkeypatch.setattr(requests, "post", boom)
with pytest.raises(ai_svc.AIError, match="所有模型均失败"):
ai_svc.generate(SUMMARY)
def test_preferred_model_is_honoured(self, keys, monkeypatch):
seen = {}
def fake_post(url, **kwargs):
seen["model"] = (kwargs.get("json") or {}).get("model")
return FakeResponse(200, openai_payload(VALID_REPLY))
monkeypatch.setattr(requests, "post", fake_post)
_, meta = ai_svc.generate(SUMMARY, preferred_model="qwen-72b")
assert meta["model"] == "qwen-72b"
assert seen["model"] == "qwen/qwen2.5-72b-instruct"
def test_no_second_call_after_the_first_succeeds(self, keys, monkeypatch):
calls = []
def fake_post(url, **kwargs):
calls.append(url)
return FakeResponse(200, gemini_payload(VALID_REPLY))
monkeypatch.setattr(requests, "post", fake_post)
ai_svc.generate(SUMMARY)
assert len(calls) == 1
# --- service + endpoint integration -----------------------------------------
class TestAiRecommendationsService:
def test_uses_the_rule_engine_when_there_is_no_data(self, db, user, keys):
out = analysis_svc.get_ai_recommendations(user["id"])
assert out["meta"]["source"] == "rules"
assert out["recommendations"][0]["id"] == "no-data"
def test_returns_ai_results_when_a_model_answers(
self, seed_health, user, keys, monkeypatch
):
seed_health([{"date": "2026-08-20", "steps": 5000}])
monkeypatch.setattr(
requests, "post", lambda *a, **k: FakeResponse(200, gemini_payload(VALID_REPLY))
)
out = analysis_svc.get_ai_recommendations(user["id"])
assert out["meta"]["source"] == "ai"
assert len(out["recommendations"]) == 2
def test_degrades_to_rules_when_all_models_fail(
self, seed_health, user, keys, monkeypatch
):
seed_health([{"date": "2026-08-20", "steps": 5000}])
def boom(*a, **k):
raise requests.Timeout("down")
monkeypatch.setattr(requests, "post", boom)
out = analysis_svc.get_ai_recommendations(user["id"])
assert out["meta"]["source"] == "rules"
assert "所有模型均失败" in out["meta"]["reason"]
assert out["recommendations"], "must still return rule-based advice"
def test_degrades_to_rules_when_no_key_is_configured(
self, seed_health, user, no_keys
):
seed_health([{"date": "2026-08-20", "steps": 5000}])
out = analysis_svc.get_ai_recommendations(user["id"])
assert out["meta"]["source"] == "rules"
assert "GEMINI_API_KEY" in out["meta"]["reason"]
class TestEndpoints:
def test_models_requires_auth(self, client):
assert client.get("/api/analysis/models").status_code == 401
def test_ai_recommendations_requires_auth(self, client):
assert client.get("/api/analysis/ai-recommendations").status_code == 401
def test_models_endpoint_lists_catalog(self, client, auth, keys):
r = client.get("/api/analysis/models", headers=auth)
assert r.status_code == 200
assert {m["id"] for m in r.get_json()} >= {"gemini-flash", "llama-70b"}
def test_models_endpoint_never_leaks_api_keys(self, client, auth, keys):
body = client.get("/api/analysis/models", headers=auth).get_data(as_text=True)
assert "test-gemini-key" not in body
assert "test-nvidia-key" not in body
def test_ai_endpoint_returns_200_even_with_no_models(self, client, auth, no_keys):
r = client.get("/api/analysis/ai-recommendations", headers=auth)
assert r.status_code == 200
assert r.get_json()["meta"]["source"] == "rules"
def test_ai_endpoint_passes_model_param_through(
self, client, auth, seed_health, keys, monkeypatch
):
seed_health([{"date": "2026-08-20", "steps": 5000}])
monkeypatch.setattr(
requests, "post", lambda *a, **k: FakeResponse(200, openai_payload(VALID_REPLY))
)
r = client.get(
"/api/analysis/ai-recommendations?model=qwen-72b", headers=auth
)
assert r.get_json()["meta"]["model"] == "qwen-72b"
def test_unknown_model_param_degrades_to_rules(
self, client, auth, seed_health, keys
):
seed_health([{"date": "2026-08-20", "steps": 5000}])
r = client.get("/api/analysis/ai-recommendations?model=bogus", headers=auth)
assert r.status_code == 200
assert r.get_json()["meta"]["source"] == "rules"