feat(fusion): add litellm.fusion() and litellm.afusion()

Adds multi-model Fusion support to the litellm SDK, resolving #30456.

## What

-  — sync
-  — async

## Behavior

- **Panel only** ( omitted): fan out to N models in
  parallel via asyncio.gather, return
- **With judge** ( set): synthesize a final answer from
  all panel responses, return single
  -  (default): panel responses attached to

  - : panel responses omitted from output

## Strategies (when judge_model is set)

-  (default): judge synthesizes a new combined answer
- : judge selects the single best response verbatim
- : judge scores each response, returns highest-scored

## Other details

- Partial panel failure: failed models are skipped; fusion continues
  with remaining responses
- All panel failure: raises RuntimeError
- Usage: token counts from all panel + judge calls are summed
- @overload signatures for precise return-type inference (list vs single)

Closes #30456
This commit is contained in:
leecoder 2026-08-11 14:14:38 +09:00
parent 6f36bee6ba
commit 3045762c73
4 changed files with 800 additions and 0 deletions

View file

@ -1356,6 +1356,7 @@ from .batches.main import *
from .images.main import *
from .videos.main import *
from .batch_completion.main import *
from .fusion.main import fusion, afusion, FusionStrategy
from .rerank_api.main import *
from .llms.anthropic.experimental_pass_through.messages.handler import *
from .responses.main import *

View file

@ -0,0 +1,3 @@
from .main import FusionStrategy, afusion, fusion
__all__ = ["fusion", "afusion", "FusionStrategy"]

430
litellm/fusion/main.py Normal file
View file

