Skip to content
Merged
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
154 changes: 74 additions & 80 deletions db/jvm/src/JdbcDecoder.scala
Original file line number Diff line number Diff line change
Expand Up @@ -9,9 +9,30 @@ import java.util.UUID
/** Typeclass for reading a single column from a [[java.sql.ResultSet]] by 1-based index.
*
* Instances are provided for all types supported by the [[Flavour]] type mapping. Custom types can be supported by providing a given instance.
*
* ===Nullability===
* A column is read from the `ResultSet` exactly once, because JDBC only guarantees that each column of a row is read once, in left-to-right order. That is why nullability is a
* second method on this typeclass rather than a wrapper around `decode`: the wrapper would have to either read the column twice, or let `decode` throw before it could ask whether
* the value was NULL.
*
* [[decode]] is strict - a SQL NULL is a schema disagreement and is reported as one. [[decodeOption]] is the same read with the opposite answer for NULL, and is what the
* `JdbcDecoder[Option[T]]` instance calls.
*/
trait JdbcDecoder[T]:
/** Read the column, failing if it holds SQL NULL. */
def decode(rs: ResultSet, index: Int): T

/** Read the column, returning `None` if it holds SQL NULL.
*
* The default is correct for any `decode` that tolerates a NULL read without throwing, which is the case for a decoder built on a JDBC getter that returns a sentinel (`0`,
* `false`, `null`). Override it whenever `decode` inspects `rs.wasNull()` itself, or dereferences the value it read - otherwise `Option[T]` inherits the very failure the
* `Option` was meant to express.
*/
def decodeOption(rs: ResultSet, index: Int): Option[T] =
val v = decode(rs, index)
if rs.wasNull() then None else Some(v)
end if
end decodeOption
end JdbcDecoder

object JdbcDecoder:
Expand All @@ -21,99 +42,72 @@ object JdbcDecoder:
try rs.getMetaData.getColumnLabel(index)
catch case _: Exception => index.toString

given JdbcDecoder[Int] with
def decode(rs: ResultSet, index: Int): Int = rs.getInt(index)
end given
private def nullError(rs: ResultSet, index: Int, typeName: String): java.sql.SQLDataException =
new java.sql.SQLDataException(
s"Column \"${columnLabel(rs, index)}\" (index $index) is NULL but was mapped to a non-nullable $typeName. Use Option[$typeName] for nullable columns."
)

given JdbcDecoder[Long] with
def decode(rs: ResultSet, index: Int): Long = rs.getLong(index)
end given
/** Builds a decoder that reads the column exactly once and then decides what NULL means.
*
* `read` is the raw JDBC getter; `convert` runs only after the NULL check, so a converter may dereference its argument freely - `BigDecimal(_)` and `Timestamp#toInstant` would
* both NPE if they ran on the sentinel a NULL read returns.
*
* @param typeName
* How the type is spelled in the error message, so it reads as the Scala type the user actually wrote
*/
private def strict[R, T](typeName: String)(read: (ResultSet, Int) => R)(convert: R => T): JdbcDecoder[T] =
new JdbcDecoder[T]:
def decode(rs: ResultSet, index: Int): T =
val raw = read(rs, index)
if rs.wasNull() then throw nullError(rs, index, typeName)
end if
convert(raw)
end decode

given JdbcDecoder[Double] with
def decode(rs: ResultSet, index: Int): Double = rs.getDouble(index)
end given
override def decodeOption(rs: ResultSet, index: Int): Option[T] =
val raw = read(rs, index)
if rs.wasNull() then None else Some(convert(raw))
end if
end decodeOption

given JdbcDecoder[Float] with
def decode(rs: ResultSet, index: Int): Float = rs.getFloat(index)
end given
private def strict[T](typeName: String)(read: (ResultSet, Int) => T): JdbcDecoder[T] =
strict[T, T](typeName)(read)(identity)

given JdbcDecoder[Boolean] with
def decode(rs: ResultSet, index: Int): Boolean = rs.getBoolean(index)
end given
given JdbcDecoder[Int] = strict("Int")(_.getInt(_))

given JdbcDecoder[String] with
def decode(rs: ResultSet, index: Int): String =
val v = rs.getString(index)
// If the DB returns NULL for a column typed as String (not Option[String]),
// that means the schema is inconsistent. Throw rather than silently returning "".
// For nullable columns, use Option[String] (which uses the JdbcDecoder[Option[T]] instance).
if rs.wasNull() then
throw new java.sql.SQLDataException(
s"Column \"${columnLabel(rs, index)}\" (index $index) is NULL but was mapped to a non-nullable String. Use Option[String] for nullable columns."
)
end if
v
end decode
end given
given JdbcDecoder[Long] = strict("Long")(_.getLong(_))

