Restore simpler call mechanism with locals(), avoiding the need for forbidden params

This commit is contained in:
Dominik Jain committed 2025-06-03 12:28:40 +02:00
1 parent a5edf6eb2c
commit 5ae12bcaa2
4 files changed
+27 -30

No files matched your search

+2 -2
View File
@@ -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
+10 -14
View File
@@ -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)
+7 -6
View File
@@ -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"""
@@ -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())