From 89b06cd74e12a73f316ba6321c830df0e25122c6 Mon Sep 17 00:00:00 2001 From: RheagalFire Date: Thu, 23 Apr 2026 23:54:40 +0530 Subject: [PATCH] feat: add LiteLLM as AI gateway engine --- gui_agents/s2/core/engine.py | 27 ++++ gui_agents/s2/core/mllm.py | 6 +- gui_agents/s2_5/core/engine.py | 27 ++++ gui_agents/s2_5/core/mllm.py | 6 +- gui_agents/s3/core/engine.py | 31 +++++ gui_agents/s3/core/mllm.py | 6 +- requirements.txt | 1 + setup.py | 1 + tests/__init__.py | 0 tests/test_litellm_engine.py | 235 +++++++++++++++++++++++++++++++++ 10 files changed, 337 insertions(+), 3 deletions(-) create mode 100644 tests/__init__.py create mode 100644 tests/test_litellm_engine.py diff --git a/gui_agents/s2/core/engine.py b/gui_agents/s2/core/engine.py index cef83fd8..20cd3e14 100644 --- a/gui_agents/s2/core/engine.py +++ b/gui_agents/s2/core/engine.py @@ -451,6 +451,33 @@ def generate(self, messages, temperature=0.0, max_new_tokens=None, **kwargs): ) +class LMMEngineLiteLLM(LMMEngine): + def __init__(self, api_key=None, model=None, rate_limit=-1, **kwargs): + assert model is not None, "model must be provided" + self.model = model + self.api_key = api_key + self.request_interval = 0 if rate_limit == -1 else 60.0 / rate_limit + + @backoff.on_exception( + backoff.expo, (APIConnectionError, APIError, RateLimitError), max_time=60 + ) + def generate(self, messages, temperature=0.0, max_new_tokens=None, **kwargs): + import litellm + + params = { + "model": self.model, + "messages": messages, + "max_tokens": max_new_tokens if max_new_tokens else 4096, + "temperature": temperature, + "drop_params": True, + **kwargs, + } + if self.api_key: + params["api_key"] = self.api_key + response = litellm.completion(**params) + return response.choices[0].message.content + + class LMMEngineParasail(LMMEngine): def __init__(self, api_key=None, model=None, rate_limit=-1, **kwargs): assert model is not None, "Parasail model id must be provided" diff --git a/gui_agents/s2/core/mllm.py b/gui_agents/s2/core/mllm.py index 6f2a516d..9bc49736 100644 --- a/gui_agents/s2/core/mllm.py +++ b/gui_agents/s2/core/mllm.py @@ -5,12 +5,13 @@ from gui_agents.s2.core.engine import ( LMMEngineAnthropic, LMMEngineAzureOpenAI, + LMMEngineGemini, LMMEngineHuggingFace, + LMMEngineLiteLLM, LMMEngineOpenAI, LMMEngineOpenRouter, LMMEngineParasail, LMMEnginevLLM, - LMMEngineGemini, ) @@ -35,6 +36,8 @@ def __init__(self, engine_params=None, system_prompt=None, engine=None): self.engine = LMMEngineOpenRouter(**engine_params) elif engine_type == "parasail": self.engine = LMMEngineParasail(**engine_params) + elif engine_type == "litellm": + self.engine = LMMEngineLiteLLM(**engine_params) else: raise ValueError("engine_type is not supported") else: @@ -127,6 +130,7 @@ def add_message( LMMEngineAzureOpenAI, LMMEngineHuggingFace, LMMEngineGemini, + LMMEngineLiteLLM, LMMEngineOpenRouter, LMMEngineParasail, ), diff --git a/gui_agents/s2_5/core/engine.py b/gui_agents/s2_5/core/engine.py index 3b17de60..66e96d7e 100644 --- a/gui_agents/s2_5/core/engine.py +++ b/gui_agents/s2_5/core/engine.py @@ -398,6 +398,33 @@ def generate(self, messages, temperature=0.0, max_new_tokens=None, **kwargs): ) +class LMMEngineLiteLLM(LMMEngine): + def __init__(self, api_key=None, model=None, rate_limit=-1, **kwargs): + assert model is not None, "model must be provided" + self.model = model + self.api_key = api_key + self.request_interval = 0 if rate_limit == -1 else 60.0 / rate_limit + + @backoff.on_exception( + backoff.expo, (APIConnectionError, APIError, RateLimitError), max_time=60 + ) + def generate(self, messages, temperature=0.0, max_new_tokens=None, **kwargs): + import litellm + + params = { + "model": self.model, + "messages": messages, + "max_tokens": max_new_tokens if max_new_tokens else 4096, + "temperature": temperature, + "drop_params": True, + **kwargs, + } + if self.api_key: + params["api_key"] = self.api_key + response = litellm.completion(**params) + return response.choices[0].message.content + + class LMMEngineParasail(LMMEngine): def __init__( self, base_url=None, api_key=None, model=None, rate_limit=-1, **kwargs diff --git a/gui_agents/s2_5/core/mllm.py b/gui_agents/s2_5/core/mllm.py index cd7cd038..1ee30e55 100644 --- a/gui_agents/s2_5/core/mllm.py +++ b/gui_agents/s2_5/core/mllm.py @@ -5,12 +5,13 @@ from gui_agents.s2_5.core.engine import ( LMMEngineAnthropic, LMMEngineAzureOpenAI, + LMMEngineGemini, LMMEngineHuggingFace, + LMMEngineLiteLLM, LMMEngineOpenAI, LMMEngineOpenRouter, LMMEngineParasail, LMMEnginevLLM, - LMMEngineGemini, ) @@ -35,6 +36,8 @@ def __init__(self, engine_params=None, system_prompt=None, engine=None): self.engine = LMMEngineOpenRouter(**engine_params) elif engine_type == "parasail": self.engine = LMMEngineParasail(**engine_params) + elif engine_type == "litellm": + self.engine = LMMEngineLiteLLM(**engine_params) else: raise ValueError("engine_type is not supported") else: @@ -127,6 +130,7 @@ def add_message( LMMEngineAzureOpenAI, LMMEngineHuggingFace, LMMEngineGemini, + LMMEngineLiteLLM, LMMEngineOpenRouter, LMMEngineParasail, ), diff --git a/gui_agents/s3/core/engine.py b/gui_agents/s3/core/engine.py index 7bf90f14..1d0a05f4 100644 --- a/gui_agents/s3/core/engine.py +++ b/gui_agents/s3/core/engine.py @@ -402,6 +402,37 @@ def generate(self, messages, temperature=0.0, max_new_tokens=None, **kwargs): ) +class LMMEngineLiteLLM(LMMEngine): + def __init__( + self, api_key=None, model=None, rate_limit=-1, temperature=None, **kwargs + ): + assert model is not None, "model must be provided" + self.model = model + self.api_key = api_key + self.request_interval = 0 if rate_limit == -1 else 60.0 / rate_limit + self.temperature = temperature + + @backoff.on_exception( + backoff.expo, (APIConnectionError, APIError, RateLimitError), max_time=60 + ) + def generate(self, messages, temperature=0.0, max_new_tokens=None, **kwargs): + import litellm + + temp = self.temperature if self.temperature is not None else temperature + params = { + "model": self.model, + "messages": messages, + "max_tokens": max_new_tokens if max_new_tokens else 4096, + "temperature": temp, + "drop_params": True, + **kwargs, + } + if self.api_key: + params["api_key"] = self.api_key + response = litellm.completion(**params) + return response.choices[0].message.content + + class LMMEngineParasail(LMMEngine): def __init__( self, base_url=None, api_key=None, model=None, rate_limit=-1, **kwargs diff --git a/gui_agents/s3/core/mllm.py b/gui_agents/s3/core/mllm.py index fb49e4b0..e96392cc 100644 --- a/gui_agents/s3/core/mllm.py +++ b/gui_agents/s3/core/mllm.py @@ -5,12 +5,13 @@ from gui_agents.s3.core.engine import ( LMMEngineAnthropic, LMMEngineAzureOpenAI, + LMMEngineGemini, LMMEngineHuggingFace, + LMMEngineLiteLLM, LMMEngineOpenAI, LMMEngineOpenRouter, LMMEngineParasail, LMMEnginevLLM, - LMMEngineGemini, ) @@ -35,6 +36,8 @@ def __init__(self, engine_params=None, system_prompt=None, engine=None): self.engine = LMMEngineOpenRouter(**engine_params) elif engine_type == "parasail": self.engine = LMMEngineParasail(**engine_params) + elif engine_type == "litellm": + self.engine = LMMEngineLiteLLM(**engine_params) else: raise ValueError(f"engine_type '{engine_type}' is not supported") else: @@ -127,6 +130,7 @@ def add_message( LMMEngineAzureOpenAI, LMMEngineHuggingFace, LMMEngineGemini, + LMMEngineLiteLLM, LMMEngineOpenRouter, LMMEngineParasail, ), diff --git a/requirements.txt b/requirements.txt index 8d1dc9f9..eae877e5 100644 --- a/requirements.txt +++ b/requirements.txt @@ -16,6 +16,7 @@ toml black pytesseract google-genai +litellm>=1.60.0,<2.0.0 # Platform-specific dependencies pyobjc; platform_system == "Darwin" diff --git a/setup.py b/setup.py index 355a9062..8481661d 100644 --- a/setup.py +++ b/setup.py @@ -29,6 +29,7 @@ "toml", "pytesseract", "google-genai", + "litellm>=1.60.0,<2.0.0", 'pywinauto; platform_system == "Windows"', # Only for Windows 'pywin32; platform_system == "Windows"', # Only for Windows ], diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/test_litellm_engine.py b/tests/test_litellm_engine.py new file mode 100644 index 00000000..9c7687e6 --- /dev/null +++ b/tests/test_litellm_engine.py @@ -0,0 +1,235 @@ +"""Tests for the LiteLLM engine across all Agent-S versions (s2, s2_5, s3).""" + +import types as builtin_types +from unittest import mock + +import pytest + + +# --------------------------------------------------------------------------- +# Fake response helpers +# --------------------------------------------------------------------------- + + +class _Msg: + def __init__(self, content="hello"): + self.content = content + + +class _Choice: + def __init__(self, content="hello"): + self.message = _Msg(content=content) + + +class _Response: + def __init__(self, content="hello"): + self.choices = [_Choice(content=content)] + + +# --------------------------------------------------------------------------- +# Helpers to inject a fake litellm module +# --------------------------------------------------------------------------- + + +def _install_fake_litellm(response_content="hello"): + import sys + + fake = builtin_types.ModuleType("litellm") + fake.completion = mock.MagicMock(return_value=_Response(response_content)) + sys.modules["litellm"] = fake + return fake + + +def _uninstall_fake_litellm(): + import sys + + sys.modules.pop("litellm", None) + + +# --------------------------------------------------------------------------- +# s3 engine tests +# --------------------------------------------------------------------------- + + +class TestS3EngineLiteLLM: + def setup_method(self): + self.fake = _install_fake_litellm("s3 response") + + def teardown_method(self): + _uninstall_fake_litellm() + + def test_generate_returns_content(self): + from gui_agents.s3.core.engine import LMMEngineLiteLLM + + engine = LMMEngineLiteLLM(api_key="test-key", model="openai/gpt-4o") + result = engine.generate( + messages=[{"role": "user", "content": "hi"}], + temperature=0.5, + max_new_tokens=100, + ) + assert result == "s3 response" + + def test_generate_passes_drop_params(self): + from gui_agents.s3.core.engine import LMMEngineLiteLLM + + engine = LMMEngineLiteLLM(api_key="test-key", model="anthropic/claude-3-haiku") + engine.generate( + messages=[{"role": "user", "content": "test"}], + ) + call_kwargs = self.fake.completion.call_args[1] + assert call_kwargs["drop_params"] is True + + def test_generate_passes_model(self): + from gui_agents.s3.core.engine import LMMEngineLiteLLM + + engine = LMMEngineLiteLLM(model="bedrock/anthropic.claude-v2") + engine.generate(messages=[{"role": "user", "content": "hi"}]) + call_kwargs = self.fake.completion.call_args[1] + assert call_kwargs["model"] == "bedrock/anthropic.claude-v2" + + def test_generate_forwards_api_key(self): + from gui_agents.s3.core.engine import LMMEngineLiteLLM + + engine = LMMEngineLiteLLM(api_key="sk-test", model="openai/gpt-4o") + engine.generate(messages=[{"role": "user", "content": "hi"}]) + call_kwargs = self.fake.completion.call_args[1] + assert call_kwargs["api_key"] == "sk-test" + + def test_generate_omits_api_key_when_none(self): + from gui_agents.s3.core.engine import LMMEngineLiteLLM + + engine = LMMEngineLiteLLM(api_key=None, model="openai/gpt-4o") + engine.generate(messages=[{"role": "user", "content": "hi"}]) + call_kwargs = self.fake.completion.call_args[1] + assert "api_key" not in call_kwargs + + def test_instance_temperature_overrides_param(self): + from gui_agents.s3.core.engine import LMMEngineLiteLLM + + engine = LMMEngineLiteLLM(model="openai/gpt-4o", temperature=0.9) + engine.generate( + messages=[{"role": "user", "content": "hi"}], + temperature=0.1, + ) + call_kwargs = self.fake.completion.call_args[1] + assert call_kwargs["temperature"] == 0.9 + + def test_default_max_tokens(self): + from gui_agents.s3.core.engine import LMMEngineLiteLLM + + engine = LMMEngineLiteLLM(model="openai/gpt-4o") + engine.generate(messages=[{"role": "user", "content": "hi"}]) + call_kwargs = self.fake.completion.call_args[1] + assert call_kwargs["max_tokens"] == 4096 + + +# --------------------------------------------------------------------------- +# s2 engine tests +# --------------------------------------------------------------------------- + + +class TestS2EngineLiteLLM: + def setup_method(self): + self.fake = _install_fake_litellm("s2 response") + + def teardown_method(self): + _uninstall_fake_litellm() + + def test_generate_returns_content(self): + from gui_agents.s2.core.engine import LMMEngineLiteLLM + + engine = LMMEngineLiteLLM(api_key="test-key", model="openai/gpt-4o") + result = engine.generate( + messages=[{"role": "user", "content": "hi"}], + ) + assert result == "s2 response" + + def test_drop_params_default(self): + from gui_agents.s2.core.engine import LMMEngineLiteLLM + + engine = LMMEngineLiteLLM(model="openai/gpt-4o") + engine.generate(messages=[{"role": "user", "content": "hi"}]) + call_kwargs = self.fake.completion.call_args[1] + assert call_kwargs["drop_params"] is True + + +# --------------------------------------------------------------------------- +# s2_5 engine tests +# --------------------------------------------------------------------------- + + +class TestS25EngineLiteLLM: + def setup_method(self): + self.fake = _install_fake_litellm("s2_5 response") + + def teardown_method(self): + _uninstall_fake_litellm() + + def test_generate_returns_content(self): + from gui_agents.s2_5.core.engine import LMMEngineLiteLLM + + engine = LMMEngineLiteLLM(api_key="test-key", model="openai/gpt-4o") + result = engine.generate( + messages=[{"role": "user", "content": "hi"}], + ) + assert result == "s2_5 response" + + +# --------------------------------------------------------------------------- +# LMMAgent registration tests +# --------------------------------------------------------------------------- + + +class TestLMMAgentRegistration: + def setup_method(self): + self.fake = _install_fake_litellm("agent response") + + def teardown_method(self): + _uninstall_fake_litellm() + + def test_s3_agent_creates_litellm_engine(self): + from gui_agents.s3.core.engine import LMMEngineLiteLLM + from gui_agents.s3.core.mllm import LMMAgent + + agent = LMMAgent( + engine_params={"engine_type": "litellm", "model": "openai/gpt-4o"}, + ) + assert isinstance(agent.engine, LMMEngineLiteLLM) + + def test_s3_agent_get_response(self): + from gui_agents.s3.core.mllm import LMMAgent + + agent = LMMAgent( + engine_params={ + "engine_type": "litellm", + "model": "openai/gpt-4o", + "api_key": "test", + }, + system_prompt="You are helpful.", + ) + resp = agent.get_response(user_message="hi") + assert resp == "agent response" + + def test_s2_agent_creates_litellm_engine(self): + from gui_agents.s2.core.engine import LMMEngineLiteLLM + from gui_agents.s2.core.mllm import LMMAgent + + agent = LMMAgent( + engine_params={"engine_type": "litellm", "model": "openai/gpt-4o"}, + ) + assert isinstance(agent.engine, LMMEngineLiteLLM) + + def test_s2_5_agent_creates_litellm_engine(self): + from gui_agents.s2_5.core.engine import LMMEngineLiteLLM + from gui_agents.s2_5.core.mllm import LMMAgent + + agent = LMMAgent( + engine_params={"engine_type": "litellm", "model": "openai/gpt-4o"}, + ) + assert isinstance(agent.engine, LMMEngineLiteLLM) + + def test_unsupported_engine_raises(self): + from gui_agents.s3.core.mllm import LMMAgent + + with pytest.raises(ValueError, match="not supported"): + LMMAgent(engine_params={"engine_type": "nonexistent", "model": "x"})