given JdbcDecoder[BigDecimal] with
def decode(rs: ResultSet, index: Int): BigDecimal =
val v = rs.getBigDecimal(index)
if rs.wasNull() then
throw new java.sql.SQLDataException(
s"Column \"${columnLabel(rs, index)}\" (index $index) is NULL but was mapped to a non-nullable BigDecimal. Use Option[BigDecimal] for nullable columns."
)
end if
BigDecimal(v)
end decode
end given
given JdbcDecoder[Double] = strict("Double")(_.getDouble(_))

given JdbcDecoder[Array[Byte]] with
def decode(rs: ResultSet, index: Int): Array[Byte] =
val v = rs.getBytes(index)
if rs.wasNull() then
throw new java.sql.SQLDataException(
s"Column \"${columnLabel(rs, index)}\" (index $index) is NULL but was mapped to a non-nullable Array[Byte]. Use Option[Array[Byte]] for nullable columns."
)
end if
v
end decode
end given
given JdbcDecoder[Float] = strict("Float")(_.getFloat(_))

given JdbcDecoder[LocalDate] with
def decode(rs: ResultSet, index: Int): LocalDate =
rs.getObject(index, classOf[LocalDate])
end given
given JdbcDecoder[Boolean] = strict("Boolean")(_.getBoolean(_))

given JdbcDecoder[LocalDateTime] with
def decode(rs: ResultSet, index: Int): LocalDateTime =
rs.getObject(index, classOf[LocalDateTime])
end given
given JdbcDecoder[String] = strict("String")(_.getString(_))

given JdbcDecoder[Instant] with
def decode(rs: ResultSet, index: Int): Instant =
val ts = rs.getTimestamp(index)
if rs.wasNull() then
throw new java.sql.SQLDataException(
s"Column \"${columnLabel(rs, index)}\" (index $index) is NULL but was mapped to a non-nullable Instant. Use Option[Instant] for nullable columns."
)
end if
ts.toInstant
end decode
end given
given JdbcDecoder[BigDecimal] = strict("BigDecimal")(_.getBigDecimal(_))(BigDecimal(_))

given JdbcDecoder[UUID] with
def decode(rs: ResultSet, index: Int): UUID =
rs.getObject(index, classOf[UUID])
end given
given JdbcDecoder[Array[Byte]] = strict("Array[Byte]")(_.getBytes(_))

/** Nullable column: returns `None` when the DB value is SQL NULL. */
given JdbcDecoder[LocalDate] = strict("LocalDate")(_.getObject(_, classOf[LocalDate]))

given JdbcDecoder[LocalDateTime] = strict("LocalDateTime")(_.getObject(_, classOf[LocalDateTime]))

given JdbcDecoder[Instant] = strict("Instant")(_.getTimestamp(_))(_.toInstant)

given JdbcDecoder[UUID] = strict("UUID")(_.getObject(_, classOf[UUID]))

/** Nullable column: returns `None` when the DB value is SQL NULL.
*
* Delegates to [[JdbcDecoder.decodeOption]] rather than calling `decode` and testing `rs.wasNull()` afterwards. The latter cannot work: a strict decoder throws on NULL, so the
* test would never be reached, and every reference-typed `Option` column would fail on its first NULL row.
*/
given [T](using inner: JdbcDecoder[T]): JdbcDecoder[Option[T]] with
def decode(rs: ResultSet, index: Int): Option[T] =
val v = inner.decode(rs, index)
if rs.wasNull() then None else Some(v)
end if
end decode
def decode(rs: ResultSet, index: Int): Option[T] = inner.decodeOption(rs, index)

/** `Option[Option[T]]` is not a column shape, but the default would read the column a second time. */
override def decodeOption(rs: ResultSet, index: Int): Option[Option[T]] =
Some(inner.decodeOption(rs, index))
end given

end JdbcDecoder
85 changes: 85 additions & 0 deletions db/test/src/DbSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -162,6 +162,91 @@ class JdbcDecoderSuite extends H2Fixture:
}
}

