diff --git a/NEWS.md b/NEWS.md index c9113cd17..a0fb83c77 100644 --- a/NEWS.md +++ b/NEWS.md @@ -1,5 +1,6 @@ # scoringutils (development version) +- Fixed `as_forecast_()` functions silently creating forecast objects with duplicate column names when asked to rename a column onto a name that already exists in the data (e.g. `predicted = "prob"` while a stale `predicted` column is present). This produced corrupted objects that passed validation and were scored on the wrong column. The constructors now error with a clear message, and `assert_forecast_generic()` rejects data with duplicate column names (#1199). - Fixed `assert_forecast()` for nominal and ordinal forecasts: the error message for incomplete forecasts named the first *complete* forecast instead of the first incomplete one (and `NA` when all forecasts were incomplete), and the methods visibly returned the forecast object instead of `invisible(NULL)` as documented and as all other forecast types do (#1195). - Fixed `as_forecast_quantile()` for sample-based forecasts producing silently wrong quantiles or erroring when `probs` was not symmetric around 0.5 (e.g. `probs = 0.4` or `probs = c(0.1, 0.2)`). Quantiles are now computed at exactly the requested `probs` (deduplicated), and out-of-range `probs` produce a clear assertion error (#1196). - Fixed `interval_coverage()` erroring on quantile levels generated with `seq()` (e.g. `seq(0.05, 0.95, 0.05)`) because the required quantile levels were matched with an exact floating point comparison. Quantile levels are now rounded to 10 decimal places before matching, consistent with the rest of the package. Also fixed `wis()`, `interval_score()` and `quantile_score(weigh = FALSE)` returning `NaN` for forecasts that include the quantile levels 0 and 1 (which form a 100% prediction interval where alpha = 0). Scores are now finite when the observation falls inside the interval, restoring the identity between the WIS and the mean of the quantile scores; the unweighted scores return `Inf` when the observation falls outside a 100% prediction interval (#1202). diff --git a/R/class-forecast.R b/R/class-forecast.R index 6525193f1..a8186d5f9 100644 --- a/R/class-forecast.R +++ b/R/class-forecast.R @@ -7,6 +7,7 @@ #' @param ... Named arguments that are used to rename columns. The names of the #' arguments are the names of the columns that should be renamed. The values #' are the new names. +#' @importFrom cli cli_abort #' @keywords as_forecast as_forecast_generic <- function(data, forecast_unit = NULL, @@ -26,6 +27,24 @@ as_forecast_generic <- function(data, oldnames <- unlist(oldnames[provided]) newnames <- unlist(newnames[provided]) if (!is.null(oldnames) && length(oldnames) > 0) { + # renaming a column onto a name that already exists (and is not itself + # being renamed away) would create duplicate column names + remaining <- setdiff(colnames(data), oldnames) + collides <- newnames %in% remaining + if (any(collides)) { + # sources/targets are used inside the cli glue strings below + sources <- oldnames[collides] # nolint: object_usage_linter. + targets <- newnames[collides] # nolint: object_usage_linter. + cli_abort( + c( + `!` = "Cannot rename {cli::qty(sources)} column{?s} {.val {sources}} + to {.val {targets}}: {?a column/columns} with {?this name/these + names} already exist{?s/} in the data.", + i = "Rename or remove the existing {cli::qty(targets)} + column{?s} first." + ) + ) + } setnames(data, old = oldnames, new = newnames) } @@ -110,6 +129,16 @@ assert_forecast.default <- function( assert_forecast_generic <- function(data, verbose = TRUE) { # check that data is a data.table and that the columns look fine assert_data_table(data, min.rows = 1) + duplicated_cols <- unique(colnames(data)[duplicated(colnames(data))]) + if (length(duplicated_cols) > 0) { + cli_abort( + c( + `!` = "Found duplicate column{?s} in the data: + {.val {duplicated_cols}}.", + i = "Column names must be unique." + ) + ) + } assert_subset(c("observed", "predicted"), colnames(data)) problem <- test_subset(c("sample_id", "quantile_level"), colnames(data)) if (problem) { diff --git a/tests/testthat/test-class-forecast.R b/tests/testthat/test-class-forecast.R index dc6d550ea..c0ab399c1 100644 --- a/tests/testthat/test-class-forecast.R +++ b/tests/testthat/test-class-forecast.R @@ -4,6 +4,66 @@ # see tests for each forecast type for more specific tests. +test_that("as_forecast_generic() errors when renaming onto an existing column", { + # stale `predicted` column alongside the column that should be renamed + dt <- data.table::data.table( + model = "m", + id = 1:2, + observed = factor(c(0, 1)), + predicted = c(0.9, 0.9), + prob = c(0.3, 0.7) + ) + expect_error( + as_forecast_binary(dt, predicted = "prob"), + 'rename column "prob" to "predicted".*already exists' + ) + + # same for other renameable columns, e.g. `quantile_level` + quantile_dt <- data.table::data.table( + model = "m", + target = "t", + observed = 5, + predicted = c(1, 5, 9), + quantile_level = c(0.1, 0.5, 0.9), + q = c(0.1, 0.5, 0.9) + ) + expect_error( + as_forecast_quantile(quantile_dt, quantile_level = "q"), + 'rename column "q" to "quantile_level".*already exists' + ) + + # multiple collisions produce a correctly pluralised message naming all + # source and target columns + multi_dt <- data.table::data.table( + model = "m", + id = 1:2, + observed = 1, + obs = 2, + predicted = 3, + prob = 4 + ) + expect_error( + as_forecast_binary(multi_dt, observed = "obs", predicted = "prob"), + paste0( + 'rename\\s+columns\\s+"obs"\\s+and\\s+"prob"\\s+to\\s+"observed"', + '\\s+and\\s+"predicted".*already\\s+exist\\s+in\\s+the\\s+data' + ) + ) +}) + +test_that("as_forecast_generic() still allows identity renames", { + dt <- data.table::data.table( + model = "m", + id = 1:2, + observed = factor(c(0, 1)), + predicted = c(0.3, 0.7) + ) + expect_no_condition( + as_forecast_binary(dt, observed = "observed", predicted = "predicted") + ) +}) + + # ============================================================================== # is_forecast() # nolint: commented_code_linter # ============================================================================== @@ -39,6 +99,21 @@ test_that("assert_forecast_generic() works as expected with a data.frame", { ) }) +test_that("assert_forecast_generic() errors on duplicate column names", { + dt <- data.table::data.table( + model = "m", + id = 1:2, + observed = factor(c(0, 1)), + predicted = c(0.3, 0.7), + stale = c(0.9, 0.9) + ) + data.table::setnames(dt, "stale", "predicted") + expect_error( + assert_forecast_generic(dt), + "duplicate" + ) +}) + # ============================================================================== # new_forecast() # nolint: commented_code_linter