mirror of
https://github.com/simular-ai/Agent-S.git
synced 2026-09-01 15:02:27 +08:00
feat: add DeepSeek/Qwen support and improve Ollama config
This commit is contained in:
@@ -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
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user