Central factory for creating and managing language model instances.
| 24 | |
| 25 | |
| 26 | class ModelFactory: |
| 27 | """Central factory for creating and managing language model instances.""" |
| 28 | |
| 29 | registered_models = {} |
| 30 | _model_cache = {} |
| 31 | _model_stats = {} |
| 32 | |
| 33 | @staticmethod |
| 34 | def create_model(config: Dict[str, Any]) -> BaseModel: |
| 35 | """ |
| 36 | Create a model instance based on provided configuration. |
| 37 | |
| 38 | Args: |
| 39 | config: Configuration dictionary for model creation |
| 40 | |
| 41 | Returns: |
| 42 | Configured model instance |
| 43 | |
| 44 | Raises: |
| 45 | ValueError: For unsupported providers or configuration errors |
| 46 | """ |
| 47 | provider = config.get("provider") |
| 48 | if provider == "default" or provider is None: |
| 49 | provider = config.get("default_provider", "openai") |
| 50 | |
| 51 | if "models" in config and provider in config["models"]: |
| 52 | provider_config = config["models"][provider].copy() |
| 53 | logger.info(f"Using provider-specific config for {provider}: {provider_config}") |
| 54 | |
| 55 | # Merge with general configuration |
| 56 | provider_config.update({ |
| 57 | k: v for k, v in config.items() |
| 58 | if k not in ["models", "agents", "workflow", "tools", "memory", "execution"] |
| 59 | }) |
| 60 | config_to_use = provider_config |
| 61 | else: |
| 62 | logger.info(f"Using fallback config for provider {provider}") |
| 63 | config_to_use = config |
| 64 | |
| 65 | cache_key = ModelFactory._create_cache_key(provider, config_to_use) |
| 66 | |
| 67 | if cache_key in ModelFactory._model_cache: |
| 68 | logger.debug(f"Reusing cached model for provider: {provider}") |
| 69 | return ModelFactory._model_cache[cache_key] |
| 70 | |
| 71 | if provider not in MODEL_PROVIDER_MAP and provider not in ModelFactory.registered_models: |
| 72 | raise ValueError(f"Unsupported model provider: {provider}") |
| 73 | |
| 74 | if provider in ModelFactory.registered_models: |
| 75 | model_class = ModelFactory.registered_models[provider] |
| 76 | else: |
| 77 | try: |
| 78 | module_path, class_name = MODEL_PROVIDER_MAP[provider].rsplit(".", 1) |
| 79 | module = importlib.import_module(module_path) |
| 80 | model_class = getattr(module, class_name) |
| 81 | except (ImportError, AttributeError) as e: |
| 82 | logger.error(f"Failed to load model implementation: {e}") |
| 83 | raise ValueError(f"Implementation not available for provider: {provider}") |