diff --git a/fire/core.py b/fire/core.py index 6367262d..55d8a58e 100644 --- a/fire/core.py +++ b/fire/core.py @@ -78,7 +78,7 @@ def main(argv): import asyncio # pylint: disable=import-error,g-import-not-at-top # pytype: disable=import-error -def Fire(component=None, command=None, name=None, serialize=None): +def Fire(component=None, command=None, name=None, serialize=None, context=None): """This function, Fire, is the main entrypoint for Python Fire. Executes a command either from the `command` argument or from sys.argv by @@ -130,8 +130,9 @@ def Fire(component=None, command=None, name=None, serialize=None): argparser = parser.CreateParser() parsed_flag_args, unused_args = argparser.parse_known_args(flag_args) - context = {} - if parsed_flag_args.interactive or component is None: + if context is None: + context = {} + if parsed_flag_args.interactive or not context: # Determine the calling context. caller = inspect.stack()[1] caller_frame = caller[0] @@ -436,6 +437,7 @@ def _Fire(component, args, parsed_flag_args, context, name=None): instance = None remaining_args = args + variable_env = dict() while True: last_component = component initial_args = remaining_args @@ -452,14 +454,19 @@ def _Fire(component, args, parsed_flag_args, context, name=None): saved_args = [] used_separator = False - if separator in remaining_args: - # For the current component, only use arguments up to the separator. - separator_index = remaining_args.index(separator) - saved_args = remaining_args[separator_index + 1:] + separator_index = -1 + for ai, ra in enumerate(remaining_args): + if ra[0] == "@" or ra == "-": + separator_index = ai + break + if separator_index != -1: + if remaining_args[separator_index][0] == "@": + saved_args = remaining_args[separator_index:] + else: + saved_args = remaining_args[separator_index+1:] remaining_args = remaining_args[:separator_index] used_separator = True assert separator not in remaining_args - handled = False candidate_errors = [] @@ -468,6 +475,19 @@ def _Fire(component, args, parsed_flag_args, context, name=None): is_sequence = isinstance(component, (list, tuple)) is_map = isinstance(component, dict) or inspectutils.IsNamedTuple(component) + if not (is_callable or is_callable_object) and len(remaining_args)==0 and len(saved_args) > 0 and isinstance(saved_args[0], str): + if saved_args[0] == "@": + variable_env["_"] = component + saved_args.pop(0) + component = context.copy() + component.update(variable_env) + elif saved_args[0][0] == "@": + variable_env[saved_args[0][1:]] = component + saved_args.pop(0) + component = context.copy() + component.update(variable_env) + + if not handled and is_callable: # The component is a class or a routine; we'll try to initialize it or # call it. @@ -478,6 +498,7 @@ def _Fire(component, args, parsed_flag_args, context, name=None): component, remaining_args, component_trace, + variable_env, treatment='class' if is_class else 'routine', target=component.__name__) handled = True @@ -569,6 +590,7 @@ def _Fire(component, args, parsed_flag_args, context, name=None): component, remaining_args, component_trace, + variable_env, treatment='callable') handled = True except FireError as error: @@ -611,6 +633,7 @@ def _Fire(component, args, parsed_flag_args, context, name=None): if interactive: variables = context.copy() + variables.update(variable_env) if name is not None: variables[name] = initial_component @@ -658,7 +681,7 @@ def _GetMember(component, args): raise FireError('Could not consume arg:', arg) -def _CallAndUpdateTrace(component, args, component_trace, treatment='class', +def _CallAndUpdateTrace(component, args, component_trace, variable_env, treatment='class', target=None): """Call the component by consuming args from args, and update the FireTrace. @@ -682,7 +705,7 @@ def _CallAndUpdateTrace(component, args, component_trace, treatment='class', filename, lineno = inspectutils.GetFileAndLine(component) metadata = decorators.GetMetadata(component) fn = component.__call__ if treatment == 'callable' else component - parse = _MakeParseFn(fn, metadata) + parse = _MakeParseFn(fn, metadata, variable_env) (varargs, kwargs), consumed_args, remaining_args, capacity = parse(args) # Call the function. @@ -705,7 +728,7 @@ def _CallAndUpdateTrace(component, args, component_trace, treatment='class', return component, remaining_args -def _MakeParseFn(fn, metadata): +def _MakeParseFn(fn, metadata, variable_env): """Creates a parse function for fn. Args: @@ -731,7 +754,7 @@ def _ParseFn(args): # Note: _ParseArgs modifies kwargs. parsed_args, kwargs, remaining_args, capacity = _ParseArgs( fn_spec.args, fn_spec.defaults, num_required_args, kwargs, - remaining_args, metadata) + remaining_args, metadata, variable_env) if fn_spec.varargs or fn_spec.varkw: # If we're allowed *varargs or **kwargs, there's always capacity. @@ -752,7 +775,7 @@ def _ParseFn(args): varargs = [] for index, value in enumerate(varargs): - varargs[index] = _ParseValue(value, None, None, metadata) + varargs[index] = _ParseValue(value, None, None, metadata, variable_env) varargs = parsed_args + varargs remaining_args += remaining_kwargs @@ -764,7 +787,7 @@ def _ParseFn(args): def _ParseArgs(fn_args, fn_defaults, num_required_args, kwargs, - remaining_args, metadata): + remaining_args, metadata, variable_env): """Parses the positional and named arguments from the available supplied args. Modifies kwargs, removing args as they are used. @@ -798,13 +821,13 @@ def _ParseArgs(fn_args, fn_defaults, num_required_args, kwargs, for index, arg in enumerate(fn_args): value = kwargs.pop(arg, None) if value is not None: # A value is specified at the command line. - value = _ParseValue(value, index, arg, metadata) + value = _ParseValue(value, index, arg, metadata, variable_env) parsed_args.append(value) else: # No value has been explicitly specified. if remaining_args and accepts_positional_args: # Use a positional arg. value = remaining_args.pop(0) - value = _ParseValue(value, index, arg, metadata) + value = _ParseValue(value, index, arg, metadata, variable_env) parsed_args.append(value) elif index < num_required_args: raise FireError( @@ -817,7 +840,7 @@ def _ParseArgs(fn_args, fn_defaults, num_required_args, kwargs, parsed_args.append(fn_defaults[default_index]) for key, value in kwargs.items(): - kwargs[key] = _ParseValue(value, None, key, metadata) + kwargs[key] = _ParseValue(value, None, key, metadata, variable_env) return parsed_args, kwargs, remaining_args, capacity @@ -862,7 +885,7 @@ def _ParseKeywordArgs(args, fn_spec): skip_argument = False continue - if _IsFlag(argument): + if isinstance(argument, str) and _IsFlag(argument): # This is a named argument. We get its value from this arg or the next. # Terminology: @@ -963,7 +986,7 @@ def _IsMultiCharFlag(argument): return argument.startswith('--') or re.match('^-[a-zA-Z]', argument) -def _ParseValue(value, index, arg, metadata): +def _ParseValue(value, index, arg, metadata, variable_env): """Parses value, a string, into the appropriate type. The function used to parse value is determined by the remaining arguments. @@ -976,6 +999,8 @@ def _ParseValue(value, index, arg, metadata): Returns: value, parsed into the appropriate type for calling a function. """ + if value in variable_env: + return variable_env[value] parse_fn = parser.DefaultParseValue # We check to see if any parse function from the fn metadata applies here.