@ -0,0 +1,430 @@
"""
litellm.fusion — Multi-model Fusion with optional judge synthesis.
Sends the same prompt to N models in parallel. Without a judge model, returns
all panel responses as a list. With a judge model, synthesizes a final answer
and includes panel responses by default (set include_panel=False to suppress).
Inspired by OpenRouter Fusion / Sakana Fugu / Self-Consistency (Wang et al., 2022).
Usage:
# Panel only — get all responses, pick yourself
responses = litellm.fusion(
models=["gpt-4o", "claude-3-5-sonnet", "gemini-2.0-flash"],
messages=[{"role": "user", "content": "Explain quantum entanglement"}],
)
# returns: list[ModelResponse]
# Judge synthesis — synthesized answer + panel included by default
response = litellm.fusion(
models=["gpt-4o", "claude-3-5-sonnet", "gemini-2.0-flash"],
judge_model="gpt-4o",
messages=[{"role": "user", "content": "Explain quantum entanglement"}],
)
# returns: ModelResponse
# response._hidden_params["fusion"]["panel_responses"] ← individual responses
# Judge synthesis — panel excluded
response = litellm.fusion(
models=["gpt-4o", "claude-3-5-sonnet", "gemini-2.0-flash"],
judge_model="gpt-4o",
messages=[{"role": "user", "content": "Explain quantum entanglement"}],
include_panel=False,
)
# returns: ModelResponse (no panel_responses in metadata)
# Async variants
responses = await litellm.afusion(models=[...], messages=[...])
response = await litellm.afusion(models=[...], judge_model="...", messages=[...])
"""
from __future__ import annotations
import asyncio
import time
import uuid
from typing import List, Literal, Optional, Union, overload
import litellm
from litellm.types.utils import Choices, Message, ModelResponse, Usage
# ---------------------------------------------------------------------------
# Judge prompt templates
# ---------------------------------------------------------------------------
_JUDGE_PROMPT_SINGLE = """\
You received the following responses from {n} AI models for this user request:
{responses}
Your task: synthesize a single, comprehensive, and accurate final answer.
- Identify the points of agreement and the best reasoning across responses.
- Resolve contradictions by selecting the most well-supported position.
- Fill in gaps where some models provided information others missed.
- Write the final answer as if you generated it directly — no meta-commentary about the synthesis process.
"""
_JUDGE_PROMPT_MAJORITY = """\
You received the following responses from {n} AI models for this user request:
{responses}
Select the single best response. Output ONLY the text of the chosen response, verbatim.
"""
_JUDGE_PROMPT_BEST_OF_N = """\
You received the following responses from {n} AI models for this user request:
{responses}
Score each response on a scale of 1-10 for accuracy, completeness, and clarity.
Then output the highest-scored response verbatim.
"""
_JUDGE_PROMPTS: dict[str, str] = {
"single_judge": _JUDGE_PROMPT_SINGLE,
"majority_vote": _JUDGE_PROMPT_MAJORITY,
"best_of_n": _JUDGE_PROMPT_BEST_OF_N,
}
FusionStrategy = Literal["single_judge", "majority_vote", "best_of_n"]
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _build_judge_messages(
original_messages: list[dict],
panel_responses: list[ModelResponse],
panel_models: list[str],
strategy: FusionStrategy,
) -> list[dict]:
"""Build the message list for the judge call."""
user_request = next(
(m["content"] for m in reversed(original_messages) if m.get("role") == "user"),
"",
)
response_blocks = []
for i, (model, resp) in enumerate(zip(panel_models, panel_responses), start=1):
content = ""
if resp and resp.choices:
content = resp.choices[0].message.content or ""
response_blocks.append(f"[Response {i} — {model}]\n{content}")
responses_text = "\n\n".join(response_blocks)
template = _JUDGE_PROMPTS[strategy]
judge_system = template.format(n=len(panel_responses), responses=responses_text)
messages: list[dict] = []
# Preserve any existing system prompt
for m in original_messages:
if m.get("role") == "system":
messages.append(m)
break
messages.append({"role": "system", "content": judge_system})
messages.append({"role": "user", "content": user_request})
return messages
def _sum_usage(responses: list[ModelResponse]) -> Usage:
"""Sum token usage across a list of ModelResponse objects."""
total_prompt = 0
total_completion = 0
total_tokens = 0
for r in responses:
if r and r.usage:
total_prompt += r.usage.prompt_tokens or 0
total_completion += r.usage.completion_tokens or 0
total_tokens += r.usage.total_tokens or 0
return Usage(
prompt_tokens=total_prompt,
completion_tokens=total_completion,
total_tokens=total_tokens,
)
def _merge_usage(
panel_responses: list[ModelResponse],
judge_response: ModelResponse,
) -> Usage:
"""Sum token usage across all panel calls + judge call."""
return _sum_usage(panel_responses + [judge_response])
def _build_fusion_response(
judge_response: ModelResponse,
panel_responses: list[ModelResponse],
panel_models: list[str],
judge_model: str,
include_panel: bool,
original_model_tag: str,
) -> ModelResponse:
"""Wrap judge response in a ModelResponse tagged with fusion metadata."""
merged_usage = _merge_usage(panel_responses, judge_response)
response = ModelResponse(
id=f"fusion-{uuid.uuid4().hex}",
choices=judge_response.choices,
created=int(time.time()),
model=original_model_tag,
usage=merged_usage,
object="chat.completion",
)
fusion_meta: dict = {
"panel_models": panel_models,
"judge_model": judge_model,
}
if include_panel:
fusion_meta["panel_responses"] = panel_responses
response._hidden_params = {"fusion": fusion_meta}
return response
async def _run_panel(
models: list[str],
messages: list,
panel_kwargs: dict,
) -> tuple[list[ModelResponse], list[str]]:
"""Fan out to all panel models in parallel; filter failed calls."""
tasks = [
litellm.acompletion(model=m, messages=messages, stream=False, **panel_kwargs)
for m in models
]
raw: list = list(await asyncio.gather(*tasks, return_exceptions=True))
valid_responses: list[ModelResponse] = []
valid_models: list[str] = []
for model, resp in zip(models, raw):
if isinstance(resp, Exception):
litellm.utils.print_verbose(f"fusion: panel model {model!r} failed: {resp}")
else:
valid_responses.append(resp)
valid_models.append(model)
if not valid_responses:
raise RuntimeError(f"fusion: all panel models failed. Errors: {raw}")
return valid_responses, valid_models
# ---------------------------------------------------------------------------
# Overloads for precise return-type inference
# ---------------------------------------------------------------------------
@overload
async def afusion(
models: List[str],
messages: list,
*,
judge_model: None = ...,
strategy: FusionStrategy = ...,
include_panel: bool = ...,
timeout: Optional[float] = ...,
temperature: Optional[float] = ...,
max_tokens: Optional[int] = ...,
**kwargs,
) -> list[ModelResponse]: ...
@overload
async def afusion(
models: List[str],
messages: list,
*,
judge_model: str,
strategy: FusionStrategy = ...,
include_panel: bool = ...,
timeout: Optional[float] = ...,
temperature: Optional[float] = ...,
max_tokens: Optional[int] = ...,
**kwargs,
) -> ModelResponse: ...
# ---------------------------------------------------------------------------
# Core async implementation
# ---------------------------------------------------------------------------
async def afusion(
models: List[str],
messages: list,
*,
judge_model: Optional[str] = None,
strategy: FusionStrategy = "single_judge",
include_panel: bool = True,
timeout: Optional[float] = None,
temperature: Optional[float] = None,
max_tokens: Optional[int] = None,
**kwargs,
) -> Union[list[ModelResponse], ModelResponse]:
"""
Async fusion: call panel models in parallel, optionally synthesize with judge.
Args:
models: List of panel model identifiers (≥ 2 recommended).
messages: OpenAI-compatible message list.
judge_model: If provided, synthesizes a final answer from panel responses.
If omitted, returns the raw list of panel responses.
strategy: How the judge processes responses (only used when judge_model is set).
- "single_judge": synthesize a new combined answer (default)
- "majority_vote": pick the single best response verbatim
- "best_of_n": score and select highest-scored response
include_panel: When judge_model is set, whether to include individual panel
responses in response._hidden_params["fusion"]["panel_responses"].
Defaults to True. Ignored when judge_model is None.
timeout: Per-call timeout in seconds (applied to panel + judge calls).
temperature: Forwarded to panel models. Judge always uses temperature=0.
max_tokens: Forwarded to panel models.
**kwargs: Any additional litellm.acompletion kwargs forwarded to panel calls.
Returns:
- list[ModelResponse] when judge_model is None
- ModelResponse (synthesized) when judge_model is provided
"""
if not models:
raise ValueError("fusion: `models` must be a non-empty list")
if strategy not in _JUDGE_PROMPTS:
raise ValueError(
f"fusion: unknown strategy {strategy!r}. Choose from {list(_JUDGE_PROMPTS)}"
)
panel_kwargs: dict = dict(kwargs)
if temperature is not None:
panel_kwargs["temperature"] = temperature
if max_tokens is not None:
panel_kwargs["max_tokens"] = max_tokens
if timeout is not None:
panel_kwargs["timeout"] = timeout
# 1. Fan out to all panel models in parallel
valid_responses, valid_models = await _run_panel(models, messages, panel_kwargs)
# 2. No judge — return panel responses as-is
if judge_model is None:
return valid_responses
# 3. Judge synthesis
judge_messages = _build_judge_messages(
original_messages=messages,
panel_responses=valid_responses,
panel_models=valid_models,
strategy=strategy,
)
judge_kwargs: dict = {}
if timeout is not None:
judge_kwargs["timeout"] = timeout
judge_response: ModelResponse = await litellm.acompletion(
model=judge_model,
messages=judge_messages,
stream=False,
temperature=0, # deterministic synthesis
**judge_kwargs,
)
# 4. Wrap into a single ModelResponse with merged metadata
return _build_fusion_response(
judge_response=judge_response,
panel_responses=valid_responses,
panel_models=valid_models,
judge_model=judge_model,
include_panel=include_panel,
original_model_tag=f"fusion/{'+'.join(valid_models)}",
)
# ---------------------------------------------------------------------------
# Sync overloads
# ---------------------------------------------------------------------------
@overload
def fusion(
models: List[str],
messages: list,
*,
judge_model: None = ...,
strategy: FusionStrategy = ...,
include_panel: bool = ...,
timeout: Optional[float] = ...,
temperature: Optional[float] = ...,
max_tokens: Optional[int] = ...,
**kwargs,
) -> list[ModelResponse]: ...
@overload
def fusion(
models: List[str],
messages: list,
*,
judge_model: str,
strategy: FusionStrategy = ...,
include_panel: bool = ...,
timeout: Optional[float] = ...,
temperature: Optional[float] = ...,
max_tokens: Optional[int] = ...,
**kwargs,
) -> ModelResponse: ...
def fusion(
models: List[str],
messages: list,
*,
judge_model: Optional[str] = None,
strategy: FusionStrategy = "single_judge",
include_panel: bool = True,
timeout: Optional[float] = None,
temperature: Optional[float] = None,
max_tokens: Optional[int] = None,
**kwargs,
) -> Union[list[ModelResponse], ModelResponse]:
"""
Sync fusion: call panel models in parallel, optionally synthesize with judge.
Thin sync wrapper around :func:`afusion`. For async contexts, prefer
:func:`afusion` directly.
Args:
models: List of panel model identifiers (≥ 2 recommended).
messages: OpenAI-compatible message list.
judge_model: If provided, synthesizes a final answer from panel responses.
If omitted, returns the raw list of panel responses.
strategy: Synthesis strategy ("single_judge", "majority_vote", "best_of_n").
Only used when judge_model is provided.
include_panel: When judge_model is set, whether to attach individual panel
responses to response._hidden_params["fusion"]["panel_responses"].
Defaults to True. Ignored when judge_model is None.
timeout: Per-call timeout in seconds.
temperature: Panel model temperature (judge always uses 0).
max_tokens: Forwarded to panel models.
**kwargs: Any additional litellm.completion kwargs.
Returns:
- list[ModelResponse] when judge_model is None
- ModelResponse (synthesized) when judge_model is provided
"""
return asyncio.run(
afusion(
models=models,
messages=messages,
judge_model=judge_model,
strategy=strategy,
include_panel=include_panel,
timeout=timeout,
temperature=temperature,
max_tokens=max_tokens,
**kwargs,
)
)
__all__ = ["fusion", "afusion", "FusionStrategy"]

