diff --git a/.gitignore b/.gitignore index cff0601..f97ff98 100644 --- a/.gitignore +++ b/.gitignore @@ -5,5 +5,4 @@ test.ipynb .vscode .venv .pytest_cache -.coverage -build +.coverage \ No newline at end of file diff --git a/corrlib/__init__.py b/corrlib/__init__.py index 7cd9c9d..4e1b364 100644 --- a/corrlib/__init__.py +++ b/corrlib/__init__.py @@ -15,11 +15,10 @@ For now, we are interested in collecting primary IObservables only, as these are __app_name__ = "corrlib" -from . import input as input -from .find import find_project as find_project -from .find import find_record as find_record -from .find import list_projects as list_projects +from .import input as input from .initialization import create as create from .meas_io import load_record as load_record from .meas_io import load_records as load_records -from .toml import import_toml \ No newline at end of file +from .find import find_project as find_project +from .find import find_record as find_record +from .find import list_projects as list_projects diff --git a/corrlib/__main__.py b/corrlib/__main__.py index e719ee7..24f9c83 100644 --- a/corrlib/__main__.py +++ b/corrlib/__main__.py @@ -1,4 +1,4 @@ -from corrlib import __app_name__, cli +from corrlib import cli, __app_name__ def main() -> None: diff --git a/corrlib/cli.py b/corrlib/cli.py index 9d7fcfe..a28c837 100644 --- a/corrlib/cli.py +++ b/corrlib/cli.py @@ -1,18 +1,19 @@ +from typing import Optional +import typer +from corrlib import __app_name__ + +from .initialization import create +from .toml import import_tomls, update_project, reimport_project +from .find import find_record, list_projects, list_ensembles, get_stat +from .tools import str2list +from .main import update_aliases +from .meas_io import drop_cache as mio_drop_cache +from .integrity import full_integrity_check + import os from importlib.metadata import version from pathlib import Path -import typer - -from corrlib import __app_name__ - -from .find import find_record, get_stat, list_ensembles, list_projects -from .initialization import create -from .integrity import full_integrity_check -from .main import update_aliases -from .meas_io import drop_cache as mio_drop_cache -from .toml import import_tomls, reimport_project, update_project -from .tools import str2list app = typer.Typer() @@ -25,7 +26,7 @@ def _version_callback(value: bool) -> None: @app.command() def update( - path: Path = typer.Option( # noqa: B008 + path: Path = typer.Option( Path('.'), "--dataset", "-d", @@ -41,7 +42,7 @@ def update( @app.command() def lister( - path: Path = typer.Option( # noqa: B008 + path: Path = typer.Option( Path('.'), "--dataset", "-d", @@ -72,7 +73,7 @@ def lister( @app.command() def alias_add( - path: Path = typer.Option( # noqa: B008 + path: Path = typer.Option( Path('.'), "--dataset", "-d", @@ -90,7 +91,7 @@ def alias_add( @app.command() def find( - path: Path = typer.Option( # noqa: B008 + path: Path = typer.Option( Path('.'), "--dataset", "-d", @@ -99,7 +100,7 @@ def find( corr: str = typer.Argument(), code: str = typer.Argument(), arg: str = typer.Option( - 'all', + str('all'), "--argument", "-a", ), @@ -124,7 +125,7 @@ def find( @app.command() def stat( - path: Path = typer.Option( # noqa: B008 + path: Path = typer.Option( Path('.'), "--dataset", "-d", @@ -140,7 +141,7 @@ def stat( @app.command() -def check(path: Path = typer.Option( # noqa: B008 +def check(path: Path = typer.Option( Path('.'), "--dataset", "-d", @@ -154,7 +155,7 @@ def check(path: Path = typer.Option( # noqa: B008 @app.command() def importer( - path: Path = typer.Option( # noqa: B008 + path: Path = typer.Option( Path('.'), "--dataset", "-d", @@ -162,7 +163,7 @@ def importer( files: str = typer.Argument( ), copy_file: bool = typer.Option( - True, + bool(True), "--save", "-s", ), @@ -178,7 +179,7 @@ def importer( @app.command() def reimporter( - path: Path = typer.Option( # noqa: B008 + path: Path = typer.Option( Path('.'), "--dataset", "-d", @@ -203,13 +204,13 @@ def reimporter( @app.command() def init( - path: Path = typer.Option( # noqa: B008 + path: Path = typer.Option( Path('.'), "--dataset", "-d", ), tracker: str = typer.Option( - 'datalad', + str('datalad'), "--tracker", "-t", ), @@ -223,7 +224,7 @@ def init( @app.command() def drop_cache( - path: Path = typer.Option( # noqa: B008 + path: Path = typer.Option( Path('.'), "--dataset", "-d", @@ -238,7 +239,7 @@ def drop_cache( @app.callback() def main( - version: bool | None = typer.Option( + version: Optional[bool] = typer.Option( None, "--version", "-v", diff --git a/corrlib/find.py b/corrlib/find.py index 18c4001..af21a4d 100644 --- a/corrlib/find.py +++ b/corrlib/find.py @@ -1,22 +1,20 @@ -import datetime as dt -import json -import os import sqlite3 -import warnings -from collections.abc import Callable -from pathlib import Path -from typing import Any - -import numpy as np +import os +import json import pandas as pd -from pyerrors import Corr, Obs - +import numpy as np from .input.implementations import codes -from .integrity import has_valid_times -from .meas_io import load_record -from .sql import thin_sql_wrapper -from .tools import get_db_file, k2m +from .tools import k2m, get_db_file from .tracker import get +from .integrity import has_valid_times +from .sql import thin_sql_wrapper +from typing import Any, Optional +from pathlib import Path +import datetime as dt +from collections.abc import Callable +import warnings +from .meas_io import load_record +from pyerrors import Corr, Obs def _project_lookup_by_alias(path: Path, alias: str) -> str: @@ -65,7 +63,7 @@ def _project_lookup_by_id(path: Path, uuid: str) -> list[tuple[str, ...]]: return results -def _time_filter(results: pd.DataFrame, created_before: str | None=None, created_after: str | None=None, updated_before: str | None=None, updated_after: str | None=None) -> pd.DataFrame: +def _time_filter(results: pd.DataFrame, created_before: Optional[str]=None, created_after: Optional[Any]=None, updated_before: Optional[Any]=None, updated_after: Optional[Any]=None) -> pd.DataFrame: """ Filter the results from the database in terms of the creation and update times. @@ -114,7 +112,7 @@ def _time_filter(results: pd.DataFrame, created_before: str | None=None, create return results.drop(drops) -def _db_lookup(db: Path, ensemble: str, correlator_name: str, code: str, project: str | None=None, parameters: str | None=None) -> pd.DataFrame: +def _db_lookup(db: Path, ensemble: str, correlator_name: str, code: str, project: Optional[str]=None, parameters: Optional[str]=None) -> pd.DataFrame: """ Look up a correlator record in the database by the data given to the method. @@ -288,7 +286,7 @@ def openQCD_filter(results:pd.DataFrame, **kwargs: Any) -> pd.DataFrame: The filtered results. """ - warnings.warn("A filter for openQCD parameters is no implemented yet.", Warning, 1) + warnings.warn("A filter for openQCD parameters is no implemented yet.", Warning) return results @@ -321,10 +319,10 @@ def _code_filter(results: pd.DataFrame, code: str, **kwargs: Any) -> pd.DataFram raise ValueError(f"Code {code} is not known.") -def find_record(path: Path, ensemble: str, correlator_name: str, code: str, project: str | None=None, parameters: str | None=None, - created_before: str | None=None, created_after: str | None=None, updated_before: str | None=None, updated_after: str | None=None, - revision: str | None=None, - customFilter: Callable[[pd.DataFrame], pd.DataFrame] | None = None, +def find_record(path: Path, ensemble: str, correlator_name: str, code: str, project: Optional[str]=None, parameters: Optional[str]=None, + created_before: Optional[str]=None, created_after: Optional[str]=None, updated_before: Optional[str]=None, updated_after: Optional[str]=None, + revision: Optional[str]=None, + customFilter: Optional[Callable[[pd.DataFrame], pd.DataFrame]] = None, **kwargs: Any) -> pd.DataFrame: path = Path(path) db_file = get_db_file(path) diff --git a/corrlib/git_tools.py b/corrlib/git_tools.py index 7808ada..d77f109 100644 --- a/corrlib/git_tools.py +++ b/corrlib/git_tools.py @@ -1,9 +1,7 @@ import os -from pathlib import Path - -import git - from .tracker import save +import git +from pathlib import Path GITMODULES_FILE = '.gitmodules' @@ -27,7 +25,7 @@ def move_submodule(repo_path: Path, old_path: Path, new_path: Path) -> None: gitmodules_file_path = repo_path / GITMODULES_FILE # update paths in .gitmodules - with open(gitmodules_file_path) as file: + with open(gitmodules_file_path, 'r') as file: lines = [line.strip() for line in file] updated_lines = [] diff --git a/corrlib/initialization.py b/corrlib/initialization.py index 83a3b44..99f90d1 100644 --- a/corrlib/initialization.py +++ b/corrlib/initialization.py @@ -1,10 +1,9 @@ -import os -import sqlite3 from configparser import ConfigParser +import sqlite3 +import os +from .tracker import save, init from pathlib import Path - from .tools import CONFIG_FILENAME -from .tracker import init, save def _create_db(db: Path) -> None: diff --git a/corrlib/input/__init__.py b/corrlib/input/__init__.py index 3ccbf28..be6d6b2 100644 --- a/corrlib/input/__init__.py +++ b/corrlib/input/__init__.py @@ -2,6 +2,6 @@ Import functions for different codes. """ -from . import implementations as implementations -from . import openQCD as openQCD from . import sfcf as sfcf +from . import openQCD as openQCD +from . import implementations as implementations diff --git a/corrlib/input/openQCD.py b/corrlib/input/openQCD.py index eb17fd3..9dfe85f 100644 --- a/corrlib/input/openQCD.py +++ b/corrlib/input/openQCD.py @@ -1,16 +1,16 @@ -import fnmatch -import os -from pathlib import Path -from typing import Any - -import datalad.api as dl -import matplotlib.pyplot as plt import pyerrors.input.openQCD as input - -from ..pars.openQCD import ms1, qcd2 +import datalad.api as dl +import os +import fnmatch +from typing import Any, Optional +from pathlib import Path +import matplotlib.pyplot as plt +from ..pars.openQCD import ms1 +from ..pars.openQCD import qcd2 from ..tools import get_plot_dir + def load_ms1_infile(path: Path, project: str, file_in_project: str) -> dict[str, Any]: """ Read the parameters for ms1 measurements from a parameter file in the project. @@ -31,18 +31,16 @@ def load_ms1_infile(path: Path, project: str, file_in_project: str) -> dict[str, """ file = os.path.join(path, "projects", project, file_in_project) - if not os.path.exists(file): - raise OSError(f"File {file} does not exist.") ds = os.path.join(path, "projects", project) dl.get(file, dataset=ds) - with open(file) as fp: + with open(file, 'r') as fp: lines = fp.readlines() fp.close() param: dict[str, Any] = {} param['rw_fcts'] = [] param['rand'] = {} - for line in lines: + for i, line in enumerate(lines): if line.startswith('#'): continue if line.startswith('\n'): @@ -99,7 +97,7 @@ def load_ms3_infile(path: Path, project: str, file_in_project: str) -> dict[str, file = os.path.join(path, "projects", project, file_in_project) ds = os.path.join(path, "projects", project) dl.get(file, dataset=ds) - with open(file) as fp: + with open(file, 'r') as fp: lines = fp.readlines() fp.close() param = {} @@ -111,7 +109,7 @@ def load_ms3_infile(path: Path, project: str, file_in_project: str) -> dict[str, return param -def read_rwms(path: Path, project: str, dir_in_project: str, param: dict[str, Any], prefix: str, postfix: str="ms1", version: str='2.0', names: list[str] | None=None, files: list[str] | None=None) -> dict[str, Any]: +def read_rwms(path: Path, project: str, dir_in_project: str, param: dict[str, Any], prefix: str, postfix: str="ms1", version: str='2.0', names: Optional[list[str]]=None, files: Optional[list[str]]=None) -> dict[str, Any]: """ Read reweighting factor measurements from the project. @@ -146,7 +144,7 @@ def read_rwms(path: Path, project: str, dir_in_project: str, param: dict[str, An directory = os.path.join(dataset, dir_in_project) if files is None: files = [] - for _root, _ds, fs in os.walk(directory): + for root, ds, fs in os.walk(directory): for f in fs: if fnmatch.fnmatch(f, prefix + "*" + postfix + ".dat"): files.append(f) @@ -168,8 +166,8 @@ def read_rwms(path: Path, project: str, dir_in_project: str, param: dict[str, An return rw_dict -def extract_t0(path: Path, project: str, dir_in_project: str, ensemble: str, param: dict[str, Any], prefix: str, dtr_read: int, xmin: int, spatial_extent: int, fit_range: int = 5, postfix: str="", names: list[str] | None=None, files: list[str] | None=None, - r_start: list[int] | None=None, r_stop: list[int] | None=None, r_step:int=1) -> dict[str, Any]: +def extract_t0(path: Path, project: str, dir_in_project: str, ensemble: str, param: dict[str, Any], prefix: str, dtr_read: int, xmin: int, spatial_extent: int, fit_range: int = 5, postfix: str="", names: Optional[list[str]]=None, files: Optional[list[str]]=None, + r_start: list[int]=[], r_stop: list[int]=[], r_step:int=1) -> dict[str, Any]: """ Extract t0 measurements from the project. @@ -206,15 +204,11 @@ def extract_t0(path: Path, project: str, dir_in_project: str, ensemble: str, par Dictionary of t0 values in the pycorrlib style, with the parameters at hand. """ - if r_stop is None: - r_stop = [] - if r_start is None: - r_start = [] dataset = os.path.join(path, "projects", project) directory = os.path.join(dataset, dir_in_project) if files is None: files = [] - for _root, _ds, fs in os.walk(directory): + for root, ds, fs in os.walk(directory): for f in fs: if fnmatch.fnmatch(f, prefix + "*" + postfix + ".dat"): files.append(f) @@ -256,8 +250,8 @@ def extract_t0(path: Path, project: str, dir_in_project: str, ensemble: str, par return t0_dict -def extract_t1(path: Path, project: str, dir_in_project: str, ensemble: str, param: dict[str, Any], prefix: str, dtr_read: int, xmin: int, spatial_extent: int, fit_range: int = 5, postfix: str = "", names: list[str] | None=None, files: list[str] | None=None, - r_start: list[int] | None=None, r_stop: list[int] | None=None, r_step:int=1) -> dict[str, Any]: +def extract_t1(path: Path, project: str, dir_in_project: str, ensemble: str, param: dict[str, Any], prefix: str, dtr_read: int, xmin: int, spatial_extent: int, fit_range: int = 5, postfix: str = "", names: Optional[list[str]]=None, files: Optional[list[str]]=None, + r_start: list[int]=[], r_stop: list[int]=[], r_step:int=1) -> dict[str, Any]: """ Extract t1 measurements from the project. @@ -294,14 +288,10 @@ def extract_t1(path: Path, project: str, dir_in_project: str, ensemble: str, par Dictionary of t1 values in the pycorrlib style, with the parameters at hand. """ - if r_stop is None: - r_stop = [] - if r_start is None: - r_start = [] directory = os.path.join(path, "projects", project, dir_in_project) if files is None: files = [] - for _root, _ds, fs in os.walk(directory): + for root, ds, fs in os.walk(directory): for f in fs: if fnmatch.fnmatch(f, prefix + "*" + postfix + ".dat"): files.append(f) diff --git a/corrlib/input/sfcf.py b/corrlib/input/sfcf.py index a12661d..acd8261 100644 --- a/corrlib/input/sfcf.py +++ b/corrlib/input/sfcf.py @@ -1,11 +1,11 @@ +import pyerrors as pe +import datalad.api as dl import json import os +from typing import Any from fnmatch import fnmatch from pathlib import Path -from typing import Any -import datalad.api as dl -import pyerrors as pe bi_corrs: list[str] = ["f_P", "fP", "f_p", "g_P", "gP", "g_p", @@ -99,7 +99,7 @@ def read_param(path: Path, project: str, file_in_project: str) -> dict[str, Any] file = path / "projects" / project / file_in_project dl.get(file, dataset=path) - with open(file) as f: + with open(file, 'r') as f: lines = f.readlines() params: dict[str, Any] = {} @@ -258,7 +258,7 @@ def get_specs(key: str, parameters: dict[str, Any], sep: str = '/') -> str: return s -def read_data(path: Path, project: str, dir_in_project: str, prefix: str, param: dict[str, Any], version: str = '1.0c', cfg_separator: str = 'n', sep: str = '/', **kwargs: Any) -> dict[str, Any]: +def read_data(path: Path, project: str, dir_in_project: str, prefix: str, param: dict[str, Any], version: str = '1.0c', cfg_seperator: str = 'n', sep: str = '/', **kwargs: Any) -> dict[str, Any]: """ Extract the data from the sfcf file. @@ -274,7 +274,7 @@ def read_data(path: Path, project: str, dir_in_project: str, prefix: str, param: The parameter dictionary, as given by read_param. version: str Version of sfcf. - cfg_separator: str + cfg_seperator: str Separator of the configuration number. Needed for reading. default: "n" sep: str Seperator for the key in return dict. (default: "/) @@ -291,7 +291,7 @@ def read_data(path: Path, project: str, dir_in_project: str, prefix: str, param: appended = (version[-1] == "a") ls = [] files_to_get = [] - for _dirpath, dirnames, filenames in os.walk(directory): + for (dirpath, dirnames, filenames) in os.walk(directory): if not appended: ls.extend(dirnames) else: @@ -299,7 +299,7 @@ def read_data(path: Path, project: str, dir_in_project: str, prefix: str, param: break if not appended: compact = (version[-1] == "c") - for item in ls: + for i, item in enumerate(ls): if fnmatch(item, prefix + "*"): rep_path = directory + '/' + item sub_ls = pe.input.sfcf._find_files(rep_path, prefix, compact, []) @@ -321,10 +321,10 @@ def read_data(path: Path, project: str, dir_in_project: str, prefix: str, param: if not param['crr'] == []: if names is not None: data_crr = pe.input.sfcf.read_sfcf_multi(directory, prefix, param['crr'], param['mrr'], corr_type_list, range(len(param['wf_offsets'])), - range(len(param['wf_basis'])), range(len(param['wf_basis'])), version, cfg_separator, keyed_out=True, silent=True, names=names) + range(len(param['wf_basis'])), range(len(param['wf_basis'])), version, cfg_seperator, keyed_out=True, silent=True, names=names) else: data_crr = pe.input.sfcf.read_sfcf_multi(directory, prefix, param['crr'], param['mrr'], corr_type_list, range(len(param['wf_offsets'])), - range(len(param['wf_basis'])), range(len(param['wf_basis'])), version, cfg_separator, keyed_out=True, silent=True) + range(len(param['wf_basis'])), range(len(param['wf_basis'])), version, cfg_seperator, keyed_out=True, silent=True) for key in data_crr.keys(): data[key] = data_crr[key] diff --git a/corrlib/integrity.py b/corrlib/integrity.py index abffdb4..66ea4df 100644 --- a/corrlib/integrity.py +++ b/corrlib/integrity.py @@ -1,15 +1,14 @@ import datetime as dt -import os -import sqlite3 -from configparser import ConfigParser from pathlib import Path +from .tools import get_db_file, CONFIG_FILENAME +import pandas as pd +import sqlite3 +from .tracker import get +import pyerrors.input.json as pj +import os +from configparser import ConfigParser from typing import Any -import pandas as pd -import pyerrors.input.json as pj - -from .tools import CONFIG_FILENAME, get_db_file -from .tracker import get path_opts = ['db', 'projects_path', 'archive_path', 'toml_imports_path', 'import_scripts_path'] diff --git a/corrlib/main.py b/corrlib/main.py index a0806cf..5df8165 100644 --- a/corrlib/main.py +++ b/corrlib/main.py @@ -1,18 +1,17 @@ -import os -import shutil import sqlite3 -from pathlib import Path - import datalad.api as dl import datalad.config as dlc - -from .find import _project_lookup_by_id +import os from .git_tools import move_submodule -from .tools import get_db_file, list2str, str2list -from .tracker import clone, drop, get, save, unlock +import shutil +from .find import _project_lookup_by_id +from .tools import list2str, str2list, get_db_file +from .tracker import get, save, unlock, clone, drop +from typing import Union, Optional +from pathlib import Path -def create_project(path: Path, uuid: str, owner: str | None=None, tags: list[str] | None=None, aliases: list[str] | None=None, code: str | None=None) -> None: +def create_project(path: Path, uuid: str, owner: Union[str, None]=None, tags: Union[list[str], None]=None, aliases: Union[list[str], None]=None, code: Union[str, None]=None) -> None: """ Create a new project entry in the database. @@ -50,7 +49,7 @@ def create_project(path: Path, uuid: str, owner: str | None=None, tags: list[str return -def update_project_data(path: Path, uuid: str, prop: str, value: str | None = None) -> None: +def update_project_data(path: Path, uuid: str, prop: str, value: Union[str, None] = None) -> None: """ Update/Edit a project entry in the database. Thin wrapper around sql3 call. @@ -103,7 +102,7 @@ def update_aliases(path: Path, uuid: str, aliases: list[str]) -> None: return -def import_project(path: Path, url: str, owner: str | None=None, tags: list[str] | None=None, aliases: list[str] | None=None, code: str | None=None, isDataset: bool=True) -> str: +def import_project(path: Path, url: str, owner: Union[str, None]=None, tags: Optional[list[str]]=None, aliases: Optional[list[str]]=None, code: Optional[str]=None, isDataset: bool=True) -> str: """ Import a datalad dataset into the backlogger. diff --git a/corrlib/meas_io.py b/corrlib/meas_io.py index b5c4029..52e307e 100644 --- a/corrlib/meas_io.py +++ b/corrlib/meas_io.py @@ -1,23 +1,23 @@ -import json -import os -import shutil -import sqlite3 -from hashlib import sha256 -from pathlib import Path -from typing import Any - -from pyerrors import Corr, Obs, dump_object, load_object from pyerrors.input import json as pj - -from .input import openQCD, sfcf -from .integrity import _check_db2paths -from .tools import cache_enabled, get_db_file, get_plot_dir +import os +import sqlite3 +from .input import sfcf,openQCD +import json +from typing import Union +from pyerrors import Obs, Corr, dump_object, load_object +from hashlib import sha256 +from .tools import get_db_file, cache_enabled, get_plot_dir from .tracker import get, save, unlock +import shutil +from typing import Any +from pathlib import Path +from .integrity import _check_db2paths + CACHE_DIR = ".cache" -def write_measurement(path: Path, ensemble: str, measurement: dict[str, dict[str, dict[str, Any]]], uuid: str, code: str, parameter_file: str | None, final_write: dict[str, bool]) -> None: +def write_measurement(path: Path, ensemble: str, measurement: dict[str, dict[str, dict[str, Any]]], uuid: str, code: str, parameter_file: Union[str, None]) -> None: """ Write a measurement to the backlog. If the file for the measurement already exists, update the measurement. @@ -36,8 +36,6 @@ def write_measurement(path: Path, ensemble: str, measurement: dict[str, dict[str Name of the code that was used for the project. parameter_file: str The parameter file used for the measurement. - final_write: bool - Determmines whether this is the final ime the file is touched during the current import. """ path = Path(path) db_file = get_db_file(path) @@ -54,16 +52,12 @@ def write_measurement(path: Path, ensemble: str, measurement: dict[str, dict[str for corr in measurement.keys(): file_in_archive = Path('.') / 'archive' / ensemble / corr / str(uuid + '.json.gz') file = Path(path) / file_in_archive - tmp_file_in_archive = Path('.') / 'archive' / ensemble / corr / (str(uuid) + ".p") - tmp_file = Path(path) / tmp_file_in_archive - known_meas: dict[str, Any] = {} + known_meas = {} if not os.path.exists(path / 'archive' / ensemble / corr): os.makedirs(path / 'archive' / ensemble / corr) files_to_save.append(file_in_archive) else: - if os.path.exists(tmp_file): - known_meas = load_object(str(tmp_file)) - elif os.path.exists(file): + if os.path.exists(file): if file not in files_to_save: unlock(path, file_in_archive) files_to_save.append(file_in_archive) @@ -79,7 +73,7 @@ def write_measurement(path: Path, ensemble: str, measurement: dict[str, dict[str pars[subkey] = sfcf.get_specs(corr + "/" + subkey, parameters) elif code == "openQCD": - ms_type = next(iter(measurement.keys())) + ms_type = list(measurement.keys())[0] if ms_type == 'ms1': if parameter_file is not None: if parameter_file.endswith(".ms1.in"): @@ -138,27 +132,13 @@ def write_measurement(path: Path, ensemble: str, measurement: dict[str, dict[str c.execute("INSERT INTO backlogs (name, ensemble, code, path, project, parameters, parameter_file, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?, datetime('now'), datetime('now'))", (corr, ensemble, code, meas_path, uuid, pars[subkey], parameter_file)) conn.commit() - if final_write[str(file)]: - pj.dump_dict_to_json(known_meas, str(file)) - if os.path.exists(tmp_file): - os.remove(tmp_file) - else: - dump_object(known_meas, str(tmp_file)[:-2]) + pj.dump_dict_to_json(known_meas, str(file)) conn.close() save(path, message="Add measurements to database", files=files_to_save) return -def affected_files(corrs: list[str], ensemble: str, uuid: str) -> list[Path]: - file_list = [] - for corr in corrs: - file_in_archive = Path('.') / 'archive' / ensemble / corr / str(uuid + '.json.gz') - file_list.append(file_in_archive) - file_list = list(set(file_list)) - return file_list - - -def load_record(path: Path, meas_path: str) -> Corr | Obs: +def load_record(path: Path, meas_path: str) -> Union[Corr, Obs]: """ Load a list of records by their paths. @@ -177,7 +157,7 @@ def load_record(path: Path, meas_path: str) -> Corr | Obs: return load_records(path, [meas_path])[0] -def load_records(path: Path, meas_paths: list[str], preloaded: dict[str, Any] | None = None, dry_run: bool = False) -> list[Corr | Obs]: +def load_records(path: Path, meas_paths: list[str], preloaded: dict[str, Any] = {}, dry_run: bool = False) -> list[Union[Corr, Obs]]: """ Load a list of records by their paths. @@ -197,8 +177,6 @@ def load_records(path: Path, meas_paths: list[str], preloaded: dict[str, Any] | returned_data: list The loaded records. """ - if preloaded is None: - preloaded = {} path = Path(path) if dry_run: _check_db2paths(path, meas_paths) diff --git a/corrlib/pars/openQCD/flags.py b/corrlib/pars/openQCD/flags.py index 0ab429a..95be919 100644 --- a/corrlib/pars/openQCD/flags.py +++ b/corrlib/pars/openQCD/flags.py @@ -5,7 +5,6 @@ Reconstruct the outputs of flags. import struct from typing import Any, BinaryIO - # lat_parms.c def lat_parms_write_lat_parms(fp: BinaryIO) -> dict[str, Any]: """ @@ -30,7 +29,7 @@ def lat_parms_write_lat_parms(fp: BinaryIO) -> dict[str, Any]: kappas = [] m0s = [] # read kappas - for _ik in range(nk): + for ik in range(nk): t = fp.read(8) kappas.append(struct.unpack('d', t)[0]) t = fp.read(8) diff --git a/corrlib/pars/openQCD/ms1.py b/corrlib/pars/openQCD/ms1.py index b9b3fc1..4c2aed5 100644 --- a/corrlib/pars/openQCD/ms1.py +++ b/corrlib/pars/openQCD/ms1.py @@ -1,8 +1,8 @@ -from pathlib import Path -from typing import Any - from . import flags +from typing import Any +from pathlib import Path + def read_qcd2_ms1_par_file(fname: Path) -> dict[str, dict[str, Any]]: """ diff --git a/corrlib/pars/openQCD/qcd2.py b/corrlib/pars/openQCD/qcd2.py index 121c25f..e73c156 100644 --- a/corrlib/pars/openQCD/qcd2.py +++ b/corrlib/pars/openQCD/qcd2.py @@ -1,8 +1,8 @@ +from . import flags + from pathlib import Path from typing import Any -from . import flags - def read_qcd2_par_file(fname: Path) -> dict[str, dict[str, Any]]: """ diff --git a/corrlib/sql.py b/corrlib/sql.py index fd8e786..f45ce31 100644 --- a/corrlib/sql.py +++ b/corrlib/sql.py @@ -1,9 +1,8 @@ import sqlite3 +from .tools import get_db_file from pathlib import Path from typing import Any -from .tools import get_db_file - def thin_sql_wrapper(path: Path, stmt: str) -> list[Any]: db_file = get_db_file(path) diff --git a/corrlib/toml.py b/corrlib/toml.py index 7becb66..76533f4 100644 --- a/corrlib/toml.py +++ b/corrlib/toml.py @@ -8,20 +8,18 @@ the import of projects via TOML. """ -import os +import tomllib as toml import shutil -from pathlib import Path -from typing import Any import datalad.api as dl -import tomllib as toml - -from .input import openQCD, sfcf -from .input.implementations import codes as known_codes -from .main import import_project, update_aliases -from .meas_io import affected_files, write_measurement -from .tools import step_differences from .tracker import save +from .input import sfcf, openQCD +from .main import import_project, update_aliases +from .meas_io import write_measurement +import os +from .input.implementations import codes as known_codes +from typing import Any +from pathlib import Path def replace_string(string: str, name: str, val: str) -> str: @@ -118,7 +116,7 @@ def check_measurement_data(measurements: dict[str, dict[str, str]], code: str) - """ var_names: list[str] = [] if code == "sfcf": - var_names = ["path", "ensemble", "param_file", "version", "prefix", "cfg_separator", "names"] + var_names = ["path", "ensemble", "param_file", "version", "prefix", "cfg_seperator", "names"] elif code == "openQCD": var_names = ["path", "ensemble", "measurement", "prefix"] # , "param_file" for mname, md in measurements.items(): @@ -186,47 +184,19 @@ def import_toml(path: Path, file: str, copy_file: bool=True) -> None: uuid = import_project(path, project['url'], aliases=aliases) imeas = 1 nmeas = len(measurements.keys()) - - # preparation step - affected_file_d = {} - mname_list = list(measurements.keys()) - for mname in mname_list: - md = measurements[mname] - ensemble = md['ensemble'] - if project['code'] == 'sfcf': - param = sfcf.read_param(path, uuid, md['param_file']) - affected_by_meas = affected_files(param['crr'], ensemble, uuid) - elif project['code'] == 'openQCD': - if md['measurement'] == 'ms1': - affected_by_meas = affected_files(param['type'], ensemble, uuid) - elif md['measurement'] == 't0': - affected_by_meas = affected_files(param['type'], ensemble, uuid) - elif md['measurement'] == 't1': - affected_by_meas = affected_files(param['type'], ensemble, uuid) - affected_file_d[mname] = [str(path / f) for f in affected_by_meas] - future_affected_file_d = {} - for i,mname in enumerate(mname_list): - future_affected_file_d[mname] = [] - for mname2 in mname_list[i+1:]: - future_affected_file_d[mname].extend(affected_file_d[mname2]) - for mname in mname_list: - md = measurements[mname] + for mname, md in measurements.items(): print(f"Import measurement {imeas}/{nmeas}: {mname}") ensemble = md['ensemble'] if project['code'] == 'sfcf': param = sfcf.read_param(path, uuid, md['param_file']) if 'names' in md.keys(): measurement = sfcf.read_data(path, uuid, md['path'], md['prefix'], param, - version=md['version'], cfg_separator=md['cfg_separator'], sep='/', names=md['names']) + version=md['version'], cfg_seperator=md['cfg_seperator'], sep='/', names=md['names']) else: measurement = sfcf.read_data(path, uuid, md['path'], md['prefix'], param, - version=md['version'], cfg_separator=md['cfg_separator'], sep='/') + version=md['version'], cfg_seperator=md['cfg_seperator'], sep='/') elif project['code'] == 'openQCD': - if not (isinstance(md['files'], list)): - raise ValueError("files has to be a list of strings") - if not all(isinstance(f, str) for f in md["files"]): - raise ValueError("files has to be a list of strings") if md['measurement'] == 'ms1': if 'param_file' in md.keys(): parameter_file = md['param_file'] @@ -269,12 +239,7 @@ def import_toml(path: Path, file: str, copy_file: bool=True) -> None: measurement = openQCD.extract_t1(path, uuid, md['path'], ensemble, param, str(md["prefix"]), int(md["dtr_read"]), int(md["xmin"]), int(md["spatial_extent"]), fit_range=int(md.get('fit_range', 5)), postfix=str(md.get('postfix', '')), names=md.get('names', []), files=md.get('files', []), r_start=md.get('r_start', []), r_stop=md.get('r_stop', []), r_step=md.get('r_step', 1)) - final_write = {} - for file in affected_file_d[mname]: - final_write[str(file)] = True - if str(file) in future_affected_file_d[mname]: - final_write[str(file)] = False - write_measurement(path, ensemble, measurement, uuid, project['code'], (md['param_file'] if 'param_file' in md else None), final_write) + write_measurement(path, ensemble, measurement, uuid, project['code'], (md['param_file'] if 'param_file' in md else None)) imeas += 1 print(mname + " imported.") @@ -301,7 +266,7 @@ def reimport_project(path: Path, uuid: str) -> None: uuid of the project that is to be reimported. """ config_path = path / "import_scripts" / uuid - for _p, filenames, _dirnames in os.walk(config_path): + for p, filenames, dirnames in os.walk(config_path): for fname in filenames: import_toml(path, os.path.join(config_path, fname), copy_file=False) return diff --git a/corrlib/tools.py b/corrlib/tools.py index 12cf2ac..8dec61a 100644 --- a/corrlib/tools.py +++ b/corrlib/tools.py @@ -1,7 +1,7 @@ import os from configparser import ConfigParser -from pathlib import Path from typing import Any +from pathlib import Path CONFIG_FILENAME = ".corrlib" cached: bool = True @@ -183,22 +183,3 @@ def cache_enabled(path: Path) -> bool: raise ValueError(f"String {cached_str} is not a valid option, only True and False are allowed!") cached_bool = cached_str == ('True') return cached_bool - - -def step_differences(name_list: list[Any], dict_of_lists: dict[Any, Any]) -> list[set[Any]]: - needed_until_step = [] - for i in range(len(name_list)): - nf: set[Any] = set() - for k in range(i, len(name_list)): - nf = nf.union(dict_of_lists[name_list[k]]) - needed_until_step.append(nf) - - discard_after = [] - for i in range(len(needed_until_step)-1): - discard_after.append(needed_until_step[i].difference(needed_until_step[i+1])) - discard_after.append(needed_until_step[-1]) - - print(discard_after) - if not set(dict_of_lists[name_list[-1]]) == discard_after[-1]: - raise ValueError("Discards and last items diverge.") - return discard_after diff --git a/corrlib/tracker.py b/corrlib/tracker.py index 0a962fa..6f4ae3d 100644 --- a/corrlib/tracker.py +++ b/corrlib/tracker.py @@ -1,12 +1,10 @@ import os -import shutil -import warnings from configparser import ConfigParser -from pathlib import Path - import datalad.api as dl - -from .tools import CONFIG_FILENAME, get_db_file +from typing import Optional +import shutil +from .tools import get_db_file, CONFIG_FILENAME +from pathlib import Path def get_tracker(path: Path) -> str: @@ -61,7 +59,7 @@ def get(path: Path, file: Path) -> None: return -def save(path: Path, message: str, files: list[Path] | None=None) -> None: +def save(path: Path, message: str, files: Optional[list[Path]]=None) -> None: """ Wrapper function to save a file to the dataset located at path with the specified tracker. @@ -81,7 +79,8 @@ def save(path: Path, message: str, files: list[Path] | None=None) -> None: files = [path / f for f in files] dl.save(files, message=message, dataset=path) elif tracker == 'None': - warnings.warn("Tracker 'None' does not implement save.", Warning, 1) + Warning("Tracker 'None' does not implement save.") + pass else: raise ValueError(f"Tracker {tracker} is not supported.") @@ -123,7 +122,8 @@ def unlock(path: Path, file: Path) -> None: if tracker == 'datalad': dl.unlock(os.path.join(path, file), dataset=path) elif tracker == 'None': - warnings.warn("Tracker 'None' does not implement unlock.", Warning, 1) + Warning("Tracker 'None' does not implement unlock.") + pass else: raise ValueError(f"Tracker {tracker} is not supported.") return @@ -144,7 +144,7 @@ def clone(path: Path, source: str, target: str) -> None: path = Path(path) tracker = get_tracker(path) if tracker == 'datalad': - dl.clone(path=target, source=source, dataset=path) + dl.clone(target=target, source=source, dataset=path) elif tracker == 'None': os.makedirs(path, exist_ok=True) # Implement a simple clone by copying files @@ -154,7 +154,7 @@ def clone(path: Path, source: str, target: str) -> None: return -def drop(path: Path, reckless: str | None=None) -> None: +def drop(path: Path, reckless: Optional[str]=None) -> None: """ Wrapper function to drop data from a dataset located at path with the specified tracker. @@ -170,7 +170,8 @@ def drop(path: Path, reckless: str | None=None) -> None: if tracker == 'datalad': dl.drop(path, reckless=reckless) elif tracker == 'None': - warnings.warn("Tracker 'None' does not implement drop.", Warning, 1) + Warning("Tracker 'None' does not implement drop.") + pass else: raise ValueError(f"Tracker {tracker} is not supported.") return diff --git a/corrlib/version.py b/corrlib/version.py index 3e22db8..23dd03f 100644 --- a/corrlib/version.py +++ b/corrlib/version.py @@ -18,7 +18,7 @@ version_tuple: tuple[int | str, ...] commit_id: str | None __commit_id__: str | None -__version__ = version = '0.3.1.dev32+g906a2bdf3.d20260710' -__version_tuple__ = version_tuple = (0, 3, 1, 'dev32', 'g906a2bdf3.d20260710') +__version__ = version = '0.3.1.dev0+g08de17e6b.d20260507' +__version_tuple__ = version_tuple = (0, 3, 1, 'dev0', 'g08de17e6b.d20260507') -__commit_id__ = commit_id = 'g906a2bdf3' +__commit_id__ = commit_id = 'g08de17e6b' diff --git a/pyproject.toml b/pyproject.toml index 10756b7..68abbe8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -27,17 +27,13 @@ include = ["corrlib", "corrlib.*"] [tool.setuptools_scm] write_to = "corrlib/version.py" -[tool.ruff] -target-version = "py310" - [tool.ruff.lint] -extend-select = ["E", "W", "I", "B", "PIE", "PLE", "PLW", "UP", "NPY", "RUF"] -ignore = [ - "F403", # star imports in __init__ files are intentional - "E501", # line too long - "PLC0415", # import outside top level - "PLW2901", # redefined loop name (too noisy) - "RUF002", # ambiguous unicode in docstrings (Greek letters) +ignore = ["E501"] +extend-select = [ + "YTT", + "E", + "W", + "F", ] [tool.mypy] diff --git a/tests/find_test.py b/tests/find_test.py index b462634..2144001 100644 --- a/tests/find_test.py +++ b/tests/find_test.py @@ -306,7 +306,8 @@ def test_openQCD_filter() -> None: "updated_at"] df = pd.DataFrame(data,columns=cols) - find.openQCD_filter(df, a = "asdf") + with pytest.warns(Warning): + find.openQCD_filter(df, a = "asdf") def test_code_filter() -> None: diff --git a/tests/import_project_test.py b/tests/import_project_test.py index 8493773..685d2cf 100644 --- a/tests/import_project_test.py +++ b/tests/import_project_test.py @@ -10,7 +10,7 @@ def test_toml_check_measurement_data() -> None: "param_file": "/path/to/file", "version": "1.1", "prefix": "pref", - "cfg_separator": "n", + "cfg_seperator": "n", "names": ['list', 'of', 'names'] } }