Free Handbook · Every example compiled & verified

Higher-Order Functions

Functions that take and return functions: map, flatMap and filter, how for-comprehensions desugar, composition, partial application, closures, partial functions and your own HOFs.

0 / 142 lessons🔥 0 day streak
ShareXLinkedIn

Module 07 · what you'll be able to do

  • Write methods that take a function as a parameter and methods that return a new function
  • Explain map, flatMap and filter as higher-order functions and rewrite a for-comprehension into them by hand
  • Compose functions with andThen and compose, and fix some arguments with partial application
  • Use closures that capture state, and know when captured mutable state bites
  • Use partial functions with collect, and write your own reusable HOFs such as retry and timing wrappers
01

Taking and returning functions

A higher-order function (HOF) is a function that takes another function as a parameter, returns one, or both. In Scala a function is an ordinary value with a type such as Int => Int ("takes an Int, returns an Int") or (String, Int) => Boolean. You pass it like any other argument — you have been doing it since Module 03 every time you wrote xs.map(_ * 2).

scalaMain.scala
def applyTwice(f: Int => Int, x: Int): Int = f(f(x))

def multiplier(n: Int): Int => Int = x => x * n

def describe(label: String, f: Int => Int): String =
  s"$label: ${(1 to 4).map(f).mkString(", ")}"

@main def run(): Unit =
  println(applyTwice(_ + 3, 10))
  println(applyTwice(x => x * x, 3))

  val triple = multiplier(3)
  println(triple(7))
  println(applyTwice(triple, 2))

  println(describe("double", multiplier(2)))
  println(describe("square", x => x * x))
Outputcompiled & run with real Scala
16
81
21
18
double: 2, 4, 6, 8
square: 1, 4, 9, 16

multiplier returns a function: multiplier(3) is a value of type Int => Int you can store, call and pass on. applyTwice does not care which function it gets, only its type.

Your turn

Write def adder(n: Int): Int => Int and use it with applyTwice to add 5 twice to 100.

Error you will hit

A lambda with no parameter type to infer

scala
@main def run(): Unit =
  val inc = x => x + 1
  println(inc(4))
-- [E081] Type Error: Main.scala:2:12
2 |  val inc = x => x + 1
  |            ^
  |            Missing parameter type
  |
  |            I could not infer the type of the parameter x
1 error found
Compilation failed
Why the compiler said that

When you pass x => x + 1 to map, the compiler knows from map's signature what x must be. Stored in a bare val there is nothing to infer from, and x could be an Int, a String or anything else with a +.

The fix

Give the parameter a type, or give the val a function type so the compiler can work backwards.

scala
@main def run(): Unit =
  val inc = (x: Int) => x + 1
  val dec: Int => Int = x => x - 1
  println(inc(4))
  println(dec(4))
02

map, flatMap and filter as HOFs

The three workhorses are all higher-order functions. map(f) applies f to every element and keeps the shape. filter(p) keeps the elements for which the predicate p is true. flatMap(f) is for when f itself returns a collection: it maps and then flattens the results into one collection instead of a collection of collections.

scalaMain.scala
@main def run(): Unit =
  val lines = List("to be or", "not to be")

  println(lines.map(_.length))
  println(lines.map(_.split(" ").toList))
  println(lines.flatMap(_.split(" ").toList))
  println(lines.flatMap(_.split(" ")).filter(_.length > 2))

  val xs = List(1, 2, 3)
  println(xs.flatMap(x => List(x, x * 10)))
  println(xs.map(x => List(x, x * 10)).flatten)
  println(xs.flatMap(x => if x % 2 == 1 then List(x) else Nil))
Outputcompiled & run with real Scala
List(8, 9)
List(List(to, be, or), List(not, to, be))
List(to, be, or, not, to, be)
List(not)
List(1, 10, 2, 20, 3, 30)
List(1, 10, 2, 20, 3, 30)
List(1, 3)

map with a function returning a list gives List[List[String]]; flatMap gives List[String]. flatMap is exactly map followed by flatten. The last line shows that flatMap can do the job of filter too, by returning an empty list to drop an element.

Your turn

Given List("a,b", "c", "d,e,f"), produce List("a", "b", "c", "d", "e", "f") with one flatMap.

To see there is nothing magic here, write map and filter yourself for a list. Both are a foldRight — a fold that starts from the end — which rebuilds the list with :: in the original order.

scalaMain.scala
def myMap[A, B](xs: List[A], f: A => B): List[B] =
  xs.foldRight(List.empty[B])((x, acc) => f(x) :: acc)

def myFilter[A](xs: List[A], p: A => Boolean): List[A] =
  xs.foldRight(List.empty[A])((x, acc) => if p(x) then x :: acc else acc)

