orchestrator: 全专家失败时不调用 aggregator provider
This commit is contained in:
@@ -259,7 +259,8 @@ async def predict_match_multi(
|
||||
"match %s: 所有 %d 位专家均无有效数据,跳过终裁,标记 degraded",
|
||||
match_id, len(reports),
|
||||
)
|
||||
aggregator_provider = await _agent_provider("aggregator", tier="aggregator")
|
||||
# 无有效专家时不调用 aggregator provider,避免多余开销
|
||||
# model 使用 settings 默认值占位(无实际 LLM 调用)
|
||||
final = {
|
||||
"pred_home_goals": None,
|
||||
"pred_away_goals": None,
|
||||
@@ -272,6 +273,7 @@ async def predict_match_multi(
|
||||
}
|
||||
agg_prompt_tokens = 0
|
||||
agg_completion_tokens = 0
|
||||
aggregator_model = settings.LLM_MODEL # 占位,无实际 LLM 调用
|
||||
|
||||
latency_ms = int((time.perf_counter() - start) * 1000)
|
||||
|
||||
@@ -297,17 +299,19 @@ async def predict_match_multi(
|
||||
raise RuntimeError(f"终裁输出校验失败: {e}")
|
||||
agent_weights = validate_agent_weights(final.get("agent_weights"))
|
||||
pred_status = "success"
|
||||
model_name = aggregator_provider.model
|
||||
else:
|
||||
# 无有效报告:跳过严格校验,直接构造降级结果
|
||||
validated = None # type: ignore
|
||||
agent_weights = {}
|
||||
pred_status = "degraded"
|
||||
model_name = aggregator_model
|
||||
|
||||
pred = await _upsert_prediction(
|
||||
session,
|
||||
match_id=match_id,
|
||||
provider_name=settings.LLM_PROVIDER,
|
||||
model=aggregator_provider.model,
|
||||
model=model_name,
|
||||
mode="multi",
|
||||
run_type="backtest" if backtest else "live",
|
||||
values={
|
||||
|
||||
@@ -244,3 +244,57 @@ class TestPartialExpertsOk:
|
||||
assert captured_status.get("status") == "success", \
|
||||
f"期望 status=success,实际 {captured_status.get('status')}"
|
||||
print(f"PASS: 1 ok + 4 error → status={captured_status.get('status')}")
|
||||
|
||||
|
||||
class TestNoAggregatorCallOnDegraded:
|
||||
"""全专家失败时,aggregator provider 不应被调用。"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_all_error_no_aggregator_call(self):
|
||||
"""5 个专家全 error → 不应调用 _agent_provider('aggregator')。"""
|
||||
header = _make_header()
|
||||
aggregator_called = []
|
||||
|
||||
async def mock_run_specialists(header, *, version, before):
|
||||
return _all_error_reports()
|
||||
|
||||
async def mock_agent_provider(agent_id, *, tier):
|
||||
aggregator_called.append((agent_id, tier))
|
||||
return MagicMock(model="test-model")
|
||||
|
||||
async def mock_load_header(mid, db=None):
|
||||
return header
|
||||
|
||||
captured_values = {}
|
||||
|
||||
async def mock_upsert(session, **kw):
|
||||
captured_values.update(kw.get("values", {}))
|
||||
mock_pred = MagicMock()
|
||||
mock_pred.id = 1
|
||||
mock_pred.provider = "test"
|
||||
mock_pred.model = kw["values"].get("model")
|
||||
return mock_pred
|
||||
|
||||
class FakeUow:
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
async def __aexit__(self, *a):
|
||||
pass
|
||||
async def get(self, cls, id):
|
||||
return MagicMock()
|
||||
|
||||
with patch.object(orch_mod, "run_specialists", mock_run_specialists), \
|
||||
patch.object(orch_mod, "_agent_provider", mock_agent_provider), \
|
||||
patch.object(orch_mod, "load_match_header", mock_load_header), \
|
||||
patch.object(orch_mod, "_upsert_prediction", mock_upsert), \
|
||||
patch.object(orch_mod, "get_uow", FakeUow):
|
||||
|
||||
await orch_mod.predict_match_multi(999)
|
||||
|
||||
# 断言:aggregator provider 未被调用
|
||||
assert len(aggregator_called) == 0, \
|
||||
f"全失败时不应调用 aggregator provider,实际调用: {aggregator_called}"
|
||||
# 断言:model 使用 settings 默认值
|
||||
assert captured_values.get("model") is not None
|
||||
assert captured_values.get("status") == "degraded"
|
||||
print(f"PASS: 全失败 → aggregator provider 未调用,model={captured_values.get('model')}")
|
||||
|
||||
Reference in New Issue
Block a user