mirror of
https://github.com/tiennm99/serena.git
synced 2026-10-11 12:29:04 +00:00
Bucha typing fixes, poe type-check is green
This commit is contained in:
1 parent
d73e725a87
commit
cd380206a1
8 files changed
+44
-24
No files matched your search
@@ -59,6 +59,7 @@ dev = [
|
||||
"sphinx-toolbox>=3.5.0",
|
||||
"sphinxcontrib-bibtex",
|
||||
"sphinxcontrib-spelling>=8.0.0",
|
||||
"types-pyyaml>=6.0.12.20241230",
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
from typing import Any
|
||||
|
||||
import jinja2
|
||||
import jinja2.meta
|
||||
import jinja2.nodes
|
||||
@@ -8,21 +10,21 @@ from serena.util.class_decorators import singleton
|
||||
|
||||
@singleton
|
||||
class JinjaEnvProvider:
|
||||
def __init__(self):
|
||||
self._env = None
|
||||
def __init__(self) -> None:
|
||||
self._env: jinja2.Environment | None = None
|
||||
|
||||
def get_env(self):
|
||||
def get_env(self) -> jinja2.Environment:
|
||||
if self._env is None:
|
||||
self._env = jinja2.Environment()
|
||||
return self._env
|
||||
|
||||
|
||||
class JinjaTemplate:
|
||||
def __init__(self, template_string: str):
|
||||
def __init__(self, template_string: str) -> None:
|
||||
self._template_string = template_string
|
||||
self._template = JinjaEnvProvider().get_env().from_string(self._template_string)
|
||||
|
||||
def render(self, **kwargs) -> str:
|
||||
def render(self, **kwargs: Any) -> str:
|
||||
return self._template.render(**kwargs)
|
||||
|
||||
def get_parameters(self) -> set[str]:
|
||||
|
||||
@@ -12,7 +12,7 @@ LANG_CODES = ["en", "de"]
|
||||
|
||||
|
||||
class PromptTemplate(ToStringMixin):
|
||||
def __init__(self, name: str, jinja_template_string: str):
|
||||
def __init__(self, name: str, jinja_template_string: str) -> None:
|
||||
self.name = name
|
||||
self.jinja_template = JinjaTemplate(jinja_template_string.strip())
|
||||
self.parameters = self.jinja_template.get_parameters()
|
||||
@@ -20,15 +20,15 @@ class PromptTemplate(ToStringMixin):
|
||||
def _tostring_excludes(self) -> list[str]:
|
||||
return ["jinja_template"]
|
||||
|
||||
def instantiate(self, **kwargs) -> str:
|
||||
def instantiate(self, **kwargs: Any) -> str:
|
||||
return self.jinja_template.render(**kwargs)
|
||||
|
||||
|
||||
class PromptList:
|
||||
def __init__(self, items: list[str]):
|
||||
def __init__(self, items: list[str]) -> None:
|
||||
self.items = [x.strip() for x in items]
|
||||
|
||||
def to_string(self):
|
||||
def to_string(self) -> str:
|
||||
bullet = " * "
|
||||
indent = " " * len(bullet)
|
||||
items = [x.replace("\n", "\n" + indent) for x in self.items]
|
||||
@@ -43,7 +43,7 @@ class MultiLangContainer(Generic[T], ToStringMixin):
|
||||
Represents a container of items which are associated with different languages
|
||||
"""
|
||||
|
||||
def __init__(self, name: str):
|
||||
def __init__(self, name: str) -> None:
|
||||
self.name = name
|
||||
self.lang2item: dict[str, T] = {}
|
||||
|
||||
@@ -63,7 +63,7 @@ class MultiLangContainer(Generic[T], ToStringMixin):
|
||||
If the requested language is not found, raise an exception
|
||||
"""
|
||||
|
||||
def add_item(self, item: T, lang: str = ""):
|
||||
def add_item(self, item: T, lang: str = "") -> None:
|
||||
self.lang2item[lang] = item
|
||||
|
||||
def get_item(self, lang: str, fallback_mode: FallbackMode = FallbackMode.EXCEPTION) -> T:
|
||||
@@ -105,6 +105,7 @@ class MultiLangPromptTemplate(MultiLangContainer[PromptTemplate]):
|
||||
params == prev_params
|
||||
), f"Parameters of MLPT '{self.name}' are inconsistent: {sorted(params)} vs {sorted(prev_params)}"
|
||||
prev_params = params
|
||||
assert prev_params is not None
|
||||
return sorted(prev_params)
|
||||
|
||||
|
||||
@@ -126,7 +127,7 @@ class MultiLangPromptTemplateCollection:
|
||||
The language of all can be set by specifying the key 'lang' in addition to 'prompts'.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self) -> None:
|
||||
self.prompt_templates: dict[str, MultiLangPromptTemplate] = {}
|
||||
self.prompt_lists: dict[str, MultiLangPromptList] = {}
|
||||
prompts_dir = self._prompt_template_folder()
|
||||
@@ -141,7 +142,7 @@ class MultiLangPromptTemplateCollection:
|
||||
prompts_dir = os.path.join(dir_path, "prompts")
|
||||
if os.path.isdir(prompts_dir):
|
||||
break
|
||||
if not os.path.isdir(prompts_dir):
|
||||
if prompts_dir is None or not os.path.isdir(prompts_dir):
|
||||
raise FileNotFoundError("Could not find the 'prompts' directory")
|
||||
return prompts_dir
|
||||
|
||||
@@ -159,7 +160,7 @@ class MultiLangPromptTemplateCollection:
|
||||
|
||||
return container, lang
|
||||
|
||||
def _add_prompt_template(self, prompt_name: str, jinja_prompt_template: str):
|
||||
def _add_prompt_template(self, prompt_name: str, jinja_prompt_template: str) -> None:
|
||||
"""
|
||||
:param prompt_name: a prompt name, which may have a language shortcode suffix (e.g. "_de")
|
||||
:param jinja_prompt_template: the actual prompt string which may contain placeholders/parameters (e.g. "{name}")
|
||||
@@ -167,7 +168,7 @@ class MultiLangPromptTemplateCollection:
|
||||
multilang_prompt_template, lang = self._container_lang(prompt_name, self.prompt_templates, MultiLangPromptTemplate)
|
||||
multilang_prompt_template.add_item(PromptTemplate(prompt_name, jinja_prompt_template), lang=lang)
|
||||
|
||||
def _add_prompt_list(self, prompt_name: str, prompt_list: list[str]):
|
||||
def _add_prompt_list(self, prompt_name: str, prompt_list: list[str]) -> None:
|
||||
"""
|
||||
:param prompt_name: a prompt name, which may have a language shortcode suffix (e.g. "_de")
|
||||
:param prompt_list: a list of prompts
|
||||
@@ -175,7 +176,7 @@ class MultiLangPromptTemplateCollection:
|
||||
multilang_prompt_list, lang = self._container_lang(prompt_name, self.prompt_lists, MultiLangPromptList)
|
||||
multilang_prompt_list.add_item(PromptList(prompt_list), lang=lang)
|
||||
|
||||
def _read_prompt_templates(self, prompts_dir: str):
|
||||
def _read_prompt_templates(self, prompts_dir: str) -> None:
|
||||
for fn in os.listdir(prompts_dir):
|
||||
path = os.path.join(prompts_dir, fn)
|
||||
if fn.endswith(".txt"):
|
||||
|
||||
@@ -4,12 +4,14 @@ from .multilang_prompt import MultiLangContainer, MultiLangPromptTemplateCollect
|
||||
class PromptFactory:
|
||||
# NOTE: This class is auto-generated by gen_prompt_factory.py
|
||||
|
||||
def __init__(self, lang_shortcode: str = "en", fallback_mode=MultiLangContainer.FallbackMode.EXCEPTION):
|
||||
def __init__(
|
||||
self, lang_shortcode: str = "en", fallback_mode: MultiLangContainer.FallbackMode = MultiLangContainer.FallbackMode.EXCEPTION
|
||||
):
|
||||
self.lang_shortcode = lang_shortcode
|
||||
self.collection = MultiLangPromptTemplateCollection()
|
||||
self.fallback_mode = fallback_mode
|
||||
|
||||
def _format_prompt(self, prompt_name: str, kwargs) -> str:
|
||||
def _format_prompt(self, prompt_name: str, kwargs: dict) -> str:
|
||||
del kwargs["self"]
|
||||
mpt = self.collection.get_multilang_prompt_template(prompt_name)
|
||||
return mpt.get_item(self.lang_shortcode, self.fallback_mode).instantiate(**kwargs)
|
||||
@@ -18,5 +20,5 @@ class PromptFactory:
|
||||
mpl = self.collection.get_multilang_prompt_list(prompt_name)
|
||||
return mpl.get_item(self.lang_shortcode, self.fallback_mode)
|
||||
|
||||
def create_onboarding_prompt(self, *, onboarding_file) -> str:
|
||||
def create_onboarding_prompt(self, *, onboarding_file: str) -> str:
|
||||
return self._format_prompt("onboarding_prompt", locals())
|
||||
@@ -52,7 +52,7 @@ class MatchedConsecutiveLines:
|
||||
matched_lines: list[TextLine] = field(default_factory=list)
|
||||
lines_after_matched: list[TextLine] = field(default_factory=list)
|
||||
|
||||
def __post_init__(self):
|
||||
def __post_init__(self) -> None:
|
||||
for line in self.lines:
|
||||
if line.match_type == LineType.BEFORE_MATCH:
|
||||
self.lines_before_matched.append(line)
|
||||
|
||||
@@ -1,7 +1,10 @@
|
||||
def singleton(cls):
|
||||
from typing import Any
|
||||
|
||||
|
||||
def singleton(cls: type[Any]) -> Any:
|
||||
instance = None
|
||||
|
||||
def get_instance(*args, **kwargs):
|
||||
def get_instance(*args: Any, **kwargs: Any) -> Any:
|
||||
nonlocal instance
|
||||
if instance is None:
|
||||
instance = cls(*args, **kwargs)
|
||||
|
||||
@@ -4,7 +4,7 @@ from collections.abc import Sequence
|
||||
|
||||
def scan_directory(
|
||||
path: str,
|
||||
recursive=False,
|
||||
recursive: bool = False,
|
||||
relative_to: str | None = None,
|
||||
ignored_dirs: Sequence[str] = (),
|
||||
ignored_files: Sequence[str] = (),
|
||||
@@ -24,7 +24,7 @@ def scan_directory(
|
||||
rel_base = os.path.abspath(relative_to) if relative_to else None
|
||||
|
||||
# Helper function to check if an item should be ignored
|
||||
def is_ignored(entry_path, ignored_items):
|
||||
def is_ignored(entry_path: str, ignored_items: Sequence[str]) -> bool:
|
||||
entry_name = os.path.basename(entry_path)
|
||||
|
||||
# Check if name is directly in ignored list
|
||||
|
||||
@@ -2114,6 +2114,7 @@ dev = [
|
||||
{ name = "sphinxcontrib-bibtex" },
|
||||
{ name = "sphinxcontrib-spelling" },
|
||||
{ name = "toml-sort" },
|
||||
{ name = "types-pyyaml" },
|
||||
]
|
||||
|
||||
[package.metadata]
|
||||
@@ -2151,6 +2152,7 @@ requires-dist = [
|
||||
{ name = "sphinxcontrib-bibtex", marker = "extra == 'dev'" },
|
||||
{ name = "sphinxcontrib-spelling", marker = "extra == 'dev'", specifier = ">=8.0.0" },
|
||||
{ name = "toml-sort", marker = "extra == 'dev'", specifier = ">=0.24.2" },
|
||||
{ name = "types-pyyaml", marker = "extra == 'dev'", specifier = ">=6.0.12.20241230" },
|
||||
]
|
||||
provides-extras = ["dev"]
|
||||
|
||||
@@ -2709,6 +2711,15 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/0f/b3/ca41df24db5eb99b00d97f89d7674a90cb6b3134c52fb8121b6d8d30f15c/types_python_dateutil-2.9.0.20241206-py3-none-any.whl", hash = "sha256:e248a4bc70a486d3e3ec84d0dc30eec3a5f979d6e7ee4123ae043eedbb987f53", size = 14384 },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "types-pyyaml"
|
||||
version = "6.0.12.20241230"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/9a/f9/4d566925bcf9396136c0a2e5dc7e230ff08d86fa011a69888dd184469d80/types_pyyaml-6.0.12.20241230.tar.gz", hash = "sha256:7f07622dbd34bb9c8b264fe860a17e0efcad00d50b5f27e93984909d9363498c", size = 17078 }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/e8/c1/48474fbead512b70ccdb4f81ba5eb4a58f69d100ba19f17c92c0c4f50ae6/types_PyYAML-6.0.12.20241230-py3-none-any.whl", hash = "sha256:fa4d32565219b68e6dee5f67534c722e53c00d1cfc09c435ef04d7353e1e96e6", size = 20029 },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "typing-extensions"
|
||||
version = "4.12.2"
|
||||
|
||||
Reference in new issue
Block a user