diff --git a/mise.toml b/mise.toml new file mode 100644 index 00000000..55953276 --- /dev/null +++ b/mise.toml @@ -0,0 +1,4 @@ +[tools] +black = "latest" +python = "3.13" +uv = "latest" diff --git a/src/modelgauge/suts/openrouter_sut_factory.py b/src/modelgauge/suts/openrouter_sut_factory.py index 829d7db8..6414e34d 100644 --- a/src/modelgauge/suts/openrouter_sut_factory.py +++ b/src/modelgauge/suts/openrouter_sut_factory.py @@ -1,15 +1,76 @@ -from modelgauge.secret_values import RawSecrets -from modelgauge.suts.openai_sut_factory import OPENAI_SUT_FACTORIES, OpenAIGenericSUTFactory +from openai import OpenAI + +from modelgauge.auth.openai_compatible_secrets import OpenAICompatibleApiKey +from modelgauge.dynamic_sut_factory import ( + DynamicDriverSUTFactory, + ModelNotSupportedError, +) +from modelgauge.secret_values import InjectSecret, RawSecrets +from modelgauge.sut import SUT +from modelgauge.sut_capabilities import ( + AcceptsChatPrompt, + AcceptsTextPrompt, + ProducesPerTokenLogProbabilities, +) +from modelgauge.sut_definition import SUTDefinition +from modelgauge.sut_decorator import modelgauge_sut +from modelgauge.suts.openai_client import OpenAIResponsesSUT +from modelgauge.suts.openai_sut_factory import NUM_RETRIES, BaseOpenAISUTFactory OPENROUTER_BASE_URL = "https://openrouter.ai/api/v1" -class OpenRouterSUTFactory(OpenAIGenericSUTFactory): +@modelgauge_sut( + capabilities=[ + AcceptsTextPrompt, + AcceptsChatPrompt, + ProducesPerTokenLogProbabilities, + ] +) +class OpenRouterSUT(OpenAIResponsesSUT): + """ + Documented at https://openrouter.ai/docs + """ + + +class OpenRouterSUTFactory(BaseOpenAISUTFactory, DynamicDriverSUTFactory): + DRIVER_NAME = "openrouter" - def __init__(self, raw_secrets: RawSecrets, base_url: str | None = OPENROUTER_BASE_URL): + def __init__(self, raw_secrets: RawSecrets): super().__init__(raw_secrets) self.provider = "openrouter" - self.base_url = base_url + self.base_url = OPENROUTER_BASE_URL + + def get_secrets(self) -> list[InjectSecret]: + return [InjectSecret(OpenAICompatibleApiKey.for_provider("openrouter"))] + + def _make_client(self) -> OpenAI: + [api_key] = self.injected_secrets() + return OpenAI(api_key=api_key.value, base_url=self.base_url, max_retries=NUM_RETRIES) + + def _model_exists(self, model_name: str) -> bool: + try: + data = self.client.models.list().data + except Exception: + return False + return any(entry.id == model_name for entry in data) + def make_sut(self, sut_definition: SUTDefinition) -> SUT: + model_name = sut_definition.external_model_name() + if not self._model_exists(model_name): + raise ModelNotSupportedError(f"Model {model_name} not found or not available on openrouter.") + return OpenRouterSUT(sut_definition.uid, model_name, client=self.client) -OPENAI_SUT_FACTORIES["openrouter"] = OpenRouterSUTFactory + def list_suts(self) -> list[SUTDefinition] | None: + data = self.client.models.list().data + result = [] + for entry in data: + maker, model = entry.id.split("/", 1) + result.append( + SUTDefinition( + driver=self.DRIVER_NAME, + maker=maker, + model=model, + ) + ) + return result diff --git a/tests/modelgauge_tests/sut_tests/test_openrouter_factory.py b/tests/modelgauge_tests/sut_tests/test_openrouter_factory.py index 63df6cdb..199d7aeb 100644 --- a/tests/modelgauge_tests/sut_tests/test_openrouter_factory.py +++ b/tests/modelgauge_tests/sut_tests/test_openrouter_factory.py @@ -1,23 +1,63 @@ +from unittest.mock import MagicMock + import pytest +from openai import OpenAI + +from modelgauge.dynamic_sut_factory import ModelNotSupportedError from modelgauge.sut_definition import SUTDefinition -from modelgauge.suts.openai_client import OpenAIResponsesSUT -from modelgauge.suts.openai_sut_factory import OpenAICompatibleSUTFactory -from modelgauge.suts.openrouter_sut_factory import OPENROUTER_BASE_URL +from modelgauge.suts.openrouter_sut_factory import OPENROUTER_BASE_URL, OpenRouterSUT, OpenRouterSUTFactory +from modelgauge_tests.utilities import FakeObject @pytest.fixture def factory(): - return OpenAICompatibleSUTFactory( + factory = OpenRouterSUTFactory( raw_secrets={ "openrouter": {"api_key": "some_key"}, } ) + real_client = OpenAI(api_key="some_key", base_url=OPENROUTER_BASE_URL) + list_response = MagicMock() + list_response.data = [ + FakeObject(id="qwen/qwen3-max-0902", provider="alibaba"), + FakeObject(id="meta-llama/llama-3.1-70b", provider="meta"), + ] + real_client.models.list = MagicMock(return_value=list_response) + factory._client = real_client + return factory def test_factory_makes_correct_openrouter_sut(factory): - sut_definition = SUTDefinition(model="gpt-oss-20b", maker="openai", driver="openai", provider="openrouter") + sut_definition = SUTDefinition( + maker="qwen", + model="qwen3-max-0902", + provider="alibaba", + driver="openrouter", + ) sut = factory.make_sut(sut_definition) - assert isinstance(sut, OpenAIResponsesSUT) - assert sut.uid == "openai/gpt-oss-20b:openrouter:openai" - assert sut.model == "gpt-oss-20b" - assert str(sut.client.base_url).startswith(OPENROUTER_BASE_URL) + + assert isinstance(sut, OpenRouterSUT) + assert sut.uid == "qwen/qwen3-max-0902:alibaba:openrouter" + assert sut.model == "qwen/qwen3-max-0902" + assert sut.client is factory.client + assert str(factory.client.base_url).startswith(OPENROUTER_BASE_URL) + factory.client.models.list.assert_called() + + +def test_make_sut_bad_model(factory): + sut_definition = SUTDefinition( + maker="qwen", + model="bogus", + provider="alibaba", + driver="openrouter", + ) + with pytest.raises(ModelNotSupportedError): + factory.make_sut(sut_definition) + + +def test_list_suts(factory): + suts = factory.list_suts() + assert suts is not None + uids = [s.uid for s in suts] + assert "qwen/qwen3-max-0902:openrouter" in uids + assert "meta-llama/llama-3.1-70b:openrouter" in uids