Skip to content
aviral gupta

// I3.5 · ~35 min · Intermediate

functools and decorators

After this lesson you can write a decorator that wraps any function and keeps its name, cache repeated calls with cache and lru_cache, pre-fill arguments with partial and fold a list with reduce.

Lesson 5 of 6 in I3 Iteration and functional tools

You will be able to

  • Write a decorator that wraps a function, and keep its name and docstring with functools.wraps
  • Cache the results of a function with cache and lru_cache, and know when not to
  • Pre-fill arguments with partial, and fold an iterable into one value with reduce
  1. Warm-up · Activity 1 of 7

    Warm-up from module B3: a function is an object you can store in a variable. What does this print?

    def shout(text):
        return text.upper() + "!"
    
    say = shout
    print(say("hi"), say.__name__)
  2. Predict · Activity 2 of 7

    Predict before you read on: hello is defined once and called once. What does this print?

    def twice(func):
        def wrapper():
            func()
            func()
        return wrapper
    
    @twice
    def hello():
        print("hello", end=" ")
    
    hello()
  3. Practice · Activity 3 of 7

    Fill in the functools decorator that copies the name and docstring of func onto the wrapper.

    import functools
    
    def logged(func):
        @functools.____(func)
        def wrapper(*args, **kwargs):
            print("calling", func.__name__)
            return func(*args, **kwargs)
        return wrapper
    
    @logged
    def area(w, h):
        """Return the area of a rectangle."""
        return w * h
    
    print(area(2, 3), area.__name__, area.__doc__)
    @functools.(func)
  4. Practice · Activity 4 of 7

    calls counts how often the body of square runs. What does this print?

    from functools import cache
    
    calls = 0
    
    @cache
    def square(n):
        global calls
        calls += 1
        return n * n
    
    square(4)
    square(4)
    square(5)
    print(calls)
  5. Practice · Activity 5 of 7

    partial and reduce are imported from functools. Match each expression to its value.

  6. Brain teaser · Activity 6 of 7

    Brain teaser. This decorator does not use functools.wraps. What does the last line print?

    def logged(func):
        def wrapper(*args, **kwargs):
            return func(*args, **kwargs)
        return wrapper
    
    @logged
    def area(w, h):
        """Return the area."""
        return w * h
    
    print(area.__name__, area.__doc__)
  7. Apply · Activity 7 of 7

    Mini-task. Write a decorator trace that prints -> name(args) before each call and <- result after it, and returns the result. Use functools.wraps. Decorate add(a, b), call add(2, 3), and print add.__name__ and add.__doc__.

    Check your work against this list

Build it yourself

Read the worked example, then write the exercises. Your code runs in your browser or on your computer and is never uploaded.

Worked example

Watching a cache at work

logged is a decorator that prints every call that reaches the function. fib has two decorators: @functools.cache on the outside, @logged on the inside, so the log shows only the calls the cache could not answer. The second half uses partial to make a binary parser and reduce to multiply a list. Callable[..., Any] is the annotation for "any function".

main.py

import functools
from collections.abc import Callable
from typing import Any


def logged(func: Callable[..., Any]) -> Callable[..., Any]:
    """Print every call of func with its arguments."""

    @functools.wraps(func)
    def wrapper(*args: Any, **kwargs: Any) -> Any:
        print("call", func.__name__, args)
        return func(*args, **kwargs)

    return wrapper


@functools.cache
@logged
def fib(n: int) -> int:
    """Return the n-th Fibonacci number."""
    return n if n < 2 else fib(n - 1) + fib(n - 2)


print(fib(5))
print(fib(6))
print(fib.cache_info())
print(fib.__name__, "-", fib.__doc__)

parse_binary = functools.partial(int, base=2)
print(parse_binary("1010"), parse_binary("111"))

product = functools.reduce(lambda acc, n: acc * n, [1, 2, 3, 4, 5], 1)
print(product)

Run it with

python main.py

Output

call fib (5,)
call fib (4,)
call fib (3,)
call fib (2,)
call fib (1,)
call fib (0,)
5
call fib (6,)
8
CacheInfo(hits=5, misses=7, maxsize=None, currsize=7)
fib - Return the n-th Fibonacci number.
10 7
120
  • Each n from 5 down to 0 was computed once; every repeated fib(n) came from the cache and was not logged.
  • fib(6) needed only one new call, because fib(5) and fib(4) were already stored.
  • fib kept its name and docstring through both decorators, thanks to functools.wraps.
  • parse_binary is int with base=2 already filled in, so it takes just the string.