@main def run(): Unit =
  val nums = List(1, 2, 3, 4, 5)
  println(myMap(nums, _ * 100))
  println(myFilter(nums, _ % 2 == 0))
  println(myMap(List("a", "bb"), _.length))
Outputcompiled & run with real Scala
List(100, 200, 300, 400, 500)
List(2, 4)
List(1, 2)

[A, B] are type parameters — the function works for any element types. Module 09 covers generics properly.

03

How for-comprehensions become map and flatMap

Module 02 introduced for ... yield. The compiler never runs a for as a loop of its own: it rewrites it into calls to flatMap, map and withFilter. The rules are mechanical:

  • The last generator becomes map (with yield) or foreach (with do).
  • Every earlier generator becomes flatMap, wrapping everything after it.
  • A guard if cond becomes withFilter(cond) on the generator before it — a lazy filter.
  • A definition y = expr becomes a map that carries y along.
scalaMain.scala
@main def run(): Unit =
  val sizes = List("S", "M")
  val colours = List("red", "blue", "green")

  val withFor = for
    s <- sizes
    c <- colours
    if c != "blue"
  yield s"$s-$c"

  val byHand =
    sizes.flatMap(s =>
      colours.withFilter(c => c != "blue").map(c => s"$s-$c"))

  println(withFor)
  println(byHand)
  println(withFor == byHand)
Outputcompiled & run with real Scala
List(S-red, S-green, M-red, M-green)
List(S-red, S-green, M-red, M-green)
true

Both forms build the same list, because they are the same program. The for version is easier to read once there are more than two steps.

Your turn

Add a third generator n <- 1 to 2 and write the by-hand version too. Which call becomes the new map?

Visualizefor x <- List(1, 2); y <- List("a", "b") yield s"$x$y"Step 1 / 6
val xs = List(1, 2)
val ys = List("a", "b")
val r = xs.flatMap(x =>
ys.map(y => s"$x$y"))
println(r)
Line 3

flatMap takes the first element of xs.

Variables now
x1
All 6 steps as a table
StepLineWhat happenedVariables now
13flatMap takes the first element of xs.x = 1
24For that x, the inner map runs over all of ys and returns a list.x = 1 inner = List(1a, 1b)
33flatMap takes the next element of xs.x = 2
44The inner map runs again.x = 2 inner = List(2a, 2b)
53flatMap flattens the two inner lists into one.r = List(1a, 1b, 2a, 2b)
65Printed.

Because the rewrite only needs map and flatMap, a for works on any type that has them — Option, Either, Try, Future. The price is that all generators in one for must be the same kind of container, since each flatMap expects its function to return that kind.

Error you will hit

Mixing Option and List in one for

scala
@main def run(): Unit =
  val maybeUser: Option[String] = Some("asha")
  val roles = List("admin", "dev")
  val pairs = for
    u <- maybeUser
    r <- roles
  yield s"$u:$r"
  println(pairs)
-- [E007] Type Mismatch Error: Main.scala:6:4
6 |    r <- roles
  |    ^
  |    Found:    List[String]
  |    Required: Option[Any]
  |
7 |  yield s"$u:$r"
1 error found
Compilation failed
Why the compiler said that

The first generator is an Option, so the whole for becomes maybeUser.flatMap(u => roles.map(...)). Option.flatMap needs a function that returns an Option, but the body returns a List.

The fix

Convert so every generator is the same type — here turn the Option into a List with toList (an empty list for None).

scala
@main def run(): Unit =
  val maybeUser: Option[String] = Some("asha")
  val roles = List("admin", "dev")
  val pairs = for
    u <- maybeUser.toList
    r <- roles
  yield s"$u:$r"
  println(pairs)
04

Function composition: andThen and compose

Two functions can be glued into one. f andThen g means "run f, then feed its result to g" — it reads in execution order. f compose g is the mathematical order, f(g(x)): g runs first. Most Scala code uses andThen because it reads left to right like a pipeline.

scalaMain.scala
val trim: String => String = _.trim
val lower: String => String = _.toLowerCase
val dashes: String => String = _.replace(" ", "-")

@main def run(): Unit =
  val slugify = trim andThen lower andThen dashes
  println(slugify("  Higher Order Functions "))

  val addOne: Int => Int = _ + 1
  val double: Int => Int = _ * 2
  println((addOne andThen double)(5))
  println((addOne compose double)(5))

  val steps = List(trim, lower, dashes)
  val pipeline = steps.reduce(_ andThen _)
  println(pipeline(" Build A Pipeline "))
Outputcompiled & run with real Scala
higher-order-functions
12
11
build-a-pipeline

(addOne andThen double)(5) is (5 + 1) * 2; (addOne compose double)(5) is 5 * 2 + 1. The last pipeline is built from a list of functions, so steps can be chosen at runtime — from configuration, for example.