/** A NULL in a column mapped to `Option[T]` must reach the caller as `None`, whatever `T` is.
*
* The reference types are the ones that used to fail here. Their strict decoders inspect `rs.wasNull()` themselves and throw, so an `Option` wrapper built as "decode, then ask
* whether it was null" never got to ask - it inherited the throw, and told the user to do the very thing they had already done.
*/
test("decode Option[T] - None for NULL, for every supported T") {
withConn { conn =>
def nullOf[T](sqlType: String)(using d: JdbcDecoder[Option[T]]): Option[T] =
val rs = conn.createStatement().executeQuery(s"SELECT CAST(NULL AS $sqlType)")
rs.next()
d.decode(rs, 1)
end nullOf

assertEquals(nullOf[Int]("INT"), None)
assertEquals(nullOf[Long]("BIGINT"), None)
assertEquals(nullOf[Double]("DOUBLE"), None)
assertEquals(nullOf[Boolean]("BOOLEAN"), None)
assertEquals(nullOf[String]("VARCHAR(10)"), None)
assertEquals(nullOf[BigDecimal]("DECIMAL(10,3)"), None)
assertEquals(nullOf[Array[Byte]]("VARBINARY(10)"), None)
assertEquals(nullOf[java.time.LocalDate]("DATE"), None)
assertEquals(nullOf[java.time.LocalDateTime]("TIMESTAMP"), None)
assertEquals(nullOf[java.time.Instant]("TIMESTAMP"), None)
assertEquals(nullOf[java.util.UUID]("UUID"), None)
}
}

test("decode Option[T] - Some for a present value, for every supported T") {
withConn { conn =>
def someOf[T](expr: String)(using d: JdbcDecoder[Option[T]]): Option[T] =
val rs = conn.createStatement().executeQuery(s"SELECT $expr")
rs.next()
d.decode(rs, 1)
end someOf

assertEquals(someOf[String]("'hello'"), Some("hello"))
assertEquals(someOf[BigDecimal]("CAST(123.456 AS DECIMAL(10,3))"), Some(BigDecimal("123.456")))
assertEquals(someOf[java.time.LocalDate]("DATE '2024-03-01'"), Some(java.time.LocalDate.of(2024, 3, 1)))
assertEquals(someOf[Int]("42"), Some(42))
// Array[Byte] has no useful equals, so compare the contents
assertEquals(someOf[Array[Byte]]("CAST(X'01ff' AS VARBINARY(10))").map(_.toList), Some(List[Byte](1, -1)))
}
}

/** The other half of the contract: a NULL in a column mapped to a bare `T` is a schema disagreement, and is reported as one rather than decoded to a sentinel.
*
* This matters most for the primitives, where JDBC hands back `0` / `false` for a NULL and the mistake would otherwise be invisible - a summed column quietly short by however
* many rows were NULL.
*/
test("decode T - a NULL in a non-Option column names the column and the fix") {
withConn { conn =>
def strictNull[T](sqlType: String)(using d: JdbcDecoder[T]): Unit =
val rs = conn.createStatement().executeQuery(s"SELECT CAST(NULL AS $sqlType) AS the_column")
rs.next()
val e = intercept[java.sql.SQLDataException](d.decode(rs, 1))
// H2 upper cases unquoted identifiers, so compare case insensitively
assert(e.getMessage.toLowerCase.contains("the_column"), e.getMessage)
assert(e.getMessage.contains("Use Option["), e.getMessage)
end strictNull

strictNull[Int]("INT")
strictNull[Long]("BIGINT")
strictNull[Double]("DOUBLE")
strictNull[Boolean]("BOOLEAN")
strictNull[String]("VARCHAR(10)")
strictNull[BigDecimal]("DECIMAL(10,3)")
strictNull[Array[Byte]]("VARBINARY(10)")
strictNull[java.time.LocalDate]("DATE")
strictNull[java.time.Instant]("TIMESTAMP")
strictNull[java.util.UUID]("UUID")
}
}

/** A whole row of NULLs, decoded positionally, is the shape the bug actually reached users in. */
test("a row of nullable reference columns decodes to all None") {
withConn { conn =>
val rs = conn
.createStatement()
.executeQuery("SELECT CAST(NULL AS VARCHAR(10)), CAST(NULL AS DECIMAL(10,3)), CAST(NULL AS DATE), CAST(NULL AS BIGINT)")
rs.next()
type Row = (Option[String], Option[BigDecimal], Option[java.time.LocalDate], Option[Long])
assertEquals(summon[JdbcRowDecoder[Row]].decodeRow(rs), (None, None, None, None))
}
}

end JdbcDecoderSuite

