diff --git a/src/interprompt/jinja_template.py b/src/interprompt/jinja_template.py index bb945f46..63f4fd87 100644 --- a/src/interprompt/jinja_template.py +++ b/src/interprompt/jinja_template.py @@ -30,9 +30,9 @@ class JinjaTemplate(ParameterizedTemplateInterface): parsed_content = self._template.environment.parse(self._template_string) self._parameters = sorted(jinja2.meta.find_undeclared_variables(parsed_content)) - def render(self, **kwargs: Any) -> str: + def render(self, **params: Any) -> str: """Renders the template with the given kwargs. You can find out which parameters are required by calling get_parameter_names().""" - return self._template.render(**kwargs) + return self._template.render(**params) def get_parameters(self) -> list[str]: """A sorted list of parameter names that are extracted from the template string. It is impossible to know the types of the parameter diff --git a/src/interprompt/multilang_prompt.py b/src/interprompt/multilang_prompt.py index 0f0ab012..55777344 100644 --- a/src/interprompt/multilang_prompt.py +++ b/src/interprompt/multilang_prompt.py @@ -1,7 +1,7 @@ import logging import os from enum import Enum -from typing import Any, ClassVar, Generic, TypeVar +from typing import Any, Generic, TypeVar import yaml from sensai.util.string import ToStringMixin @@ -19,8 +19,8 @@ class PromptTemplate(ToStringMixin, ParameterizedTemplateInterface): def _tostring_exclude_private(self) -> bool: return True - def render(self, **kwargs: Any) -> str: - return self._jinja_template.render(**kwargs) + def render(self, **params: Any) -> str: + return self._jinja_template.render(**params) def get_parameters(self) -> list[str]: return self._jinja_template.get_parameters() @@ -123,8 +123,6 @@ class _MultiLangContainer(Generic[T], ToStringMixin): class MultiLangPromptTemplate(ParameterizedTemplateInterface): - _FORBIDDEN_PARAM_NAMES: ClassVar[set[str]] = {"lang_code", "fallback_mode"} # avoid collisions with **kwargs in calls to render - """ Represents a prompt template with support for multiple languages. The parameters of all prompt templates (for all languages) are (must be) the same. @@ -153,11 +151,6 @@ class MultiLangPromptTemplate(ParameterizedTemplateInterface): :param allow_overwrite: whether to allow overwriting an existing entry for the same language """ incoming_parameters = prompt_template.get_parameters() - if contained_fobidden_parameters := self._FORBIDDEN_PARAM_NAMES.intersection(incoming_parameters): - raise ValueError( - f"Cannot add prompt template for language '{lang_code}' to MultiLangPromptTemplate '{self.name}'" - f"since it contains forbidden parameter names: {contained_fobidden_parameters}" - ) if len(self) > 0: parameters = self.get_parameters() if parameters != incoming_parameters: @@ -182,10 +175,13 @@ class MultiLangPromptTemplate(ParameterizedTemplateInterface): return first_prompt_template.get_parameters() def render( - self, lang_code: str = DEFAULT_LANG_CODE, fallback_mode: LanguageFallbackMode = LanguageFallbackMode.EXCEPTION, **kwargs: Any + self, + params: dict[str, Any], + lang_code: str = DEFAULT_LANG_CODE, + fallback_mode: LanguageFallbackMode = LanguageFallbackMode.EXCEPTION, ) -> str: prompt_template = self.get_prompt_template(lang_code, fallback_mode) - return prompt_template.render(**kwargs) + return prompt_template.render(**params) class MultiLangPromptList(_MultiLangContainer[PromptList]): @@ -313,8 +309,8 @@ class MultiLangPromptCollection: def render_prompt_template( self, prompt_name: str, + params: dict[str, Any], lang_code: str = DEFAULT_LANG_CODE, - **kwargs: Any, ) -> str: """Renders the prompt template for the given prompt name and language code.""" - return self.get_prompt_template(prompt_name, lang_code=lang_code).render(**kwargs) + return self.get_prompt_template(prompt_name, lang_code=lang_code).render(**params) diff --git a/src/interprompt/prompt_factory.py b/src/interprompt/prompt_factory.py index 791ee259..3aa9da08 100644 --- a/src/interprompt/prompt_factory.py +++ b/src/interprompt/prompt_factory.py @@ -1,7 +1,8 @@ import logging import os -from .multilang_prompt import MultiLangPromptCollection, DEFAULT_LANG_CODE, LanguageFallbackMode, PromptList +from typing import Any +from .multilang_prompt import DEFAULT_LANG_CODE, LanguageFallbackMode, MultiLangPromptCollection, PromptList log = logging.getLogger(__name__) @@ -20,8 +21,9 @@ class PromptFactoryBase: self.lang_code = lang_code self._prompt_collection = MultiLangPromptCollection(prompts_dir, fallback_mode=fallback_mode) - def _render_prompt(self, prompt_name: str, **kwargs) -> str: - return self._prompt_collection.render_prompt_template(prompt_name, self.lang_code, **kwargs) + def _render_prompt(self, prompt_name: str, params: dict[str, Any]) -> str: + del params["self"] + return self._prompt_collection.render_prompt_template(prompt_name, params, lang_code=self.lang_code) def _get_prompt_list(self, prompt_name: str) -> PromptList: return self._prompt_collection.get_prompt_list(prompt_name, self.lang_code) @@ -48,6 +50,7 @@ from interprompt.multilang_prompt import PromptList from interprompt.prompt_factory import PromptFactoryBase from typing import Any + class PromptFactory(PromptFactoryBase): \""" A class for retrieving and rendering prompt templates and prompt lists. @@ -58,15 +61,13 @@ class PromptFactory(PromptFactoryBase): for template_name in prompt_collection.get_prompt_template_names(): template_parameters = prompt_collection.get_prompt_template_parameters(template_name) - render_call_str = f'"{template_name}"' if len(template_parameters) == 0: method_params_str = "" else: method_params_str = ", *, " + ", ".join([f"{param}: Any" for param in template_parameters]) - render_call_str += ", " + ", ".join([f"{param}={param}" for param in template_parameters]) generated_code += f""" def create_{template_name}(self{method_params_str}) -> str: - return self._render_prompt({render_call_str}) + return self._render_prompt('{template_name}', locals()) """ for prompt_list_name in prompt_collection.get_prompt_list_names(): generated_code += f""" diff --git a/src/serena/generated/generated_prompt_factory.py b/src/serena/generated/generated_prompt_factory.py index 8bbf2c6d..6bb7a76b 100644 --- a/src/serena/generated/generated_prompt_factory.py +++ b/src/serena/generated/generated_prompt_factory.py @@ -1,4 +1,3 @@ - # ruff: noqa # black: skip # mypy: ignore-errors @@ -9,28 +8,29 @@ from interprompt.multilang_prompt import PromptList from interprompt.prompt_factory import PromptFactoryBase from typing import Any + class PromptFactory(PromptFactoryBase): """ A class for retrieving and rendering prompt templates and prompt lists. """ def create_onboarding_prompt(self, *, system: Any) -> str: - return self._render_prompt("onboarding_prompt", system=system) + return self._render_prompt("onboarding_prompt", locals()) def create_think_about_collected_information(self) -> str: - return self._render_prompt("think_about_collected_information") + return self._render_prompt("think_about_collected_information", locals()) def create_think_about_task_adherence(self) -> str: - return self._render_prompt("think_about_task_adherence") + return self._render_prompt("think_about_task_adherence", locals()) def create_think_about_whether_you_are_done(self) -> str: - return self._render_prompt("think_about_whether_you_are_done") + return self._render_prompt("think_about_whether_you_are_done", locals()) def create_summarize_changes(self) -> str: - return self._render_prompt("summarize_changes") + return self._render_prompt("summarize_changes", locals()) def create_prepare_for_new_conversation(self) -> str: - return self._render_prompt("prepare_for_new_conversation") + return self._render_prompt("prepare_for_new_conversation", locals()) def create_system_prompt(self, *, context_system_prompt: Any, mode_system_prompts: Any) -> str: - return self._render_prompt("system_prompt", context_system_prompt=context_system_prompt, mode_system_prompts=mode_system_prompts) + return self._render_prompt("system_prompt", locals())