from __future__ import annotations from typing import Any from adam.executor import ToolContext, ToolExecutionError from adam.model_plugins import ModelPluginRegistry, plugin_function, validate_settings def _plugin_for_tool(context: ToolContext, mode: str): registry = ModelPluginRegistry(context.root) for plugin in registry.all(): if mode == "training" and plugin.trainer_id == context.tool.id: return plugin if mode == "generation" and plugin.generator_id == context.tool.id: return plugin raise ToolExecutionError(f"No model plugin owns {context.tool.id}.") def train(context: ToolContext, **settings: Any) -> dict[str, Any]: plugin = _plugin_for_tool(context, "training") errors = validate_settings(plugin.training_settings, settings) if errors: raise ToolExecutionError(" ".join(errors)) function = plugin_function(plugin, "train") if function is None: raise ToolExecutionError(f"{plugin.name} does not implement train().") return function(settings=settings, callbacks=context) or {} def generate(context: ToolContext, **settings: Any) -> dict[str, Any]: plugin = _plugin_for_tool(context, "generation") errors = validate_settings(plugin.generation_settings, settings) if errors: raise ToolExecutionError(" ".join(errors)) function = plugin_function(plugin, "generate") if function is None: raise ToolExecutionError(f"{plugin.name} does not implement generate().") return function(settings=settings, callbacks=context) or {}