class JdbcRowDecoderSuite extends H2Fixture:
Expand Down
62 changes: 31 additions & 31 deletions scautable/src-jvm/ExcelIterator.scala
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,9 @@ class ExcelIterator[K <: Tuple, V <: Tuple](filePath: String, sheetName: String,
(cellRange.getFirstRow, cellRange.getLastRow, cellRange.getFirstColumn, cellRange.getLastColumn)
end parseRange

/** The range's corners, parsed once. `CellRangeAddress.valueOf` is string parsing, and this sits on the per-row path. */
private lazy val parsedRange: Option[(Int, Int, Int, Int)] = colRange.filter(_.nonEmpty).map(parseRange)

/** Validates that headers are unique (no duplicates)
*/
private def validateUniqueHeaders(headers: List[String]): Unit =
Expand All @@ -59,32 +62,28 @@ class ExcelIterator[K <: Tuple, V <: Tuple](filePath: String, sheetName: String,
)
val sheet = workbook.getSheet(sheetName)
// Create an iterator that gives us rows by index for the specified range
colRange match
case Some(range) if range.nonEmpty =>
val (firstRow, lastRow, _, _) = parseRange(range)
parsedRange match
case Some((firstRow, lastRow, _, _)) =>
val dataStartRow = firstRow + 1
val dataRowIndices = (dataStartRow to lastRow).toIterator
dataRowIndices.map(rowIndex => sheet.getRow(rowIndex)).filter(_ != null)
case _ =>
case None =>
sheet.iterator().asScala // For no range, use default iterator
end match
end sheetIterator

// Track current row number for error reporting - starts where data begins
private var currentRowIndex: Int = colRange match
case None => 0
case Some(range) if range.nonEmpty =>
val (firstRow, _, _, _) = parseRange(range)
firstRow + 1 // Skip the header row - data starts at firstRow + 1
case _ => 0
/** The spreadsheet row number (1 based, as Excel shows it) of the row last returned by `next()`, for error reporting.
*
* Taken from the row itself rather than counted, because the iterator skips rows POI reports as absent - a counter would drift from the sheet the moment it met one, and name
* the wrong row in an error.
*/
private var currentRowIndex: Int = 0

// Extract headers from the first row or specified range
private val headers: List[String] =
colRange match
case Some(range) if range.nonEmpty =>
extractHeadersFromRange(range)
case _ =>
extractHeadersFromFirstRow()
parsedRange match
case Some(_) => extractHeadersFromRange()
case None => extractHeadersFromFirstRow()

private lazy val numCellsPerRow = headers.size

Expand All @@ -93,8 +92,8 @@ class ExcelIterator[K <: Tuple, V <: Tuple](filePath: String, sheetName: String,

/** Extract headers from a specified cell range This accesses the header row directly by index
*/
private def extractHeadersFromRange(range: String): List[String] =
val (firstRow, _, firstCol, lastCol) = parseRange(range)
private def extractHeadersFromRange(): List[String] =
val (firstRow, _, firstCol, lastCol) = parsedRange.get
val workbook = ExcelWorkbookCache
.getOrCreate(filePath)
.getOrElse(
Expand All @@ -118,14 +117,13 @@ class ExcelIterator[K <: Tuple, V <: Tuple](filePath: String, sheetName: String,
/** Extract cell values from a row based on the column range
*/
private def extractCellValues(row: org.apache.poi.ss.usermodel.Row): List[String] =
colRange match
case Some(range) if range.nonEmpty =>
val (_, _, firstCol, lastCol) = parseRange(range)
parsedRange match
case Some((_, _, firstCol, lastCol)) =>
val cells =
for i <- firstCol.to(lastCol)
yield row.getCell(i, Row.MissingCellPolicy.CREATE_NULL_AS_BLANK).toString
cells.toList
case _ =>
case None =>
row.cellIterator().asScala.toList.map(_.toString)
end extractCellValues

Expand All @@ -134,6 +132,7 @@ class ExcelIterator[K <: Tuple, V <: Tuple](filePath: String, sheetName: String,
end if

val row = sheetIterator.next()
currentRowIndex = row.getRowNum + 1
val cellValues = extractCellValues(row)

// Validate row has expected number of cells
Expand All @@ -152,17 +151,18 @@ class ExcelIterator[K <: Tuple, V <: Tuple](filePath: String, sheetName: String,
)
)

currentRowIndex += 1
NamedTuple.build[K]()(decodedTuple)
end next

override def hasNext: Boolean =
colRange match
case Some(range) if range.nonEmpty =>
val (_, lastRow, _, _) = parseRange(range)
currentRowIndex <= lastRow
case _ =>
sheetIterator.hasNext
end hasNext
/** Whether another row is actually available.
*
* Asks the row iterator, rather than comparing a counter against the range's last row. The two are not the same: `sheetIterator` drops rows POI reports as absent, which is what
* an entirely blank row inside the range is - a visual separator between groups, say. Counting row numbers therefore promised more rows than existed, and `next()` fell off the
* end of the underlying iterator with a bare `NoSuchElementException`.
*
* The range still bounds the iteration, because `sheetIterator` is built from `firstRow + 1 to lastRow`. Skipping blank rows also keeps reading consistent with compile time
* type inference, which walks the range the same way.
*/
override def hasNext: Boolean = sheetIterator.hasNext

end ExcelIterator
Loading
Loading