Change it and run it

Tab indents and Shift+Tab outdents. To leave the editor with the keyboard, press Esc, then Tab.

The first run downloads Python for your browser (up to 6.5 MB) and keeps it cached. Your code stays on your device.

Exercises

Exercise 1 of 3

A shouting decorator

Finish the decorator shout. It wraps a function that returns a string and makes the result upper case. The wrapper must pass on any positional and keyword arguments, and the decorated function must keep its name and docstring. greet("ada") then returns "HELLO, ADA!".

Tab indents and Shift+Tab outdents. To leave the editor with the keyboard, press Esc, then Tab.

The first run downloads Python for your browser (up to 6.5 MB) and keeps it cached. Your code stays on your device.

Hints
  1. Hint 1

    Call the original, then change what it returned: func(*args, **kwargs).upper().

  2. Hint 2

    Put @functools.wraps(func) on the line above def wrapper.

  3. Hint 3

    The wrapper takes *args and **kwargs and passes both on unchanged, so it fits any function.

Show a solution

One way to solve it. Yours can look different and still pass the checks.

import functools
from collections.abc import Callable
from typing import Any


def shout(func: Callable[..., str]) -> Callable[..., str]:
    """Make the string result of func upper case."""

    @functools.wraps(func)
    def wrapper(*args: Any, **kwargs: Any) -> str:
        return func(*args, **kwargs).upper()

    return wrapper


@shout
def greet(name: str, punctuation: str = "!") -> str:
    """Greet someone by name."""
    return f"hello, {name}{punctuation}"


if __name__ == "__main__":
    print(greet("ada"))
Run it on your computer

Install Python 3.14 or newer. Save these files in one folder, open a terminal in that folder, and run the commands below.

main.py

import functools
from collections.abc import Callable
from typing import Any


def shout(func: Callable[..., str]) -> Callable[..., str]:
    """Make the string result of func upper case."""

    def wrapper(*args: Any, **kwargs: Any) -> str:
        return func(*args, **kwargs)

    return wrapper


@shout
def greet(name: str, punctuation: str = "!") -> str:
    """Greet someone by name."""
    return f"hello, {name}{punctuation}"


if __name__ == "__main__":
    print(greet("ada"))

test_main.py

from main import greet, shout


def test_upper():
    """greet('ada') is shouted"""
    got = greet("ada")
    assert got == "HELLO, ADA!", f"greet('ada') returned {got!r}, expected 'HELLO, ADA!'"


def test_keyword_argument():
    """Keyword arguments reach the function"""
    got = greet("bo", punctuation="?")
    assert got == "HELLO, BO?", f"greet('bo', punctuation='?') returned {got!r}, expected 'HELLO, BO?'"


def test_name_and_docstring():
    """greet keeps its name and docstring"""
    got = greet.__name__, greet.__doc__
    assert got == ("greet", "Greet someone by name."), f"greet.__name__ and __doc__ are {got!r}: use functools.wraps"


def test_other_function():
    """shout works on any function that returns a string"""

    @shout
    def join(*words):
        return "-".join(words)

    got = join("a", "b")
    assert got == "A-B", f"a shouted join('a', 'b') returned {got!r}, expected 'A-B'"

On macOS and Linux, type python3 wherever these commands say python, as in the first lesson.

Run the program:

python main.py

Run the checks (needs learnrun.py in the same folder):

python learnrun.py test
Download learnrun.py

Exercise 2 of 3

Counting grid paths with a cache

count_paths(rows, cols) counts the ways across a grid of rows steps down and cols steps right, moving only down or right. It is correct, but it recomputes the same smaller grids again and again. Add @cache from functools, so each grid is computed once. The tests call count_paths.cache_info(), which only a cached function has.

Tab indents and Shift+Tab outdents. To leave the editor with the keyboard, press Esc, then Tab.

The first run downloads Python for your browser (up to 6.5 MB) and keeps it cached. Your code stays on your device.

