refactor: Remove ModelProvider enum
linuztx committed
Jul 16, 2025 at 18:00 UTC
4c45b5cf8c107d4f4e02fee09bbb496405e3243f
5 files changed
+49
-51
initialize.py
+6
-6
@@ -1,5 +1,5 @@
1
-import models
1
from agent import AgentConfig
2
+import models
3
from python.helpers import runtime, settings, defer
4
from python.helpers.print_style import PrintStyle
5
@@ -28,7 +28,7 @@ def initialize_agent():
28
# chat model from user settings
29
chat_llm = models.ModelConfig(
30
type=models.ModelType.CHAT,
31
- provider=models.ModelProvider[current_settings["chat_model_provider"]],
31
+ provider=current_settings["chat_model_provider"],
32
name=current_settings["chat_model_name"],
33
api_base=current_settings["chat_model_api_base"],
34
ctx_length=current_settings["chat_model_ctx_length"],
@@ -42,7 +42,7 @@ def initialize_agent():
42
# utility model from user settings
43
utility_llm = models.ModelConfig(
44
type=models.ModelType.CHAT,
45
- provider=models.ModelProvider[current_settings["util_model_provider"]],
45
+ provider=current_settings["util_model_provider"],
46
name=current_settings["util_model_name"],
47
api_base=current_settings["util_model_api_base"],
48
ctx_length=current_settings["util_model_ctx_length"],
@@ -54,7 +54,7 @@ def initialize_agent():
54
# embedding model from user settings
55
embedding_llm = models.ModelConfig(
56
type=models.ModelType.EMBEDDING,
57
- provider=models.ModelProvider[current_settings["embed_model_provider"]],
57
+ provider=current_settings["embed_model_provider"],
58
name=current_settings["embed_model_name"],
59
api_base=current_settings["embed_model_api_base"],
60
limit_requests=current_settings["embed_model_rl_requests"],
@@ -63,7 +63,7 @@ def initialize_agent():
63
# browser model from user settings
64
browser_llm = models.ModelConfig(
65
type=models.ModelType.CHAT,
66
- provider=models.ModelProvider[current_settings["browser_model_provider"]],
66
+ provider=current_settings["browser_model_provider"],
67
name=current_settings["browser_model_name"],
68
api_base=current_settings["browser_model_api_base"],
69
vision=current_settings["browser_model_vision"],
@@ -77,7 +77,7 @@ def initialize_agent():
77
browser_model=browser_llm,
78
prompts_subdir=current_settings["agent_prompts_subdir"],
79
memory_subdir=current_settings["agent_memory_subdir"],
80
- knowledge_subdirs=["default", current_settings["agent_knowledge_subdir"]],
80
+ knowledge_subdirs=[current_settings["agent_knowledge_subdir"], "default"],
81
mcp_servers=current_settings["mcp_servers"],
82
code_exec_docker_enabled=False,
83
# code_exec_docker_name = "A0-dev",
models.py
+11
-27
@@ -58,26 +58,10 @@ class ModelType(Enum):
58
EMBEDDING = "Embedding"
59
60
61
-class ModelProvider(Enum):
62
- ANTHROPIC = "Anthropic"
63
- DEEPSEEK = "DeepSeek"
64
- GEMINI = "Google"
65
- GROQ = "Groq"
66
- HUGGINGFACE = "HuggingFace"
67
- LM_STUDIO = "LM Studio"
68
- MISTRAL = "Mistral AI"
69
- OLLAMA = "Ollama"
70
- OPENAI = "OpenAI"
71
- AZURE = "OpenAI Azure"
72
- OPENROUTER = "OpenRouter"
73
- SAMBANOVA = "Sambanova"
74
- OTHER = "Other OpenAI compatible"
75
-
76
-
61
@dataclass
62
class ModelConfig:
63
type: ModelType
80
- provider: ModelProvider
64
+ provider: str
65
name: str
66
api_base: str = ""
67
ctx_length: int = 0
@@ -114,9 +98,9 @@ def get_api_key(service: str) -> str:
98
99
100
def get_rate_limiter(
117
- provider: ModelProvider, name: str, requests: int, input: int, output: int
101
+ provider: str, name: str, requests: int, input: int, output: int
102
) -> RateLimiter:
119
- key = f"{provider.name}\\{name}"
103
+ key = f"{provider}\\{name}"
104
rate_limiters[key] = limiter = rate_limiters.get(key, RateLimiter(seconds=60))
105
limiter.limits["requests"] = requests or 0
106
limiter.limits["input"] = input or 0
@@ -467,8 +451,8 @@ def _adjust_call_args(provider_name: str, model_name: str, kwargs: dict):
451
return provider_name, model_name, kwargs
452
453
470
-def get_model(type: ModelType, provider: ModelProvider, name: str, **kwargs: Any):
471
- provider_name = provider.name.lower()
454
+def get_model(type: ModelType, provider: str, name: str, **kwargs: Any):
455
+ provider_name = provider.lower()
456
if type == ModelType.CHAT:
457
return _get_litellm_chat(LiteLLMChatWrapper, name, provider_name, **kwargs)
458
elif type == ModelType.EMBEDDING:
@@ -478,17 +462,17 @@ def get_model(type: ModelType, provider: ModelProvider, name: str, **kwargs: Any
462
463
464
def get_chat_model(
481
- provider: ModelProvider, name: str, **kwargs: Any
465
+ provider: str, name: str, **kwargs: Any
466
) -> LiteLLMChatWrapper:
483
- provider_name = provider.name.lower()
467
+ provider_name = provider.lower()
468
model = _get_litellm_chat(LiteLLMChatWrapper, name, provider_name, **kwargs)
469
return model
470
471
472
def get_browser_model(
489
- provider: ModelProvider, name: str, **kwargs: Any
473
+ provider: str, name: str, **kwargs: Any
474
) -> BrowserCompatibleChatWrapper:
491
- provider_name = provider.name.lower()
475
+ provider_name = provider.lower()
476
model = _get_litellm_chat(
477
BrowserCompatibleChatWrapper, name, provider_name, **kwargs
478
)
@@ -496,8 +480,8 @@ def get_browser_model(
480
481
482
def get_embedding_model(
499
- provider: ModelProvider, name: str, **kwargs: Any
483
+ provider: str, name: str, **kwargs: Any
484
) -> LiteLLMEmbeddingWrapper | LocalSentenceTransformerWrapper:
501
- provider_name = provider.name.lower()
485
+ provider_name = provider.lower()
486
model = _get_litellm_embedding(name, provider_name, **kwargs)
487
return model
preload.py
+2
-2
@@ -21,11 +21,11 @@ async def preload():
21
22
# preload embedding model
23
async def preload_embedding():
24
- if set["embed_model_provider"] == models.ModelProvider.HUGGINGFACE.name:
24
+ if set["embed_model_provider"] == "HuggingFace":
25
try:
26
# Use the new LiteLLM-based model system
27
emb_mod = models.get_embedding_model(
28
- models.ModelProvider.HUGGINGFACE, set["embed_model_name"]
28
+ "HuggingFace", set["embed_model_name"]
29
)
30
emb_txt = await emb_mod.aembed_query("test")
31
return emb_txt
python/helpers/memory.py
+3
-3
@@ -129,7 +129,7 @@ class Memory:
129
**model_config.build_kwargs(),
130
)
131
embeddings_model_id = files.safe_file_name(
132
- model_config.provider.name + "_" + model_config.name
132
+ model_config.provider + "_" + model_config.name
133
)
134
135
# here we setup the embeddings model with the chosen cache storage
@@ -160,7 +160,7 @@ class Memory:
160
if files.exists(emb_set_file):
161
embedding_set = json.loads(files.read_file(emb_set_file))
162
if (
163
- embedding_set["model_provider"] == model_config.provider.name
163
+ embedding_set["model_provider"] == model_config.provider
164
and embedding_set["model_name"] == model_config.name
165
):
166
# model matches
@@ -200,7 +200,7 @@ class Memory:
200
meta_file_path,
201
json.dumps(
202
{
203
- "model_provider": model_config.provider.name,
203
+ "model_provider": model_config.provider,
204
"model_name": model_config.name,
205
}
206
),
python/helpers/settings.py
+27
-13
@@ -121,9 +121,25 @@ PASSWORD_PLACEHOLDER = "****PSWD****"
121
SETTINGS_FILE = files.get_abs_path("tmp/settings.json")
122
_settings: Settings | None = None
123
124
+# TODO: this is temporary, will be replaced by a proper solution
125
+PROVIDERS: list[FieldOption] = [
126
+ {"value": "ANTHROPIC", "label": "Anthropic"},
127
+ {"value": "DEEPSEEK", "label": "DeepSeek"},
128
+ {"value": "GEMINI", "label": "Google"},
129
+ {"value": "GROQ", "label": "Groq"},
130
+ {"value": "HUGGINGFACE", "label": "HuggingFace"},
131
+ {"value": "LM_STUDIO", "label": "LM Studio"},
132
+ {"value": "MISTRAL", "label": "Mistral AI"},
133
+ {"value": "OLLAMA", "label": "Ollama"},
134
+ {"value": "OPENAI", "label": "OpenAI"},
135
+ {"value": "AZURE", "label": "OpenAI Azure"},
136
+ {"value": "OPENROUTER", "label": "OpenRouter"},
137
+ {"value": "SAMBANOVA", "label": "Sambanova"},
138
+ {"value": "OTHER", "label": "Other OpenAI compatible"},
139
+]
140
+
141
142
def convert_out(settings: Settings) -> SettingsOutput:
126
- from models import ModelProvider
143
default_settings = get_default_settings()
144
145
# main model section
@@ -135,7 +151,7 @@ def convert_out(settings: Settings) -> SettingsOutput:
151
"description": "Select provider for main chat model used by Agent Zero",
152
"type": "select",
153
"value": settings["chat_model_provider"],
138
- "options": [{"value": p.name, "label": p.value} for p in ModelProvider],
154
+ "options": PROVIDERS,
155
}
156
)
157
chat_model_fields.append(
@@ -248,7 +264,7 @@ def convert_out(settings: Settings) -> SettingsOutput:
264
"description": "Select provider for utility model used by the framework",
265
"type": "select",
266
"value": settings["util_model_provider"],
251
- "options": [{"value": p.name, "label": p.value} for p in ModelProvider],
267
+ "options": PROVIDERS,
268
}
269
)
270
util_model_fields.append(
@@ -328,7 +344,7 @@ def convert_out(settings: Settings) -> SettingsOutput:
344
"description": "Select provider for embedding model used by the framework",
345
"type": "select",
346
"value": settings["embed_model_provider"],
331
- "options": [{"value": p.name, "label": p.value} for p in ModelProvider],
347
+ "options": PROVIDERS,
348
}
349
)
350
embed_model_fields.append(
@@ -398,7 +414,7 @@ def convert_out(settings: Settings) -> SettingsOutput:
414
"description": "Select provider for web browser model used by <a href='https://github.com/browser-use/browser-use' target='_blank'>browser-use</a> framework",
415
"type": "select",
416
"value": settings["browser_model_provider"],
401
- "options": [{"value": p.name, "label": p.value} for p in ModelProvider],
417
+ "options": PROVIDERS,
418
}
419
)
420
browser_model_fields.append(
@@ -499,8 +515,8 @@ def convert_out(settings: Settings) -> SettingsOutput:
515
# api keys model section
516
api_keys_fields: list[SettingsField] = []
517
502
- for provider in ModelProvider:
503
- api_keys_fields.append(_get_api_key_field(settings, provider.name.lower(), provider.value))
518
+ for provider in PROVIDERS:
519
+ api_keys_fields.append(_get_api_key_field(settings, provider["value"].lower(), provider["label"]))
520
521
api_keys_section: SettingsSection = {
522
"id": "api_keys",
@@ -990,11 +1006,9 @@ def _write_sensitive_settings(settings: Settings):
1006
1007
1008
def get_default_settings() -> Settings:
993
- from models import ModelProvider
994
-
1009
return Settings(
1010
version=_get_version(),
997
- chat_model_provider=ModelProvider.OPENROUTER.name,
1011
+ chat_model_provider="OPENROUTER",
1012
chat_model_name="openai/gpt-4.1",
1013
chat_model_api_base="",
1014
chat_model_kwargs={"temperature": "0"},
@@ -1004,7 +1018,7 @@ def get_default_settings() -> Settings:
1018
chat_model_rl_requests=0,
1019
chat_model_rl_input=0,
1020
chat_model_rl_output=0,
1007
- util_model_provider=ModelProvider.OPENROUTER.name,
1021
+ util_model_provider="OPENROUTER",
1022
util_model_name="openai/gpt-4.1-nano",
1023
util_model_api_base="",
1024
util_model_ctx_length=100000,
@@ -1013,13 +1027,13 @@ def get_default_settings() -> Settings:
1027
util_model_rl_requests=0,
1028
util_model_rl_input=0,
1029
util_model_rl_output=0,
1016
- embed_model_provider=ModelProvider.HUGGINGFACE.name,
1030
+ embed_model_provider="HUGGINGFACE",
1031
embed_model_name="sentence-transformers/all-MiniLM-L6-v2",
1032
embed_model_api_base="",
1033
embed_model_kwargs={},
1034
embed_model_rl_requests=0,
1035
embed_model_rl_input=0,
1022
- browser_model_provider=ModelProvider.OPENROUTER.name,
1036
+ browser_model_provider="OPENROUTER",
1037
browser_model_name="openai/gpt-4.1",
1038
browser_model_api_base="",
1039
browser_model_vision=True,