Your turn

Add a step that removes every character that is not a letter, digit or dash (_.filter(c => c.isLetterOrDigit || c == '-')) and put it at the end of steps.

Composing methods
andThen is a method of function values, but Scala 3 turns a method defined with def into a function value automatically wherever one is needed (eta-expansion, Module 03). So with def parse(s: String): Int and def validate(n: Int): Boolean, parse andThen validate compiles as it is. Older Scala 2 code writes (parse _).andThen(validate); Scala 3 now warns that the trailing _ is unnecessary.
05

Partial application and closures

Partial application fixes some arguments of a function now and leaves the rest for later. With a curried method (Module 03) you just stop after the first parameter list. With an ordinary method, use _ for the arguments you are leaving open: price(_, 0.18) is a function waiting for the first argument.

scalaMain.scala
def price(base: Double, taxRate: Double): Double = base * (1 + taxRate)

def log(level: String)(msg: String): String = s"[$level] $msg"

@main def run(): Unit =
  val withGst = price(_, 0.18)
  println(withGst(100))
  println(List(50.0, 200.0).map(withGst))

  val warn = log("WARN")
  val error = log("ERROR")
  println(warn("disk 80% full"))
  println(error("disk full"))
Outputcompiled & run with real Scala
118.0
List(59.0, 236.0)
[WARN] disk 80% full
[ERROR] disk full

withGst has type Double => Double, so it slots straight into map. log("WARN") is a String => String with the level fixed.

A closure is a function that uses a variable from the scope where it was created. It keeps that variable alive after the enclosing method has returned. That is how multiplier(3) in the first lesson remembered its 3 — and a closure over a var can carry state between calls.

scalaMain.scala
def makeCounter(): () => Int =
  var count = 0
  () =>
    count += 1
    count

@main def run(): Unit =
  val a = makeCounter()
  val b = makeCounter()
  println(a())
  println(a())
  println(a())
  println(b())

  var rate = 10
  val fee = (amount: Int) => amount * rate / 100
  println(fee(500))
  rate = 20
  println(fee(500))
Outputcompiled & run with real Scala
1
2
3
1
50
100

Each call of makeCounter creates a new count, so a and b count independently. The second half is the trap: a closure captures the variable, not its value at creation time, so changing rate later changes what fee returns. Capture vals and this surprise cannot happen.

Your turn

Change makeCounter to take a step: Int and count up by it.

In real jobs
Closures over mutable state are the classic Spark bug: a function passed to rdd.map that increments a local var runs on other machines, each with its own copy, and the driver's var never changes. Keep functions you hand to a framework free of captured vars — see Spark with Scala.
06

Partial functions and collect

A PartialFunction[A, B] is a function that is only defined for some inputs. You write one as a block of case clauses. It can tell you in advance whether it handles a value (isDefinedAt), and collect uses exactly that: it maps the elements the partial function accepts and silently drops the rest — a filter and a map in one step.

scalaMain.scala
@main def run(): Unit =
  val parseDigit: PartialFunction[String, Int] =
    case s if s.length == 1 && s.head.isDigit => s.toInt

  println(parseDigit.isDefinedAt("7"))
  println(parseDigit.isDefinedAt("x"))
  println(parseDigit("7"))

  val inputs = List("3", "x", "10", "8")
  println(inputs.collect(parseDigit))

  val fallback: PartialFunction[String, Int] = { case _ => -1 }
  println(inputs.map(parseDigit.orElse(fallback)))
  println(inputs.map(parseDigit.lift))
Outputcompiled & run with real Scala
true
false
7
List(3, 8)
List(3, -1, -1, 8)
List(Some(3), None, None, Some(8))

orElse chains partial functions: the first that accepts the value wins. lift turns a partial function into an ordinary one returning Option — None where it was undefined. Option is the subject of Module 08.

scalaMain.scala
enum Event:
  case Click(x: Int, y: Int)
  case Key(ch: Char)
  case Scroll(delta: Int)

@main def run(): Unit =
  import Event.*
  val events = List(Click(1, 2), Key('a'), Scroll(-3), Click(5, 5), Key('b'))
  val clicks = events.collect { case Click(x, y) => s"($x,$y)" }
  val typed = events.collect { case Key(c) => c }.mkString
  println(clicks)
  println(typed)
Outputcompiled & run with real Scala
List((1,2), (5,5))
ab

The everyday use: pull one case out of a list of an ADT and destructure it in the same step.

Error you will hit

map with a case block that does not cover every input

scala
@main def run(): Unit =
  val xs: List[Any] = List(1, "two", 3)
  val doubled = xs.map { case n: Int => n * 2 }
  println(doubled)