366
tests/test_fusion.py Normal file
View file

@ -0,0 +1,366 @@
"""
Tests for litellm.fusion and litellm.afusion.
Uses unittest.mock to avoid real API calls.
"""
from __future__ import annotations
from unittest.mock import patch
import pytest
import litellm
from litellm.fusion.main import (
FusionStrategy,
_build_judge_messages,
_merge_usage,
_sum_usage,
afusion,
fusion,
)
from litellm.types.utils import Choices, Message, ModelResponse, Usage
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_response(content: str, model: str = "gpt-4o", prompt_tokens: int = 10) -> ModelResponse:
resp = ModelResponse(
id=f"mock-{model}",
choices=[
Choices(
index=0,
message=Message(role="assistant", content=content),
finish_reason="stop",
)
],
model=model,
usage=Usage(
prompt_tokens=prompt_tokens,
completion_tokens=len(content.split()),
total_tokens=prompt_tokens + len(content.split()),
),
)
return resp
MESSAGES = [{"role": "user", "content": "What is 2+2?"}]
PANEL_MODELS = ["gpt-4o", "claude-3-5-sonnet", "gemini-2.0-flash"]
JUDGE_MODEL = "gpt-4o-judge" # distinct from panel to avoid index confusion
def _panel_mock(panel_responses, judge_response=None):
"""Return an async mock that routes by model name."""
async def mock_acompletion(**kwargs):
m = kwargs["model"]
if judge_response is not None and m == JUDGE_MODEL:
return judge_response
if m in PANEL_MODELS:
return panel_responses[PANEL_MODELS.index(m)]
return panel_responses[0]
return mock_acompletion
# ---------------------------------------------------------------------------
# Unit: _build_judge_messages
# ---------------------------------------------------------------------------
class TestBuildJudgeMessages:
def _panel(self):
return [
_make_response("The answer is 4.", "gpt-4o"),
_make_response("4", "claude-3-5-sonnet"),
]
def test_contains_user_question(self):
msgs = _build_judge_messages(MESSAGES, self._panel(), PANEL_MODELS[:2], "single_judge")
combined = " ".join(m["content"] for m in msgs)
assert "2+2" in combined
def test_contains_panel_responses(self):
msgs = _build_judge_messages(MESSAGES, self._panel(), PANEL_MODELS[:2], "single_judge")
combined = " ".join(m["content"] for m in msgs)
assert "The answer is 4." in combined
assert "4" in combined
def test_strategy_majority_vote(self):
msgs = _build_judge_messages(MESSAGES, self._panel(), PANEL_MODELS[:2], "majority_vote")
system_msg = next(m for m in msgs if m["role"] == "system")
assert "Select the single best response" in system_msg["content"]
def test_strategy_best_of_n(self):
msgs = _build_judge_messages(MESSAGES, self._panel(), PANEL_MODELS[:2], "best_of_n")
system_msg = next(m for m in msgs if m["role"] == "system")
assert "Score each response" in system_msg["content"]
def test_ends_with_user_message(self):
msgs = _build_judge_messages(MESSAGES, self._panel(), PANEL_MODELS[:2], "single_judge")
assert msgs[-1]["role"] == "user"
# ---------------------------------------------------------------------------
# Unit: _sum_usage / _merge_usage
# ---------------------------------------------------------------------------
class TestUsageHelpers:
def test_sum_usage(self):
responses = [
_make_response("hello", prompt_tokens=10), # completion=1
_make_response("world more words", prompt_tokens=20), # completion=3
]
u = _sum_usage(responses)
assert u.prompt_tokens == 30
assert u.completion_tokens == 4
def test_merge_usage_sums_panel_and_judge(self):
panel = [
_make_response("hello", prompt_tokens=10),
_make_response("world more words", prompt_tokens=20),
]
judge = _make_response("final answer here today", prompt_tokens=5)
merged = _merge_usage(panel, judge)
assert merged.prompt_tokens == 35 # 10+20+5
assert merged.completion_tokens == 1 + 3 + 4
def test_handles_none_usage(self):
panel = [_make_response("hi")]
panel[0].usage = None
judge = _make_response("answer", prompt_tokens=5)
merged = _merge_usage(panel, judge)
assert merged.prompt_tokens == 5
# ---------------------------------------------------------------------------
# afusion — panel only (no judge)
# ---------------------------------------------------------------------------
class TestAFusionPanelOnly:
@pytest.fixture
def panel_responses(self):
return [_make_response(f"Answer {m}", m) for m in PANEL_MODELS]
@pytest.mark.asyncio
async def test_returns_list_of_model_responses(self, panel_responses):
with patch("litellm.acompletion", side_effect=_panel_mock(panel_responses)):
result = await afusion(models=PANEL_MODELS, messages=MESSAGES)
assert isinstance(result, list)
assert len(result) == len(PANEL_MODELS)
assert all(isinstance(r, ModelResponse) for r in result)
@pytest.mark.asyncio
async def test_returns_all_panel_contents(self, panel_responses):
with patch("litellm.acompletion", side_effect=_panel_mock(panel_responses)):
result = await afusion(models=PANEL_MODELS, messages=MESSAGES)
contents = {r.choices[0].message.content for r in result}
assert contents == {f"Answer {m}" for m in PANEL_MODELS}
@pytest.mark.asyncio
async def test_no_judge_call_when_judge_model_omitted(self, panel_responses):
called_models: list[str] = []
async def mock_acompletion(**kwargs):
called_models.append(kwargs["model"])
return panel_responses[PANEL_MODELS.index(kwargs["model"])]
with patch("litellm.acompletion", side_effect=mock_acompletion):
await afusion(models=PANEL_MODELS, messages=MESSAGES)
# Exactly the panel models, no judge
assert sorted(called_models) == sorted(PANEL_MODELS)
@pytest.mark.asyncio
async def test_partial_failure_returns_remaining(self, panel_responses):
async def mock_acompletion(**kwargs):
if kwargs["model"] == "gemini-2.0-flash":
raise RuntimeError("API error")
return panel_responses[PANEL_MODELS.index(kwargs["model"])]
with patch("litellm.acompletion", side_effect=mock_acompletion):
result = await afusion(models=PANEL_MODELS, messages=MESSAGES)
assert isinstance(result, list)
assert len(result) == 2 # gemini dropped
# ---------------------------------------------------------------------------
# afusion — with judge
# ---------------------------------------------------------------------------
class TestAFusionWithJudge:
@pytest.fixture
def panel_responses(self):
return [_make_response(f"Answer {m}", m) for m in PANEL_MODELS]
@pytest.fixture
def judge_response(self):
return _make_response("Synthesized final answer", JUDGE_MODEL)
@pytest.mark.asyncio
async def test_returns_single_model_response(self, panel_responses, judge_response):
with patch("litellm.acompletion", side_effect=_panel_mock(panel_responses, judge_response)):
result = await afusion(models=PANEL_MODELS, judge_model=JUDGE_MODEL, messages=MESSAGES)
assert isinstance(result, ModelResponse)
assert result.choices[0].message.content == "Synthesized final answer"
@pytest.mark.asyncio
async def test_calls_all_panel_and_judge(self, panel_responses, judge_response):
called: list[str] = []
async def mock_acompletion(**kwargs):
called.append(kwargs["model"])
return judge_response if kwargs["model"] == JUDGE_MODEL else panel_responses[PANEL_MODELS.index(kwargs["model"])]
with patch("litellm.acompletion", side_effect=mock_acompletion):
await afusion(models=PANEL_MODELS, judge_model=JUDGE_MODEL, messages=MESSAGES)
assert set(PANEL_MODELS).issubset(set(called))
assert JUDGE_MODEL in called
@pytest.mark.asyncio
async def test_include_panel_true_by_default(self, panel_responses, judge_response):
with patch("litellm.acompletion", side_effect=_panel_mock(panel_responses, judge_response)):
result = await afusion(models=PANEL_MODELS, judge_model=JUDGE_MODEL, messages=MESSAGES)
fusion_meta = result._hidden_params["fusion"]
assert "panel_responses" in fusion_meta
assert len(fusion_meta["panel_responses"]) == len(PANEL_MODELS)
assert all(isinstance(r, ModelResponse) for r in fusion_meta["panel_responses"])
@pytest.mark.asyncio
async def test_include_panel_false_excludes_panel(self, panel_responses, judge_response):
with patch("litellm.acompletion", side_effect=_panel_mock(panel_responses, judge_response)):
result = await afusion(
models=PANEL_MODELS,
judge_model=JUDGE_MODEL,
messages=MESSAGES,
include_panel=False,
)
fusion_meta = result._hidden_params["fusion"]
assert "panel_responses" not in fusion_meta
@pytest.mark.asyncio
async def test_fusion_metadata_always_present(self, panel_responses, judge_response):
with patch("litellm.acompletion", side_effect=_panel_mock(panel_responses, judge_response)):
result = await afusion(models=PANEL_MODELS, judge_model=JUDGE_MODEL, messages=MESSAGES)
meta = result._hidden_params["fusion"]
assert meta["judge_model"] == JUDGE_MODEL
assert set(meta["panel_models"]) == set(PANEL_MODELS)
@pytest.mark.asyncio
async def test_usage_merged(self, panel_responses, judge_response):
with patch("litellm.acompletion", side_effect=_panel_mock(panel_responses, judge_response)):
result = await afusion(models=PANEL_MODELS, judge_model=JUDGE_MODEL, messages=MESSAGES)
assert result.usage is not None
assert result.usage.total_tokens > 0
@pytest.mark.asyncio
async def test_partial_panel_failure_continues(self, panel_responses, judge_response):
async def mock_acompletion(**kwargs):
if kwargs["model"] == "gemini-2.0-flash":
raise RuntimeError("API error")
return judge_response if kwargs["model"] == JUDGE_MODEL else panel_responses[PANEL_MODELS.index(kwargs["model"])]
with patch("litellm.acompletion", side_effect=mock_acompletion):
result = await afusion(models=PANEL_MODELS, judge_model=JUDGE_MODEL, messages=MESSAGES)
assert isinstance(result, ModelResponse)
# Only 2 panel models survived
assert len(result._hidden_params["fusion"]["panel_models"]) == 2
@pytest.mark.asyncio
async def test_all_panel_fail_raises(self):
async def mock_acompletion(**kwargs):
raise RuntimeError("All failed")
with patch("litellm.acompletion", side_effect=mock_acompletion):
with pytest.raises(RuntimeError, match="all panel models failed"):
await afusion(models=PANEL_MODELS, judge_model=JUDGE_MODEL, messages=MESSAGES)
@pytest.mark.asyncio
@pytest.mark.parametrize("strategy", ["single_judge", "majority_vote", "best_of_n"])
async def test_all_strategies(self, strategy, panel_responses, judge_response):
with patch("litellm.acompletion", side_effect=_panel_mock(panel_responses, judge_response)):
result = await afusion(
models=PANEL_MODELS,
judge_model=JUDGE_MODEL,
messages=MESSAGES,
strategy=strategy,
)
assert isinstance(result, ModelResponse)
@pytest.mark.asyncio
async def test_invalid_strategy_raises(self, panel_responses, judge_response):
with patch("litellm.acompletion", side_effect=_panel_mock(panel_responses, judge_response)):
with pytest.raises(ValueError, match="unknown strategy"):
await afusion(
models=PANEL_MODELS,
judge_model=JUDGE_MODEL,
messages=MESSAGES,
strategy="nonexistent", # type: ignore
)
# ---------------------------------------------------------------------------
# Common error cases
# ---------------------------------------------------------------------------
class TestCommonErrors:
@pytest.mark.asyncio
async def test_empty_models_raises(self):
with pytest.raises(ValueError, match="non-empty list"):
await afusion(models=[], messages=MESSAGES)
@pytest.mark.asyncio
async def test_empty_models_with_judge_raises(self):
with pytest.raises(ValueError, match="non-empty list"):
await afusion(models=[], judge_model=JUDGE_MODEL, messages=MESSAGES)
# ---------------------------------------------------------------------------
# Sync wrapper
# ---------------------------------------------------------------------------
class TestFusionSync:
def test_callable_via_litellm(self):
assert callable(litellm.fusion)
assert callable(litellm.afusion)
def test_panel_only_returns_list(self):
panel_resps = [_make_response(f"Panel {m}", m) for m in PANEL_MODELS]
with patch("litellm.acompletion", side_effect=_panel_mock(panel_resps)):
result = fusion(models=PANEL_MODELS, messages=MESSAGES)
assert isinstance(result, list)
assert len(result) == len(PANEL_MODELS)
def test_with_judge_returns_model_response(self):
panel_resps = [_make_response(f"Panel {m}", m) for m in PANEL_MODELS]
judge_resp = _make_response("Sync result", JUDGE_MODEL)
with patch("litellm.acompletion", side_effect=_panel_mock(panel_resps, judge_resp)):
result = fusion(models=PANEL_MODELS, judge_model=JUDGE_MODEL, messages=MESSAGES)
assert isinstance(result, ModelResponse)
assert result.choices[0].message.content == "Sync result"
def test_include_panel_false(self):
panel_resps = [_make_response(f"Panel {m}", m) for m in PANEL_MODELS]
judge_resp = _make_response("Sync result", JUDGE_MODEL)
with patch("litellm.acompletion", side_effect=_panel_mock(panel_resps, judge_resp)):
result = fusion(
models=PANEL_MODELS,
judge_model=JUDGE_MODEL,
messages=MESSAGES,
include_panel=False,
)
assert "panel_responses" not in result._hidden_params["fusion"]