Skip to content
Draft
Show file tree
Hide file tree
Changes from 8 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
132 changes: 122 additions & 10 deletions R/count-tally.R
Original file line number Diff line number Diff line change
Expand Up @@ -14,8 +14,11 @@
#'
#' @param x A data frame, data frame extension (e.g. a tibble), or a
#' lazy data frame (e.g. from dbplyr or dtplyr).
#' @param ... <[`data-masking`][rlang::args_data_masking]> Variables to group
#' by.
#' @param ... For `count()`/`add_count()`:
#' <[`data-masking`][rlang::args_data_masking]> Variables to group by.
#'
#' For `tally()`/`add_tally()`: These dots are for future extensions and must
#' be empty.
#' @param wt <[`data-masking`][rlang::args_data_masking]> Frequency weights.
#' Can be `NULL` or a variable:
#'
Expand Down Expand Up @@ -103,7 +106,7 @@ count.data.frame <- function(
}

wt <- compat_wt(enquo(wt))
out <- tally(out, wt = !!wt, sort = sort, name = name)
out <- tally_dispatch(out, wt = !!wt, sort = sort, name = name)

# Ensure grouping is transient
out <- dplyr_reconstruct(out, x)
Expand All @@ -113,15 +116,34 @@ count.data.frame <- function(

#' @export
#' @rdname count
tally <- function(x, wt = NULL, sort = FALSE, name = NULL) {
tally <- function(x, ..., wt = NULL, sort = FALSE, name = NULL) {
dplyr_local_error_call()

result <- check_tally_dots(
...,
wt = enquo(wt),
sort = sort,
name = name,
fn = "tally"
)
tally_dispatch(
x,
wt = !!result$wt,
sort = result$sort,
name = result$name
)
Comment thread
krlmlr marked this conversation as resolved.
}

tally_dispatch <- function(x, wt = NULL, sort = FALSE, name = NULL) {
UseMethod("tally")
}

#' @export
tally.data.frame <- function(x, wt = NULL, sort = FALSE, name = NULL) {
name <- check_n_name(name, group_vars(x))
error_call <- dplyr_error_call()
name <- check_n_name(name, group_vars(x), call = error_call)

dplyr_local_error_call()
dplyr_local_error_call(error_call)

wt <- compat_wt(enquo(wt))
n <- tally_n(x, wt)
Expand Down Expand Up @@ -211,16 +233,35 @@ add_count_impl <- function(
}

wt <- compat_wt(enquo(wt), env = error_call, user_env = user_env)
add_tally(out, wt = !!wt, sort = sort, name = name)
add_tally_impl(out, wt = !!wt, sort = sort, name = name)
}

#' @rdname count
#' @export
add_tally <- function(x, wt = NULL, sort = FALSE, name = NULL) {
name <- check_n_name(name, tbl_vars(x))

add_tally <- function(x, ..., wt = NULL, sort = FALSE, name = NULL) {
dplyr_local_error_call()

result <- check_tally_dots(
...,
wt = enquo(wt),
sort = sort,
name = name,
fn = "add_tally"
)
add_tally_impl(
x,
wt = !!result$wt,
sort = result$sort,
name = result$name
)
}

add_tally_impl <- function(x, wt = NULL, sort = FALSE, name = NULL) {
error_call <- dplyr_error_call()
name <- check_n_name(name, tbl_vars(x), call = error_call)

dplyr_local_error_call(error_call)

wt <- compat_wt(enquo(wt))
n <- tally_n(x, wt)
out <- mutate(x, !!name := !!n)
Expand All @@ -234,6 +275,77 @@ add_tally <- function(x, wt = NULL, sort = FALSE, name = NULL) {

# Helpers -----------------------------------------------------------------

check_tally_dots <- function(
...,
wt,
sort,
name,
fn,
error_call = caller_env()
) {
if (...length() == 0L) {
return(list(wt = wt, sort = sort, name = name))

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

wt uses NSE.

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed — wt is now handled via enquo(wt) inside collect_tally_args(), preserving the NSE quosure properly. See 8286969.

}

collected <- collect_tally_args(
...,
.fn = fn,
.error_call = error_call
)

if ("wt" %in% names(collected)) {
wt <- collected$wt
}
if ("sort" %in% names(collected)) {
sort <- collected$sort
}
if ("name" %in% names(collected)) {
name <- collected$name
}

list(wt = wt, sort = sort, name = name)
}

collect_tally_args <- function(
wt = NULL,
sort = FALSE,
name = NULL,
...,
.fn,
.error_call = caller_env()
) {
if (...length() > 0L) {
cli::cli_abort(
"Extra arguments passed to `{(.fn)}()` via `...`.",
call = .error_call
)
}

mc <- match.call()
result <- list()

for (arg in c("wt", "sort", "name")) {
if (arg %in% names(mc)) {
lifecycle::deprecate_warn(
when = "1.3.0",
what = I(
glue("Passing `{arg}` as an unnamed argument to `{(.fn)}()`")
),
with = I(glue("`{(.fn)}({arg} = )`")),
env = .error_call,
user_env = caller_env(2)
)
if (arg == "wt") {
result$wt <- enquo(wt)
} else {
result[[arg]] <- get(arg)
}
}
}

result
}

tally_n <- function(x, wt) {
if (quo_is_null(wt)) {
expr(dplyr::n())
Expand Down
76 changes: 76 additions & 0 deletions tests/testthat/_snaps/count-tally.md
Original file line number Diff line number Diff line change
Expand Up @@ -102,6 +102,50 @@
<int>
1 5

# tally() warns when passing unnamed wt

Code
out <- tally(df, x)
Condition
Warning:
Passing `wt` as an unnamed argument to `tally()` was deprecated in dplyr 1.3.0.
i Please use `tally(wt = )` instead.

# tally() warns when passing multiple unnamed args

Code
out <- tally(gf, NULL, TRUE)
Condition
Warning:
Passing `wt` as an unnamed argument to `tally()` was deprecated in dplyr 1.3.0.
i Please use `tally(wt = )` instead.
Warning:
Passing `sort` as an unnamed argument to `tally()` was deprecated in dplyr 1.3.0.
i Please use `tally(sort = )` instead.

# tally() warns when passing all unnamed args

Code
out <- tally(df, x, FALSE, "count")
Condition
Warning:
Passing `wt` as an unnamed argument to `tally()` was deprecated in dplyr 1.3.0.
i Please use `tally(wt = )` instead.
Warning:
Passing `sort` as an unnamed argument to `tally()` was deprecated in dplyr 1.3.0.
i Please use `tally(sort = )` instead.
Warning:
Passing `name` as an unnamed argument to `tally()` was deprecated in dplyr 1.3.0.
i Please use `tally(name = )` instead.

# tally() errors with extra args in dots

Code
tally(df, 1, 2, 3, 4)
Condition
Error in `tally()`:
! Extra arguments passed to `tally()` via `...`.

# `.drop` is defunct

Code
Expand Down Expand Up @@ -176,3 +220,35 @@
4 4 5
5 5 5

# add_tally() warns when passing unnamed wt

Code
out <- add_tally(df, x)
Condition
Warning:
Passing `wt` as an unnamed argument to `add_tally()` was deprecated in dplyr 1.3.0.
i Please use `add_tally(wt = )` instead.

# add_tally() warns when passing all unnamed args

Code
out <- add_tally(df, x, FALSE, "count")
Condition
Warning:
Passing `wt` as an unnamed argument to `add_tally()` was deprecated in dplyr 1.3.0.
i Please use `add_tally(wt = )` instead.
Warning:
Passing `sort` as an unnamed argument to `add_tally()` was deprecated in dplyr 1.3.0.
i Please use `add_tally(sort = )` instead.
Warning:
Passing `name` as an unnamed argument to `add_tally()` was deprecated in dplyr 1.3.0.
i Please use `add_tally(name = )` instead.

# add_tally() errors with extra args in dots

Code
add_tally(df, 1, 2, 3, 4)
Condition
Error in `add_tally()`:
! Extra arguments passed to `add_tally()` via `...`.

69 changes: 68 additions & 1 deletion tests/testthat/test-count-tally.R
Original file line number Diff line number Diff line change
Expand Up @@ -158,7 +158,7 @@ test_that("tally can sort output", {

test_that("weighted tally drops NAs (#1145)", {
df <- tibble(x = c(1, 1, NA))
expect_equal(tally(df, x)$n, 2)
expect_equal(tally(df, wt = x)$n, 2)
})

test_that("tally() drops last group (#5199) ", {
Expand All @@ -182,6 +182,44 @@ test_that("tally() `wt = n()` is deprecated", {
})
})

test_that("tally() works with all named args", {
df <- tibble(x = c(1, 1, NA))
out <- tally(df, wt = x, sort = FALSE, name = "count")
expect_equal(out$count, 2)
})

test_that("tally() warns when passing unnamed wt", {
df <- tibble(x = c(1, 1, NA))
expect_snapshot({
out <- tally(df, x)
})
expect_equal(out$n, 2)
})

test_that("tally() warns when passing multiple unnamed args", {
df <- tibble(x = c(2, 1, 1))
gf <- group_by(df, x)
expect_snapshot({
out <- tally(gf, NULL, TRUE)
})
expect_equal(out$x, c(1, 2))
})

test_that("tally() warns when passing all unnamed args", {
df <- tibble(x = c(1, 1, NA))
expect_snapshot({
out <- tally(df, x, FALSE, "count")
})
expect_equal(out$count, 2)
})

test_that("tally() errors with extra args in dots", {
df <- tibble(x = 1)
expect_snapshot(error = TRUE, {
tally(df, 1, 2, 3, 4)
})
})

# add_count ---------------------------------------------------------------

test_that("add_count preserves grouping", {
Expand Down Expand Up @@ -253,3 +291,32 @@ test_that("add_tally() `wt = n()` is deprecated", {
add_tally(df, wt = n())
})
})

test_that("add_tally() works with all named args", {
df <- tibble(x = c(1, 1, NA))
out <- add_tally(df, wt = x, sort = FALSE, name = "count")
expect_equal(out$count, c(2, 2, 2))
})

test_that("add_tally() warns when passing unnamed wt", {
df <- tibble(x = c(1, 1, NA))
expect_snapshot({
out <- add_tally(df, x)
})
expect_equal(out$n, c(2, 2, 2))
})

test_that("add_tally() warns when passing all unnamed args", {
df <- tibble(x = c(1, 1, NA))
expect_snapshot({
out <- add_tally(df, x, FALSE, "count")
})
expect_equal(out$count, c(2, 2, 2))
})

test_that("add_tally() errors with extra args in dots", {
df <- tibble(x = 1)
expect_snapshot(error = TRUE, {
add_tally(df, 1, 2, 3, 4)
})
})
Loading