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,