Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 9 additions & 2 deletions fire/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@ def main(argv):
"""

import asyncio
import functools
import inspect
import json
import os
Expand Down Expand Up @@ -454,7 +455,10 @@ def _Fire(component, args, parsed_flag_args, context, name=None):
handled = False
candidate_errors = []

is_callable = inspect.isclass(component) or inspect.isroutine(component)
# Keep partials on the callable-object path, including on Python 3.14.
is_callable = (
inspect.isclass(component) or inspect.isroutine(component)
) and not isinstance(component, functools.partial)
is_callable_object = callable(component) and not is_callable
is_sequence = isinstance(component, (list, tuple))
is_map = isinstance(component, dict) or inspectutils.IsNamedTuple(component)
Expand Down Expand Up @@ -677,7 +681,10 @@ def _CallAndUpdateTrace(component, args, component_trace, treatment='class',
(varargs, kwargs), consumed_args, remaining_args, capacity = parse(args)

# Call the function.
if inspectutils.IsCoroutineFunction(fn):
coroutine_fn = component if isinstance(component, functools.partial) else fn
while isinstance(coroutine_fn, functools.partial):
coroutine_fn = coroutine_fn.func
if inspectutils.IsCoroutineFunction(coroutine_fn):
try:
loop = asyncio.get_running_loop()
except RuntimeError:
Expand Down
27 changes: 27 additions & 0 deletions fire/fire_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,8 @@

"""Tests for the fire module."""

import functools
import inspect
import os
import sys
from unittest import mock
Expand Down Expand Up @@ -716,6 +718,31 @@ def testFireAsyncio(self):
self.assertEqual(fire.Fire(tc.py3.WithAsyncio,
command=['double', '--count', '10']), 20)

def testFireAsyncPartial(self):
double = tc.py3.WithAsyncio().double
inner = functools.partial(double)
vars(inner)['description'] = 'Keep the nested partial intact.'
cases = [
(functools.partial(double, count=3), [], 6),
(functools.partial(double), ['7'], 14),
(functools.partial(double, count=3), ['--count', '10'], 20),
(functools.partial(inner, 5), [], 10),
]
for component, command, expected in cases:
with self.subTest(command=command, expected=expected):
result = fire.Fire(component, command=command)
if inspect.iscoroutine(result):
result.close()
self.assertEqual(result, expected)

def testFireSyncPartial(self):
double = tc.WithDefaults().double
self.assertEqual(fire.Fire(functools.partial(double, count=3), []), 6)
self.assertEqual(fire.Fire(functools.partial(double), ['7']), 14)
self.assertEqual(fire.Fire(functools.partial(double, count=3),
['--count', '10']), 20)
self.assertEqual(fire.Fire(functools.partial(max, 0), ['5', '3']), 5)


if __name__ == '__main__':
testutils.main()
Loading