Hints
  1. Hint 1

    A decorator goes on the line directly above def.

  2. Hint 2

    cache is already imported, so the line is just @cache.

  3. Hint 3

    The recursion calls the name count_paths, which after decorating is the cached version, so the inner calls are cached too.

Show a solution

One way to solve it. Yours can look different and still pass the checks.

from functools import cache


@cache
def count_paths(rows: int, cols: int) -> int:
    """Count the paths through a grid, moving only right or down."""
    if rows == 0 or cols == 0:
        return 1
    return count_paths(rows - 1, cols) + count_paths(rows, cols - 1)


if __name__ == "__main__":
    print(count_paths(2, 2))
Run it on your computer

Install Python 3.14 or newer. Save these files in one folder, open a terminal in that folder, and run the commands below.

main.py

from functools import cache


def count_paths(rows: int, cols: int) -> int:
    """Count the paths through a grid, moving only right or down."""
    if rows == 0 or cols == 0:
        return 1
    return count_paths(rows - 1, cols) + count_paths(rows, cols - 1)


if __name__ == "__main__":
    print(count_paths(2, 2))

test_main.py

from main import count_paths


def test_small_grids():
    """A 1 by 1 grid has 2 paths, a 2 by 2 grid has 6"""
    got = count_paths(1, 1), count_paths(2, 2)
    assert got == (2, 6), f"count_paths(1, 1) and count_paths(2, 2) returned {got!r}, expected (2, 6)"


def test_straight_line():
    """With no steps down there is one path"""
    got = count_paths(0, 5)
    assert got == 1, f"count_paths(0, 5) returned {got!r}, expected 1"


def test_bigger_grid():
    """A 10 by 10 grid has 184756 paths"""
    got = count_paths(10, 10)
    assert got == 184756, f"count_paths(10, 10) returned {got!r}, expected 184756"


def test_is_cached():
    """count_paths is cached and reuses results"""
    assert hasattr(count_paths, "cache_info"), "count_paths has no cache_info(): put @cache above def count_paths"
    count_paths(3, 3)
    hits = count_paths.cache_info().hits
    assert hits > 0, f"the cache reports {hits} hits after count_paths(3, 3)"

On macOS and Linux, type python3 wherever these commands say python, as in the first lesson.

Run the program:

python main.py

Run the checks (needs learnrun.py in the same folder):

python learnrun.py test
Download learnrun.py

Exercise 3 of 3

partial and reduce

Write two functions. parse_base(base) returns a function that reads a string as a number in that base: parse_base(16)("ff") is 255. Use partial with int. merge_all(dicts) merges a list of dicts from left to right, later values winning, with reduce and the | operator: [{"a": 1}, {"a": 3, "b": 2}] gives {"a": 3, "b": 2}. An empty list gives {}.

Tab indents and Shift+Tab outdents. To leave the editor with the keyboard, press Esc, then Tab.

The first run downloads Python for your browser (up to 6.5 MB) and keeps it cached. Your code stays on your device.

Hints
  1. Hint 1

    partial(int, base=base) is int with its base keyword already filled in.

  2. Hint 2

    a | b makes a new dict from a and b, with b’s values winning. reduce applies it pair by pair from the left.

  3. Hint 3

    Give reduce an empty dict as its initial value, so an empty list returns {}: reduce(lambda merged, d: merged | d, dicts, empty).

Show a solution

One way to solve it. Yours can look different and still pass the checks.

from collections.abc import Callable
from functools import partial, reduce


def parse_base(base: int) -> Callable[[str], int]:
    """Return a function that reads a string as a number in base."""
    return partial(int, base=base)


def merge_all(dicts: list[dict[str, int]]) -> dict[str, int]:
    """Merge dicts from left to right; later values win."""
    empty: dict[str, int] = {}
    return reduce(lambda merged, d: merged | d, dicts, empty)


if __name__ == "__main__":
    print(parse_base(16)("ff"))
    print(merge_all([{"a": 1}, {"a": 3, "b": 2}]))
Run it on your computer

Install Python 3.14 or newer. Save these files in one folder, open a terminal in that folder, and run the commands below.

main.py

from collections.abc import Callable
from functools import partial, reduce


def parse_base(base: int) -> Callable[[str], int]:
    """Return a function that reads a string as a number in base."""
    return int


def merge_all(dicts: list[dict[str, int]]) -> dict[str, int]:
    """Merge dicts from left to right; later values win."""
    return {}


