diff --git a/litellm/__init__.py b/litellm/__init__.py index bc8a13ec2cd..85417ef70be 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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 * diff --git a/litellm/fusion/__init__.py b/litellm/fusion/__init__.py new file mode 100644 index 00000000000..42e4ea3d0bb --- /dev/null +++ b/litellm/fusion/__init__.py @@ -0,0 +1,3 @@ +from .main import FusionStrategy, afusion, fusion + +__all__ = ["fusion", "afusion", "FusionStrategy"] diff --git a/litellm/fusion/main.py b/litellm/fusion/main.py new file mode 100644 index 00000000000..fb75fc57b41 --- /dev/null +++ b/litellm/fusion/main.py @@ -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"] diff --git a/tests/test_fusion.py b/tests/test_fusion.py new file mode 100644 index 00000000000..67d4737a5ec --- /dev/null +++ b/tests/test_fusion.py @@ -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"]