mirror of
https://github.com/tiennm99/serena.git
synced 2026-10-03 05:12:50 +00:00
Restore simpler call mechanism with locals(), avoiding the need for forbidden params
This commit is contained in:
1 parent
a5edf6eb2c
commit
5ae12bcaa2
4 files changed
+27
-30
No files matched your search
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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())
|
||||
Reference in new issue
Block a user