Skip to content

Commit 864e76e

Browse files
Add Pre-commit-ci (#196)
* Configure Ruff for linting in pyproject.toml Add Ruff configuration for linting with specified rules. * Add ruff pre-commit hooks for linting and formatting * do not apply on notebooks for now * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * fixes * black fixes * fixes * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * remove commended out code * more fixes * type fix * more fixes * more fixes * black fixes * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * black fixes * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * fix * black formatting --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
1 parent f6d57a4 commit 864e76e

12 files changed

Lines changed: 187 additions & 193 deletions

File tree

‎.pre-commit-config.yaml‎

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,10 @@
1+
repos:
2+
- repo: https://github.com/astral-sh/ruff-pre-commit
3+
rev: v0.15.22
4+
hooks:
5+
- id: ruff
6+
name: ruff lint
7+
args: ["--fix"]
8+
files: ^src/python_workflow_definition/
9+
- id: ruff-format
10+
name: ruff format

‎pyproject.toml‎

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,39 @@ plot = [
4242
"ipython>=7.33.0,<=9.8.0",
4343
]
4444

45+
[tool.ruff]
46+
exclude = ["documentation", "example_workflows", "tests", "binder", "_version.py"]
47+
48+
[tool.ruff.lint]
49+
select = [
50+
# pycodestyle
51+
"E",
52+
# Pyflakes
53+
"F",
54+
# pyupgrade
55+
"UP",
56+
# flake8-bugbear
57+
"B",
58+
# flake8-simplify
59+
"SIM",
60+
# isort
61+
"I",
62+
# flake8-comprehensions
63+
"C4",
64+
# eradicate
65+
"ERA",
66+
# pylint
67+
"PL",
68+
]
69+
ignore = [
70+
# ignore line-length violations
71+
"E501",
72+
# Too many statements
73+
"PLR0915",
74+
# Too many branches
75+
"PLR0912",
76+
]
77+
4578
[tool.hatch.build]
4679
include = [
4780
"src/python_workflow_definition"

‎src/python_workflow_definition/aiida.py‎

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -99,7 +99,6 @@ def write_workflow_json(wg: WorkGraph, file_name: str) -> None:
9999
GRAPH_LEVEL_NAMES = ["graph_inputs", "graph_outputs", "graph_ctx"]
100100

101101
for node in wg.tasks:
102-
103102
if node.name in GRAPH_LEVEL_NAMES:
104103
continue
105104

‎src/python_workflow_definition/cwl/__init__.py‎

Lines changed: 47 additions & 55 deletions
Original file line numberDiff line numberDiff line change
@@ -20,17 +20,15 @@
2020

2121

2222
def _get_function_argument(argument: str, position: int = 3) -> dict:
23-
return {
24-
argument
25-
+ "_file": {
26-
"type": "File",
27-
"inputBinding": {
28-
"prefix": "--arg_" + argument + "=",
29-
"separate": False,
30-
"position": position,
31-
},
23+
argument_dict = {
24+
"type": "File",
25+
"inputBinding": {
26+
"prefix": "--arg_" + argument + "=",
27+
"separate": False,
28+
"position": position,
3229
},
3330
}
31+
return {argument + "_file": argument_dict}
3432

3533

3634
def _get_function_template(function_name: str) -> dict:
@@ -44,35 +42,32 @@ def _get_function_template(function_name: str) -> dict:
4442

4543

4644
def _get_output_name(output_name: str) -> dict:
47-
return {
48-
output_name
49-
+ "_file": {"type": "File", "outputBinding": {"glob": output_name + ".pickle"}}
45+
output_dict = {
46+
"type": "File",
47+
"outputBinding": {"glob": output_name + ".pickle"},
5048
}
49+
return {output_name + "_file": output_dict}
5150

5251

5352
def _get_function(workflow):
5453
function_nodes_dict = {
5554
n["id"]: n["value"] for n in workflow[NODES_LABEL] if n["type"] == "function"
5655
}
5756
funct_dict = {}
58-
for funct_id in function_nodes_dict.keys():
57+
for funct_id in function_nodes_dict:
5958
target_ports = list(
60-
set(
61-
[
62-
e[TARGET_PORT_LABEL]
63-
for e in workflow[EDGES_LABEL]
64-
if e["target"] == funct_id
65-
]
66-
)
59+
{
60+
e[TARGET_PORT_LABEL]
61+
for e in workflow[EDGES_LABEL]
62+
if e["target"] == funct_id
63+
}
6764
)
6865
source_ports = list(
69-
set(
70-
[
71-
e[SOURCE_PORT_LABEL]
72-
for e in workflow[EDGES_LABEL]
73-
if e["source"] == funct_id
74-
]
75-
)
66+
{
67+
e[SOURCE_PORT_LABEL]
68+
for e in workflow[EDGES_LABEL]
69+
if e["source"] == funct_id
70+
}
7671
)
7772
funct_dict[funct_id] = {
7873
"targetPorts": target_ports,
@@ -86,7 +81,7 @@ def _write_function_cwl(workflow, directory_path: str = "."):
8681
export_path = Path(directory_path)
8782
export_path.mkdir(parents=True, exist_ok=True)
8883

89-
for i in function_nodes_dict.keys():
84+
for i in function_nodes_dict:
9085
template: dict[str, Any] = {
9186
"cwlVersion": "v1.2",
9287
"class": "CommandLineTool",
@@ -142,10 +137,7 @@ def _write_workflow_config(workflow, directory_path: str = "."):
142137
export_path.mkdir(parents=True, exist_ok=True)
143138
with open(export_path / "workflow.yml", "w") as f:
144139
dump(
145-
{
146-
k + "_file": {"class": "File", "path": k + ".pickle"}
147-
for k in input_dict.keys()
148-
},
140+
{k + "_file": {"class": "File", "path": k + ".pickle"} for k in input_dict},
149141
f,
150142
Dumper=Dumper,
151143
)
@@ -170,7 +162,7 @@ def _write_workflow(workflow, directory_path: str = "."):
170162
last_compute_id = [
171163
e[SOURCE_LABEL] for e in workflow[EDGES_LABEL] if e[TARGET_LABEL] == result_id
172164
][0]
173-
workflow_template["inputs"].update({k + "_file": "File" for k in input_dict.keys()})
165+
workflow_template["inputs"].update({k + "_file": "File" for k in input_dict})
174166
if funct_dict[last_compute_id]["sourcePorts"] == [None]:
175167
workflow_template["outputs"] = {
176168
"result_file": {
@@ -209,29 +201,29 @@ def _write_workflow(workflow, directory_path: str = "."):
209201
for k, v in t[1].items():
210202
if v[SOURCE_LABEL] in input_id_dict:
211203
in_dict[k + "_file"] = input_id_dict[v[SOURCE_LABEL]] + "_file"
204+
elif v["sourcePort"] is None:
205+
in_dict[k + "_file"] = (
206+
step_name_lst[v[SOURCE_LABEL]]
207+
+ "_"
208+
+ str(v[SOURCE_LABEL])
209+
+ "/result_file"
210+
)
212211
else:
213-
if v["sourcePort"] is None:
214-
in_dict[k + "_file"] = (
215-
step_name_lst[v[SOURCE_LABEL]]
216-
+ "_"
217-
+ str(v[SOURCE_LABEL])
218-
+ "/result_file"
219-
)
220-
else:
221-
in_dict[k + "_file"] = (
222-
step_name_lst[v[SOURCE_LABEL]]
223-
+ "_"
224-
+ str(v[SOURCE_LABEL])
225-
+ "/"
226-
+ v[SOURCE_PORT_LABEL]
227-
+ "_file"
228-
)
212+
in_dict[k + "_file"] = (
213+
step_name_lst[v[SOURCE_LABEL]]
214+
+ "_"
215+
+ str(v[SOURCE_LABEL])
216+
+ "/"
217+
+ v[SOURCE_PORT_LABEL]
218+
+ "_file"
219+
)
220+
step_dict = {
221+
"run": node_script,
222+
"in": in_dict,
223+
"out": output,
224+
}
229225
workflow_template["steps"].update(
230-
{
231-
step_name_lst[ind]
232-
+ "_"
233-
+ str(ind): {"run": node_script, "in": in_dict, "out": output}
234-
}
226+
{step_name_lst[ind] + "_" + str(ind): step_dict}
235227
)
236228
export_path = Path(directory_path)
237229
export_path.mkdir(parents=True, exist_ok=True)
@@ -240,7 +232,7 @@ def _write_workflow(workflow, directory_path: str = "."):
240232

241233

242234
def write_workflow(file_name: str, directory_path: str = "."):
243-
with open(file_name, "r") as f:
235+
with open(file_name) as f:
244236
workflow = json.load(f)
245237

246238
_write_function_cwl(workflow=workflow, directory_path=directory_path)

‎src/python_workflow_definition/executorlib.py‎

Lines changed: 2 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -11,10 +11,6 @@
1111
SOURCE_LABEL,
1212
SOURCE_PORT_LABEL,
1313
convert_nodes_list_to_dict,
14-
get_dict,
15-
get_kwargs,
16-
get_list,
17-
get_source_handles,
1814
remove_result,
1915
)
2016

@@ -25,9 +21,9 @@ def get_item(obj, key):
2521

2622
def _get_value(result_dict: dict, nodes_new_dict: dict, link_dict: dict, exe: Executor):
2723
source, source_handle = link_dict[SOURCE_LABEL], link_dict[SOURCE_PORT_LABEL]
28-
if source in result_dict.keys():
24+
if source in result_dict:
2925
result = result_dict[source]
30-
elif source in nodes_new_dict.keys():
26+
elif source in nodes_new_dict:
3127
result = nodes_new_dict[source]
3228
else:
3329
raise KeyError()

‎src/python_workflow_definition/jobflow.py‎

Lines changed: 26 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -27,7 +27,7 @@
2727

2828

2929
def _get_function_dict(flow: Flow):
30-
return {job.uuid: job.function for job in flow.jobs}
30+
return {j.uuid: j.function for j in flow.jobs}
3131

3232

3333
def _get_nodes_dict(function_dict: dict):
@@ -62,8 +62,8 @@ def _get_edges_and_extend_nodes(
6262
flow_dict: dict, nodes_mapping_dict: dict, nodes_dict: dict
6363
):
6464
edges_lst = []
65-
for job in flow_dict["jobs"]:
66-
for k, v in job["function_kwargs"].items():
65+
for j in flow_dict["jobs"]:
66+
for k, v in j["function_kwargs"].items():
6767
if (
6868
isinstance(v, dict)
6969
and "@module" in v
@@ -72,20 +72,18 @@ def _get_edges_and_extend_nodes(
7272
):
7373
edges_lst.append(
7474
_get_edge_from_dict(
75-
target=nodes_mapping_dict[job["uuid"]],
75+
target=nodes_mapping_dict[j["uuid"]],
7676
key=k,
7777
value_dict=v,
7878
nodes_mapping_dict=nodes_mapping_dict,
7979
)
8080
)
8181
elif isinstance(v, dict) and any(
82-
[
83-
isinstance(el, dict)
84-
and "@module" in el
85-
and "@class" in el
86-
and "@version" in el
87-
for el in v.values()
88-
]
82+
isinstance(el, dict)
83+
and "@module" in el
84+
and "@class" in el
85+
and "@version" in el
86+
for el in v.values()
8987
):
9088
node_dict_index = len(nodes_dict)
9189
nodes_dict[node_dict_index] = get_dict
@@ -122,20 +120,18 @@ def _get_edges_and_extend_nodes(
122120
)
123121
edges_lst.append(
124122
{
125-
TARGET_LABEL: nodes_mapping_dict[job["uuid"]],
123+
TARGET_LABEL: nodes_mapping_dict[j["uuid"]],
126124
TARGET_PORT_LABEL: k,
127125
SOURCE_LABEL: node_dict_index,
128126
SOURCE_PORT_LABEL: None,
129127
}
130128
)
131129
elif isinstance(v, list) and any(
132-
[
133-
isinstance(el, dict)
134-
and "@module" in el
135-
and "@class" in el
136-
and "@version" in el
137-
for el in v
138-
]
130+
isinstance(el, dict)
131+
and "@module" in el
132+
and "@class" in el
133+
and "@version" in el
134+
for el in v
139135
):
140136
node_list_index = len(nodes_dict)
141137
nodes_dict[node_list_index] = get_list
@@ -172,7 +168,7 @@ def _get_edges_and_extend_nodes(
172168
)
173169
edges_lst.append(
174170
{
175-
TARGET_LABEL: nodes_mapping_dict[job["uuid"]],
171+
TARGET_LABEL: nodes_mapping_dict[j["uuid"]],
176172
TARGET_PORT_LABEL: k,
177173
SOURCE_LABEL: node_list_index,
178174
SOURCE_PORT_LABEL: None,
@@ -186,7 +182,7 @@ def _get_edges_and_extend_nodes(
186182
node_index = {tv: tk for tk, tv in nodes_dict.items()}[v]
187183
edges_lst.append(
188184
{
189-
TARGET_LABEL: nodes_mapping_dict[job["uuid"]],
185+
TARGET_LABEL: nodes_mapping_dict[j["uuid"]],
190186
TARGET_PORT_LABEL: k,
191187
SOURCE_LABEL: node_index,
192188
SOURCE_PORT_LABEL: None,
@@ -196,10 +192,8 @@ def _get_edges_and_extend_nodes(
196192

197193

198194
def _resort_total_lst(total_dict: dict, nodes_dict: dict) -> dict:
199-
nodes_with_dep_lst = list(sorted(total_dict.keys()))
200-
nodes_without_dep_lst = [
201-
k for k in nodes_dict.keys() if k not in nodes_with_dep_lst
202-
]
195+
nodes_with_dep_lst = sorted(total_dict.keys())
196+
nodes_without_dep_lst = [k for k in nodes_dict if k not in nodes_with_dep_lst]
203197
ordered_lst: list = []
204198
total_new_dict: dict[Any, dict] = {}
205199
while len(total_new_dict) < len(total_dict):
@@ -208,7 +202,7 @@ def _resort_total_lst(total_dict: dict, nodes_dict: dict) -> dict:
208202
if ind not in ordered_lst:
209203
source_lst = [sd[SOURCE_LABEL] for sd in connect.values()]
210204
if all(
211-
[s in ordered_lst or s in nodes_without_dep_lst for s in source_lst]
205+
s in ordered_lst or s in nodes_without_dep_lst for s in source_lst
212206
):
213207
ordered_lst.append(ind)
214208
total_new_dict[ind] = connect
@@ -220,7 +214,7 @@ def _group_edges(edges_lst: list) -> dict:
220214
for ed_major in edges_lst:
221215
target_id = ed_major[TARGET_LABEL]
222216
tmp_lst = []
223-
if target_id not in total_dict.keys():
217+
if target_id not in total_dict:
224218
for ed in edges_lst:
225219
if target_id == ed[TARGET_LABEL]:
226220
tmp_lst.append(ed)
@@ -237,15 +231,15 @@ def _get_workflow(
237231
) -> list:
238232
def get_attr_helper(obj, source_handle):
239233
if source_handle is None:
240-
return getattr(obj, "output")
234+
return obj.output
241235
else:
242-
return getattr(getattr(obj, "output"), source_handle)
236+
return getattr(obj.output, source_handle)
243237

244238
memory_dict: dict[Any, Any] = {}
245-
for k in total_dict.keys():
239+
for k, subdict in total_dict.items():
246240
v = nodes_dict[k]
247241
if isfunction(v):
248-
if k in source_handles_dict.keys():
242+
if k in source_handles_dict:
249243
fn = job(
250244
method=v,
251245
data=[el for el in source_handles_dict[k] if el is not None],
@@ -261,7 +255,7 @@ def get_attr_helper(obj, source_handle):
261255
source_handle=vw[SOURCE_PORT_LABEL],
262256
)
263257
)
264-
for kw, vw in total_dict[k].items()
258+
for kw, vw in subdict.items()
265259
}
266260
memory_dict[k] = fn(**kwargs)
267261
return list(memory_dict.values())

0 commit comments

Comments
 (0)