Exception in thread "main" scala.MatchError: two (of class java.lang.String)
	at Main$package$.$anonfun$1(Main.scala:3)
	at scala.collection.immutable.List.map(List.scala:248)
	at Main$package$.run(Main.scala:3)
	at run.main(Main.scala:1)
Why the compiler said that

A { case ... } block passed to map is applied to every element. When "two" arrives no case matches, and the program throws MatchError at runtime — after compiling without complaint.

The fix

If dropping unmatched elements is what you mean, use collect. If every element must be handled, add a case for the rest.

scala
@main def run(): Unit =
  val xs: List[Any] = List(1, "two", 3)
  val doubled = xs.collect { case n: Int => n * 2 }
  println(doubled)
07

Writing your own higher-order functions

Once functions are values, recurring "wrap this work in some policy" code becomes a function you write once. Two classics are retry (call an operation until it succeeds or you run out of attempts) and timing (measure how long a block took). Both take the work to do as a parameter, often as a by-name parameter so the call site reads like a built-in block (Module 03).

scalaMain.scala
def retry[A](maxAttempts: Int)(op: Int => Option[A]): Option[A] =
  var attempt = 1
  var result: Option[A] = None
  while result.isEmpty && attempt <= maxAttempts do
    result = op(attempt)
    if result.isEmpty then println(s"attempt $attempt failed")
    attempt += 1
  result

@main def run(): Unit =
  // a fake service that only answers on its third call
  val flaky: Int => Option[String] =
    n => if n >= 3 then Some(s"data on attempt $n") else None

  println(retry(5)(flaky))
  println(retry(2)(flaky))
Outputcompiled & run with real Scala
attempt 1 failed
attempt 2 failed
Some(data on attempt 3)
attempt 1 failed
attempt 2 failed
None

The retry policy lives in one place; the operation is anything of type Int => Option[A]. A real version would also sleep between attempts and use Try or Either to keep the last error (Module 08).

Your turn

Make retry return Either[String, A] with Left("gave up after N attempts") when every attempt fails.

A timing wrapper needs a clock. Taking the clock as a parameter, instead of calling System.nanoTime() inside, keeps the function testable: production passes the real clock, and a test (or this page) passes a fake one that returns known values.

scalaMain.scala
def timed[A](label: String, clock: () => Long)(block: => A): A =
  val start = clock()
  val result = block
  val end = clock()
  println(s"$label took ${end - start} ms")
  result

@main def run(): Unit =
  var fakeNow = 1000L
  val fakeClock = () =>
    fakeNow += 40
    fakeNow

  val total = timed("sum", fakeClock) {
    (1 to 100).sum
  }
  println(total)
Outputcompiled & run with real Scala
sum took 40 ms
5050

The call site timed("sum", clock) { ... } looks like a language feature, but it is an ordinary curried method whose second parameter list takes a by-name block. In production you would pass () => System.currentTimeMillis().

Higher-order function
A function that takes a function as a parameter, returns a function, or both.
Function type
The type of a function value, such as Int => Int or (String, Int) => Boolean.
flatMap
Maps each element to a collection (or Option, etc.) and flattens the results into one.
withFilter
The lazy filter a for-comprehension guard is rewritten into.
andThen / compose
Combine two functions into one: f andThen g runs f first; f compose g runs g first.
Partial application
Fixing some arguments of a function to get a new function of the remaining ones.
Closure
A function that captures variables from the scope where it was created and keeps them alive.
PartialFunction
A function defined only for some inputs, written as case clauses, with isDefinedAt.
collect
Applies a partial function to the elements it accepts and drops the rest.
lift
Turns a partial function into a total one returning Option.
Quick check

What does for { a <- xs; b <- ys } yield a + b compile to?

Quick check

With val f: Int => Int = _ + 1 and val g: Int => Int = _ * 10, what is (f andThen g)(2)?

Frequently asked questions

What is a higher-order function in Scala?
A function that takes another function as a parameter or returns one. map, filter, flatMap, foldLeft and sortBy are all higher-order functions: you pass them the behaviour, and they handle the iteration.
What is the difference between map and flatMap in Scala?
map applies a function to each element and keeps one result per element. flatMap is for functions that return a collection (or an Option, Either, Future): it maps and then flattens, so you get one flat result instead of nested ones. A for-comprehension with several generators is a chain of flatMap calls ending in map.
What is the difference between andThen and compose?
f andThen g runs f and passes its result to g. f compose g runs g first and passes its result to f. They build the same kind of function; andThen reads in execution order, which is why it is more common.

Finish the Scala handbook, then get hired

Sit the exam for your certificate, run your resume through the ATS checker, and see the jobs that ask for exactly this.

Check my resume
Found this course useful? Share it.
ShareXLinkedIn

Comments

0

Join the conversation. Sign in to leave a comment — we'd love to hear your thoughts.