From 06b6e1e7c34abe51a8aebef1e38aeecb5f084938 Mon Sep 17 00:00:00 2001 From: Michael Granitzer Date: Fri, 21 Jun 2024 15:02:26 +0200 Subject: [PATCH] fix: added typecasting for parameter calls --- owilix/cmd/base.py | 80 +++++++++++++++++++++++++++++++++++++++-- owilix/cmd/query.py | 4 +-- owilix/core/metadata.py | 1 + pyproject.toml | 1 + 4 files changed, 81 insertions(+), 5 deletions(-) diff --git a/owilix/cmd/base.py b/owilix/cmd/base.py index 52a1f3c..a915032 100644 --- a/owilix/cmd/base.py +++ b/owilix/cmd/base.py @@ -5,7 +5,10 @@ import re from collections import Counter from logging.handlers import RotatingFileHandler from statistics import stdev -from typing import Iterable +from typing import Iterable, Tuple, Dict +import inspect +from pydantic import BaseModel, create_model, ValidationError, parse_obj_as, TypeAdapter +from typing import Any, Optional, get_origin, get_args, Union import fsspec from rich.columns import Columns @@ -150,6 +153,7 @@ class CommandResult: + class BaseCommand(metaclass=SubCommandMeta): """ Base class for command line interfaces. @@ -214,16 +218,86 @@ class BaseCommand(metaclass=SubCommandMeta): """ Handle calls to the OWILocal instance as command invocations. """ return self.do(cmd, *args, **kwargs) + def _cast_args(self, func, args: Tuple[Any, ...], kwargs: Dict[str, Any]) -> Tuple[Tuple[Any, ...], Dict[str, Any]]: + """ + Converts args and kwargs to match the types specified in the function's signature using pydantic. + + Parameters: + func (callable): The function whose signature will be used for type conversion. + args (tuple): The positional arguments to convert. + kwargs (dict): The keyword arguments to convert. + + Returns: + tuple: A tuple containing the converted args and kwargs. + """ + sig = inspect.signature(func) + parameters = list(sig.parameters.values()) + + converted_args = [] + converted_kwargs = {} + + # Convert positional arguments (*args) that match the function signature + for i, arg in enumerate(args): + if i < len(parameters): + param = parameters[i] + expected_type = param.annotation + + if expected_type == inspect.Parameter.empty: + converted_args.append(arg) + else: + try: + # Use TypeAdapter to convert the argument + type_adapter = TypeAdapter(expected_type) + converted_args.append(type_adapter.validate_python(arg)) + except (ValueError, TypeError) as e: + print(f"WARNING - Could not convert arg[{i}]='{arg}' to {expected_type}: {e}") + converted_args.append(arg) + + # Include remaining *args as-is + if len(args) > len(parameters): + converted_args.extend(args[len(parameters):]) + + # Convert keyword arguments (*kwargs) that match the function signature + for name, param in sig.parameters.items(): + if param.kind in (inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.POSITIONAL_OR_KEYWORD): + if name in kwargs: + expected_type = param.annotation + value = kwargs[name] + if expected_type != inspect.Parameter.empty: + try: + type_adapter = TypeAdapter(expected_type) + converted_kwargs[name] = type_adapter.validate_python(value) + except (ValueError, TypeError) as e: + print(f"WARNING - Could not convert kwarg '{name}'='{value}' to {expected_type}: {e}") + converted_kwargs[name] = value + else: + converted_kwargs[name] = value + + # Include any additional kwargs that weren't in the function signature + for k, v in kwargs.items(): + if k not in converted_kwargs: + converted_kwargs[k] = v + + # Ensure no positional argument conflicts with keyword arguments + for i, arg in enumerate(converted_args): + if i < len(parameters): + param_name = parameters[i].name + if param_name in converted_kwargs: + raise TypeError(f"Got multiple values for argument '{param_name}'") + + return tuple(converted_args), converted_kwargs + def do(self, cmd, *args, **kwargs): """ Dispatch to the appropriate sub-command. """ if cmd in self.commands: _param = {"group": str(self.__class__.__name__)} |get_func_params(self.commands[cmd], *args, **kwargs) try: _param["success"] = False + args, converted_kwargs = self._cast_args(self.commands[cmd], args, kwargs) if cmd =="help": - self.console.print(self.commands[cmd](*args, **kwargs)) + self.console.print(self.commands[cmd](*args, **converted_kwargs)) return - returns = self.commands[cmd](*args, **kwargs) + returns = self.commands[cmd](*args, **converted_kwargs) _param["success"] = returns.success if isinstance(returns, CommandResult) else True _param["msg"] = returns.msg if isinstance(returns, CommandResult) else "" except Exception as e: diff --git a/owilix/cmd/query.py b/owilix/cmd/query.py index bc58805..5416778 100644 --- a/owilix/cmd/query.py +++ b/owilix/cmd/query.py @@ -140,7 +140,7 @@ def slice(self, local_specifier, remote_specifier, batch_size=500, prefetch=3, partitioned_by="", access="public", overwrite_ignore=True, import_collection="userslice", chunk_size=1000000, internalID=None, - **kwargs ): + **kwargs): """ Executes the query over the datasets selected by specified local and remote specifiers and applies select and where clause provided in kwargs @@ -246,7 +246,7 @@ def slice(self, local_specifier, remote_specifier, _changelog.append(str(results)) except Exception as e: - self.console.exception(e) + self.console.print_exception() self.console.log(f"Dataset with id {_ds.internalID} at {_ds.path} is most likely corrupted and should be removed.") finally: _ds.append_changelog("Slice Result: "+",".join(_changelog), True) diff --git a/owilix/core/metadata.py b/owilix/core/metadata.py index a68e11a..ae3f2b6 100644 --- a/owilix/core/metadata.py +++ b/owilix/core/metadata.py @@ -574,6 +574,7 @@ class Dataset: _provenance_new = set([f"{create_provenance_url(d, files, select=select, where=where)}" for d in datasets]) _provenance_old = set(self.metadata["provenance"]) _provenance = list(_provenance_old.union(_provenance_new)) + self.metadata["provenance"] = _provenance _overlapping = set([parse_provenance_url(u).get("internalID", None) for u in _provenance_old.intersection(_provenance_new)]) diff --git a/pyproject.toml b/pyproject.toml index e24e42f..a49f0da 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -35,6 +35,7 @@ click= "8.1.6" irods-fsspec="^0.0.1" duckdb = {version = "^1.0.0"} ipython = {version = "^8.24.0"} +pydantic = "^2.5.3" rich = "^13.7.0" numpy = "^1.26.4" -- 2.51.2