mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
6f36bee6ba
commit
3045762c73
4 changed files with 800 additions and 0 deletions
|
|
@ -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 *
|
||||
|
|
|
|||
3
litellm/fusion/__init__.py
Normal file
3
litellm/fusion/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from .main import FusionStrategy, afusion, fusion
|
||||
|
||||
__all__ = ["fusion", "afusion", "FusionStrategy"]
|
||||
430
litellm/fusion/main.py
Normal file
430
litellm/fusion/main.py
Normal 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
366
tests/test_fusion.py
Normal 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"]
|
||||
Loading…
Add table
Reference in a new issue