diff --git a/tests/test_litellm_models.py b/tests/test_litellm_models.py new file mode 100644 index 00000000..58fa8504 --- /dev/null +++ b/tests/test_litellm_models.py @@ -0,0 +1,65 @@ +import importlib.util +from pathlib import Path + +import litellm +import pytest + + +MODULE_PATH = ( + Path(__file__).parents[1] / "wdoc" / "utils" / "customs" / "litellm_models.py" +) +SPEC = importlib.util.spec_from_file_location("wdoc_litellm_models", MODULE_PATH) +assert SPEC is not None and SPEC.loader is not None +litellm_models = importlib.util.module_from_spec(SPEC) +SPEC.loader.exec_module(litellm_models) + + +def test_registers_minimax_models_with_current_metadata(): + litellm_models.register_wdoc_models(litellm) + + expected = { + "minimax/MiniMax-M3": { + "max_tokens": 1_000_000, + "input_cost_per_token": 0.6 / 1_000_000, + "output_cost_per_token": 2.4 / 1_000_000, + "cache_read_input_token_cost": 0.12 / 1_000_000, + "cache_creation_input_token_cost": None, + "input_modalities": ["text", "image", "video"], + "thinking": ["adaptive", "disabled"], + }, + "minimax/MiniMax-M2.7": { + "max_tokens": 204_800, + "input_cost_per_token": 0.3 / 1_000_000, + "output_cost_per_token": 1.2 / 1_000_000, + "cache_read_input_token_cost": 0.06 / 1_000_000, + "cache_creation_input_token_cost": 0.375 / 1_000_000, + "input_modalities": ["text"], + "thinking": ["always_on"], + }, + } + + for model_id, metadata in expected.items(): + assert model_id in litellm.models_by_provider["minimax"] + registered = litellm.model_cost[model_id] + assert registered["litellm_provider"] == "minimax" + assert registered["mode"] == "chat" + for key, value in metadata.items(): + assert registered[key] == value + + +@pytest.mark.parametrize( + ("region", "protocol", "expected"), + [ + ("global_en", "openai", "https://api.minimax.io/v1"), + ("global_en", "anthropic", "https://api.minimax.io/anthropic"), + ("cn_zh", "openai", "https://api.minimaxi.com/v1"), + ("cn_zh", "anthropic", "https://api.minimaxi.com/anthropic"), + ], +) +def test_minimax_endpoint_recipes(region, protocol, expected): + assert litellm_models.get_minimax_api_base(region, protocol) == expected + + +def test_minimax_endpoint_recipe_rejects_unknown_selection(): + with pytest.raises(ValueError, match="Unsupported MiniMax endpoint selection"): + litellm_models.get_minimax_api_base("unknown", "openai") diff --git a/wdoc/utils/customs/litellm_models.py b/wdoc/utils/customs/litellm_models.py new file mode 100644 index 00000000..008f3cfe --- /dev/null +++ b/wdoc/utils/customs/litellm_models.py @@ -0,0 +1,68 @@ +"""wdoc-specific model metadata registered with LiteLLM.""" + +from typing import Any, Literal + + +MINIMAX_PROVIDER = "minimax" + +MINIMAX_ENDPOINTS = { + "global_en": { + "openai_base_url": "https://api.minimax.io/v1", + "anthropic_base_url": "https://api.minimax.io/anthropic", + }, + "cn_zh": { + "openai_base_url": "https://api.minimaxi.com/v1", + "anthropic_base_url": "https://api.minimaxi.com/anthropic", + }, +} + +MINIMAX_MODELS = { + "minimax/MiniMax-M3": { + "litellm_provider": MINIMAX_PROVIDER, + "mode": "chat", + "max_tokens": 1_000_000, + "max_input_tokens": 1_000_000, + "input_cost_per_token": 0.6 / 1_000_000, + "output_cost_per_token": 2.4 / 1_000_000, + "cache_read_input_token_cost": 0.12 / 1_000_000, + "cache_creation_input_token_cost": None, + "input_modalities": ["text", "image", "video"], + "thinking": ["adaptive", "disabled"], + "supports_vision": True, + "supports_reasoning": True, + "supports_adaptive_thinking": True, + }, + "minimax/MiniMax-M2.7": { + "litellm_provider": MINIMAX_PROVIDER, + "mode": "chat", + "max_tokens": 204_800, + "max_input_tokens": 204_800, + "input_cost_per_token": 0.3 / 1_000_000, + "output_cost_per_token": 1.2 / 1_000_000, + "cache_read_input_token_cost": 0.06 / 1_000_000, + "cache_creation_input_token_cost": 0.375 / 1_000_000, + "input_modalities": ["text"], + "thinking": ["always_on"], + "supports_reasoning": True, + }, +} + + +def get_minimax_api_base( + region: Literal["global_en", "cn_zh"], + protocol: Literal["openai", "anthropic"] = "openai", +) -> str: + """Return the configured MiniMax base URL for a region and protocol.""" + try: + return MINIMAX_ENDPOINTS[region][f"{protocol}_base_url"] + except KeyError as err: + raise ValueError( + f"Unsupported MiniMax endpoint selection: {region}/{protocol}" + ) from err + + +def register_wdoc_models(litellm: Any) -> None: + """Register wdoc's model recipes in LiteLLM's cost and provider catalogs.""" + litellm.register_model(MINIMAX_MODELS) + provider_models = litellm.models_by_provider.setdefault(MINIMAX_PROVIDER, set()) + provider_models.update(MINIMAX_MODELS) diff --git a/wdoc/wdoc.py b/wdoc/wdoc.py index 75bd16e2..053b1f2c 100644 --- a/wdoc/wdoc.py +++ b/wdoc/wdoc.py @@ -32,6 +32,7 @@ ) from wdoc.utils.batch_file_loader import batch_load_doc +from wdoc.utils.customs.litellm_models import register_wdoc_models from wdoc.utils.env import env, is_out_piped from wdoc.utils.errors import ( NoDocumentsAfterLLMEvalFiltering, @@ -120,6 +121,8 @@ def __init__( """ import litellm + register_wdoc_models(litellm) + if version: print(self.VERSION) return