49 lines
1.5 KiB
Python
49 lines
1.5 KiB
Python
"""Factory for creating model adapters from configuration."""
|
|
|
|
from app.config import get_settings
|
|
from app.model_adapters.base import ModelAdapter
|
|
from app.model_adapters.providers import DeepSeekAdapter, KimiAdapter, QwenAdapter
|
|
|
|
|
|
_ADAPTER_MAP = {
|
|
"deepseek": DeepSeekAdapter,
|
|
"kimi": KimiAdapter,
|
|
"qwen": QwenAdapter,
|
|
}
|
|
|
|
|
|
def create_adapter(model_id: str, concurrency: int = 5, max_retries: int = 3) -> ModelAdapter:
|
|
settings = get_settings()
|
|
model_id = model_id.lower()
|
|
adapter_cls = _ADAPTER_MAP.get(model_id)
|
|
if not adapter_cls:
|
|
raise ValueError(f"Unknown model_id: {model_id}. Available: {list(_ADAPTER_MAP.keys())}")
|
|
|
|
if model_id == "deepseek":
|
|
return adapter_cls(
|
|
api_key=settings.deepseek_api_key or "",
|
|
model=settings.deepseek_model,
|
|
base_url=settings.deepseek_base_url,
|
|
concurrency=concurrency,
|
|
max_retries=max_retries,
|
|
)
|
|
if model_id == "kimi":
|
|
return adapter_cls(
|
|
api_key=settings.kimi_api_key or "",
|
|
model=settings.kimi_model,
|
|
base_url=settings.kimi_base_url,
|
|
concurrency=concurrency,
|
|
max_retries=max_retries,
|
|
)
|
|
return adapter_cls(
|
|
api_key=settings.qwen_api_key or "",
|
|
model=settings.qwen_model,
|
|
base_url=settings.qwen_base_url,
|
|
concurrency=concurrency,
|
|
max_retries=max_retries,
|
|
)
|
|
|
|
|
|
def list_models() -> list[str]:
|
|
return list(_ADAPTER_MAP.keys())
|