From d2271de1704fb042cc42ab1748451e931f7d1056 Mon Sep 17 00:00:00 2001 From: lif <1835304752@qq.com> Date: Thu, 15 Jan 2026 17:16:13 +0800 Subject: [PATCH] feat: add DeepSeek/Qwen support and improve Ollama config --- gui_agents/s3/core/engine.py | 93 ++++++++++++++++++++++++++++++++++++ gui_agents/s3/core/mllm.py | 16 +++++-- tests/test_providers.py | 68 ++++++++++++++++++++++++++ 3 files changed, 174 insertions(+), 3 deletions(-) create mode 100644 tests/test_providers.py diff --git a/gui_agents/s3/core/engine.py b/gui_agents/s3/core/engine.py index 21aa019..50829eb 100644 --- a/gui_agents/s3/core/engine.py +++ b/gui_agents/s3/core/engine.py @@ -445,3 +445,96 @@ class LMMEngineParasail(LMMEngine): ) +class LMMEngineDeepSeek(LMMEngine): + def __init__( + self, base_url=None, api_key=None, model=None, rate_limit=-1, **kwargs + ): + assert model is not None, "DeepSeek model id must be provided" + self.base_url = base_url + self.model = model + self.api_key = api_key + self.request_interval = 0 if rate_limit == -1 else 60.0 / rate_limit + self.llm_client = None + + @backoff.on_exception( + backoff.expo, (APIConnectionError, APIError, RateLimitError), max_time=60 + ) + def generate(self, messages, temperature=0.0, max_new_tokens=None, **kwargs): + api_key = self.api_key or os.getenv("DEEPSEEK_API_KEY") + if api_key is None: + raise ValueError( + "A DeepSeek API key needs to be provided in either the api_key parameter or as an environment variable named DEEPSEEK_API_KEY" + ) + base_url = self.base_url or os.getenv("DEEPSEEK_ENDPOINT_URL") + if base_url is None: + base_url = "https://api.deepseek.com" + + if not self.llm_client: + self.llm_client = OpenAI( + base_url=base_url, + api_key=api_key, + ) + return ( + self.llm_client.chat.completions.create( + model=self.model, + messages=messages, + max_tokens=max_new_tokens if max_new_tokens else 4096, + temperature=temperature, + **kwargs, + ) + .choices[0] + .message.content + ) + + +class LMMEngineQwen(LMMEngine): + def __init__( + self, base_url=None, api_key=None, model=None, rate_limit=-1, **kwargs + ): + assert model is not None, "Qwen model id must be provided" + self.base_url = base_url + self.model = model + self.api_key = api_key + self.request_interval = 0 if rate_limit == -1 else 60.0 / rate_limit + self.llm_client = None + + @backoff.on_exception( + backoff.expo, (APIConnectionError, APIError, RateLimitError), max_time=60 + ) + def generate(self, messages, temperature=0.0, max_new_tokens=None, **kwargs): + api_key = self.api_key or os.getenv("QWEN_API_KEY") # Or DASHSCOPE_API_KEY + if api_key is None: + raise ValueError( + "A Qwen API key needs to be provided in either the api_key parameter or as an environment variable named QWEN_API_KEY" + ) + base_url = self.base_url or os.getenv("QWEN_ENDPOINT_URL") + if base_url is None: + # Alibaba Qwen often uses DashScope, but for compatible APIs let's assume standard or user provided + # If strictly Qwen (via DashScope compatible), valid URL is needed. + # defaulting to DashScope compatible endpoint as placeholder or rely on user. + # For this strict implementation, ensuring we have a URL is better if known, + # but generic "Qwen" usually implies usage via compatible interface (like vLLM serving Qwen or DashScope). + # Let's require it or default to a common one if reasonable. + # Given the other engines, let's enforce user providing it or env var if we don't have a single canonical one (DashScope is common). + # Let's default to DashScope's openai compatible endpoint if none provided? + # https://dashscope.aliyuncs.com/compatible-mode/v1 + base_url = "https://dashscope.aliyuncs.com/compatible-mode/v1" + + if not self.llm_client: + self.llm_client = OpenAI( + base_url=base_url, + api_key=api_key, + ) + return ( + self.llm_client.chat.completions.create( + model=self.model, + messages=messages, + max_tokens=max_new_tokens if max_new_tokens else 4096, + temperature=temperature, + **kwargs, + ) + .choices[0] + .message.content + ) + + diff --git a/gui_agents/s3/core/mllm.py b/gui_agents/s3/core/mllm.py index d74ebad..90c1c22 100644 --- a/gui_agents/s3/core/mllm.py +++ b/gui_agents/s3/core/mllm.py @@ -12,6 +12,8 @@ from gui_agents.s3.core.engine import ( LMMEngineParasail, LMMEnginevLLM, LMMEngineGemini, + LMMEngineDeepSeek, + LMMEngineQwen, ) @@ -37,7 +39,7 @@ class LMMAgent: elif engine_type == "parasail": self.engine = LMMEngineParasail(**engine_params) elif engine_type == "ollama": - # Reuse LMMEngineOpenAI for Ollama, defaulting to localhost if not specified + # Reuse LMMEngineOpenAI for Ollama if "base_url" not in engine_params: import os base_url = os.getenv("OLLAMA_HOST") @@ -46,10 +48,17 @@ class LMMAgent: base_url = base_url.rstrip("/") + "/v1" engine_params["base_url"] = base_url else: - engine_params["base_url"] = "http://localhost:11434/v1" + # RAISE ERROR instead of default + raise ValueError( + "Ollama endpoint must be provided via 'base_url' parameter or 'OLLAMA_HOST' environment variable." + ) if "api_key" not in engine_params: engine_params["api_key"] = "ollama" self.engine = LMMEngineOpenAI(**engine_params) + elif engine_type == "deepseek": + self.engine = LMMEngineDeepSeek(**engine_params) + elif engine_type == "qwen": + self.engine = LMMEngineQwen(**engine_params) else: raise ValueError(f"engine_type '{engine_type}' is not supported") else: @@ -144,7 +153,8 @@ class LMMAgent: LMMEngineGemini, LMMEngineOpenRouter, LMMEngineParasail, - + LMMEngineDeepSeek, + LMMEngineQwen, ), ): # infer role from previous message diff --git a/tests/test_providers.py b/tests/test_providers.py new file mode 100644 index 0000000..9543266 --- /dev/null +++ b/tests/test_providers.py @@ -0,0 +1,68 @@ + +import os +import unittest +from unittest.mock import patch, MagicMock +from gui_agents.s3.core.mllm import LMMAgent +from gui_agents.s3.core.engine import LMMEngineOpenAI, LMMEngineDeepSeek, LMMEngineQwen + +class TestProviders(unittest.TestCase): + def setUp(self): + # Clear env vars before each test + if "OLLAMA_HOST" in os.environ: + del os.environ["OLLAMA_HOST"] + if "DEEPSEEK_API_KEY" in os.environ: + del os.environ["DEEPSEEK_API_KEY"] + if "QWEN_API_KEY" in os.environ: + del os.environ["QWEN_API_KEY"] + + def test_ollama_missing_config(self): + """Test that Ollama raises ValueError if no endpoint is provided""" + with self.assertRaises(ValueError) as cm: + LMMAgent(engine_params={"engine_type": "ollama", "model": "llama3"}) + self.assertIn("Ollama endpoint must be provided", str(cm.exception)) + + def test_ollama_valid_config_param(self): + """Test Ollama init with base_url param""" + agent = LMMAgent(engine_params={ + "engine_type": "ollama", + "model": "llama3", + "base_url": "http://example.com/v1" + }) + self.assertIsInstance(agent.engine, LMMEngineOpenAI) + self.assertEqual(agent.engine.base_url, "http://example.com/v1") + + def test_ollama_valid_config_env(self): + """Test Ollama init with OLLAMA_HOST env var""" + with patch.dict(os.environ, {"OLLAMA_HOST": "http://env-host:11434"}): + agent = LMMAgent(engine_params={ + "engine_type": "ollama", + "model": "llama3" + }) + self.assertIsInstance(agent.engine, LMMEngineOpenAI) + # Check for /v1 addition + self.assertEqual(agent.engine.base_url, "http://env-host:11434/v1") + + def test_deepseek_init(self): + """Test DeepSeek initialization""" + with patch.dict(os.environ, {"DEEPSEEK_API_KEY": "sk-test"}): + agent = LMMAgent(engine_params={ + "engine_type": "deepseek", + "model": "deepseek-coder" + }) + self.assertIsInstance(agent.engine, LMMEngineDeepSeek) + # Default URL + self.assertEqual(agent.engine.base_url, None) + # (Note: engine.py logic resolves default at generate() time or if client created, + # but init just stores what's passed. Let's verify prompt generation to ensure it doesn't crash on init) + + def test_qwen_init(self): + """Test Qwen initialization""" + with patch.dict(os.environ, {"QWEN_API_KEY": "sk-qwen"}): + agent = LMMAgent(engine_params={ + "engine_type": "qwen", + "model": "qwen-max" + }) + self.assertIsInstance(agent.engine, LMMEngineQwen) + +if __name__ == "__main__": + unittest.main()