if __name__ == "__main__":
    print(parse_base(16)("ff"))
    print(merge_all([{"a": 1}, {"a": 3, "b": 2}]))

test_main.py

from main import merge_all, parse_base


def test_parse_base():
    """Base 2 reads 101 as 5, base 16 reads ff as 255"""
    got = parse_base(2)("101"), parse_base(16)("ff")
    assert got == (5, 255), f"parse_base(2)('101') and parse_base(16)('ff') returned {got!r}, expected (5, 255)"


def test_merge():
    """Later dicts win"""
    got = merge_all([{"a": 1}, {"b": 2}, {"a": 3}])
    assert got == {"a": 3, "b": 2}, f"merge_all returned {got!r}, expected {{'a': 3, 'b': 2}}"


def test_merge_empty():
    """No dicts give an empty dict"""
    got = merge_all([])
    assert got == {}, f"merge_all([]) returned {got!r}, expected {{}}"


def test_inputs_unchanged():
    """merge_all does not change the dicts it is given"""
    first = {"a": 1}
    merge_all([first, {"a": 2}])
    assert first == {"a": 1}, f"the first dict was changed to {first!r}"

On macOS and Linux, type python3 wherever these commands say python, as in the first lesson.

Run the program:

python main.py

Run the checks (needs learnrun.py in the same folder):

python learnrun.py test
Download learnrun.py

Common mistakes

A decorator that forgets to return the wrapper

def logged(func):
    def wrapper(*args, **kwargs):
        print("calling", func.__name__)
        return func(*args, **kwargs)


@logged
def greet(name):
    return "hi " + name


print(greet("Ada"))

What Python prints

TypeError: 'NoneType' object is not callable

Why, and the fix

logged defines wrapper but never returns it, so logged returns None, and @logged sets greet = None. Add return wrapper as the last line of the decorator, at the same indentation as def wrapper.

Caching a function called with a list

from functools import cache


@cache
def total(prices):
    return sum(prices)


print(total([1, 2, 3]))

What Python prints

TypeError: unhashable type: 'list'

Why, and the fix

The cache is a dict with the arguments as its key, and a list cannot be a dict key because it can change. Pass a tuple instead, total((1, 2, 3)), or leave the function uncached.

reduce on an empty list without a start value

from functools import reduce

print(reduce(lambda a, b: a + b, []))

What Python prints

TypeError: reduce() of empty iterable with no initial value

Why, and the fix

With nothing to fold and no start value, reduce has nothing to return. Pass the start value as the third argument: reduce(lambda a, b: a + b, [], 0) gives 0.

Python in the browser: Pyodide 314.0.7, MPL-2.0. Licence and source

Exit ticket

5 questions, no hints. Score 80% or more to complete the lesson.

Finish every activity above to unlock the exit ticket.

Report a problem

Spotted something wrong or unclear? Say what, and it will be checked and fixed.

#

At least 20 characters.

Only if you want a reply.

Key ideas

A decorator wraps a function

A function is an object: you can pass it to another function and get a new one back. A decorator takes a function and returns a replacement, usually an inner wrapper(*args, **kwargs) that does something before or after calling the original. @logged above def area is the same as writing area = logged(area) after it. Decorate the wrapper with @functools.wraps(func): without it, area.__name__ becomes "wrapper" and the docstring is lost, which confuses help(), tracebacks and tests.

cache and lru_cache remember results

@cache stores each result under its arguments, so a second call with the same arguments returns the stored result without running the body. That suits functions whose result depends only on their arguments, such as a recursive Fibonacci, which drops from exponential to linear time. @lru_cache(maxsize=128) keeps only the most recently used results. The arguments must be hashable, so a list raises TypeError. Never cache a function that reads the clock, a file or input(). cache_info() shows hits and misses.

partial and reduce

partial(func, *args, **kwargs) returns a new callable with some arguments already filled in: partial(int, base=2)("101") is 5. It is handy where an API wants a function of one argument, such as a key= or a callback. reduce(function, iterable, initial) folds an iterable from the left into one value: reduce(f, [a, b, c]) is f(f(a, b), c). Pass initial, so an empty iterable gives a result instead of a TypeError. For sums, use sum; reduce is for other folds, such as merging dicts.

Sources

Last reviewed September 29, 2026