Bucha typing fixes, poe type-check is green

This commit is contained in:
Michael Panchenko committed 2025-03-25 16:48:40 +01:00
1 parent d73e725a87
commit cd380206a1
8 files changed
+44 -24

No files matched your search

+1
View File
@@ -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]
+7 -5
View File
@@ -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 -11
View File
@@ -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"):
+5 -3
View File
@@ -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())
+1 -1
View File
@@ -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)
+5 -2
View File
@@ -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)
+2 -2
View File
@@ -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
Generated
+11
View File
@@ -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"