Vendor dependencies

This commit is contained in:
2026-08-01 16:11:49 +03:00
parent 7f139a0241
commit 6b5e7f0f8b
29706 changed files with 9575646 additions and 0 deletions
File diff suppressed because one or more lines are too long
+6
View File
@@ -0,0 +1,6 @@
{
"git": {
"sha1": "f6e8e897ab207443d46eebb4c2ccd94b1f91a164"
},
"path_in_vcs": "mea"
}
+349
View File
@@ -0,0 +1,349 @@
# This file is automatically @generated by Cargo.
# It is not intended for manual editing.
version = 4
[[package]]
name = "bitflags"
version = "2.10.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "812e12b5285cc515a9c72a5c1d3b6d46a19dac5acfef5265968c166106e31dd3"
[[package]]
name = "bytes"
version = "1.11.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b35204fbdc0b3f4446b89fc1ac2cf84a8a68971995d0bf2e925ec7cd960f9cb3"
[[package]]
name = "cfg-if"
version = "1.0.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801"
[[package]]
name = "errno"
version = "0.3.14"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb"
dependencies = [
"libc",
"windows-sys 0.61.2",
]
[[package]]
name = "futures-core"
version = "0.3.31"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "05f29059c0c2090612e8d742178b0580d2dc940c837851ad723096f87af6663e"
[[package]]
name = "libc"
version = "0.2.180"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "bcc35a38544a891a5f7c865aca548a982ccb3b8650a5b06d0fd33a10283c56fc"
[[package]]
name = "lock_api"
version = "0.4.14"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "224399e74b87b5f3557511d98dff8b14089b3dadafcab6bb93eab67d3aace965"
dependencies = [
"scopeguard",
]
[[package]]
name = "mea"
version = "0.6.5"
dependencies = [
"pollster",
"slab",
"tokio",
"tokio-test",
]
[[package]]
name = "mio"
version = "1.1.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a69bcab0ad47271a0234d9422b131806bf3968021e5dc9328caf2d4cd58557fc"
dependencies = [
"libc",
"wasi",
"windows-sys 0.61.2",
]
[[package]]
name = "parking_lot"
version = "0.12.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "93857453250e3077bd71ff98b6a65ea6621a19bb0f559a85248955ac12c45a1a"
dependencies = [
"lock_api",
"parking_lot_core",
]
[[package]]
name = "parking_lot_core"
version = "0.9.12"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2621685985a2ebf1c516881c026032ac7deafcda1a2c9b7850dc81e3dfcb64c1"
dependencies = [
"cfg-if",
"libc",
"redox_syscall",
"smallvec",
"windows-link",
]
[[package]]
name = "pin-project-lite"
version = "0.2.16"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3b3cff922bd51709b605d9ead9aa71031d81447142d828eb4a6eba76fe619f9b"
[[package]]
name = "pollster"
version = "0.4.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2f3a9f18d041e6d0e102a0a46750538147e5e8992d3b4873aaafee2520b00ce3"
dependencies = [
"pollster-macro",
]
[[package]]
name = "pollster-macro"
version = "0.4.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ac5da421106a50887c5b51d20806867db377fbb86bacf478ee0500a912e0c113"
dependencies = [
"proc-macro2",
"quote",
"syn",
]
[[package]]
name = "proc-macro2"
version = "1.0.105"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "535d180e0ecab6268a3e718bb9fd44db66bbbc256257165fc699dadf70d16fe7"
dependencies = [
"unicode-ident",
]
[[package]]
name = "quote"
version = "1.0.43"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "dc74d9a594b72ae6656596548f56f667211f8a97b3d4c3d467150794690dc40a"
dependencies = [
"proc-macro2",
]
[[package]]
name = "redox_syscall"
version = "0.5.18"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ed2bf2547551a7053d6fdfafda3f938979645c44812fbfcda098faae3f1a362d"
dependencies = [
"bitflags",
]
[[package]]
name = "scopeguard"
version = "1.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49"
[[package]]
name = "signal-hook-registry"
version = "1.4.8"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c4db69cba1110affc0e9f7bcd48bbf87b3f4fc7c61fc9155afd4c469eb3d6c1b"
dependencies = [
"errno",
"libc",
]
[[package]]
name = "slab"
version = "0.4.11"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7a2ae44ef20feb57a68b23d846850f861394c2e02dc425a50098ae8c90267589"
[[package]]
name = "smallvec"
version = "1.15.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "67b1b7a3b5fe4f1376887184045fcf45c69e92af734b7aaddc05fb777b6fbd03"
[[package]]
name = "socket2"
version = "0.6.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "17129e116933cf371d018bb80ae557e889637989d8638274fb25622827b03881"
dependencies = [
"libc",
"windows-sys 0.60.2",
]
[[package]]
name = "syn"
version = "2.0.114"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d4d107df263a3013ef9b1879b0df87d706ff80f65a86ea879bd9c31f9b307c2a"
dependencies = [
"proc-macro2",
"quote",
"unicode-ident",
]
[[package]]
name = "tokio"
version = "1.49.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "72a2903cd7736441aac9df9d7688bd0ce48edccaadf181c3b90be801e81d3d86"
dependencies = [
"bytes",
"libc",
"mio",
"parking_lot",
"pin-project-lite",
"signal-hook-registry",
"socket2",
"tokio-macros",
"windows-sys 0.61.2",
]
[[package]]
name = "tokio-macros"
version = "2.6.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "af407857209536a95c8e56f8231ef2c2e2aff839b22e07a1ffcbc617e9db9fa5"
dependencies = [
"proc-macro2",
"quote",
"syn",
]
[[package]]
name = "tokio-stream"
version = "0.1.18"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "32da49809aab5c3bc678af03902d4ccddea2a87d028d86392a4b1560c6906c70"
dependencies = [
"futures-core",
"pin-project-lite",
"tokio",
]
[[package]]
name = "tokio-test"
version = "0.4.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3f6d24790a10a7af737693a3e8f1d03faef7e6ca0cc99aae5066f533766de545"
dependencies = [
"futures-core",
"tokio",
"tokio-stream",
]
[[package]]
name = "unicode-ident"
version = "1.0.22"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9312f7c4f6ff9069b165498234ce8be658059c6728633667c526e27dc2cf1df5"
[[package]]
name = "wasi"
version = "0.11.1+wasi-snapshot-preview1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b"
[[package]]
name = "windows-link"
version = "0.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5"
[[package]]
name = "windows-sys"
version = "0.60.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f2f500e4d28234f72040990ec9d39e3a6b950f9f22d3dba18416c35882612bcb"
dependencies = [
"windows-targets",
]
[[package]]
name = "windows-sys"
version = "0.61.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ae137229bcbd6cdf0f7b80a31df61766145077ddf49416a728b02cb3921ff3fc"
dependencies = [
"windows-link",
]
[[package]]
name = "windows-targets"
version = "0.53.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "4945f9f551b88e0d65f3db0bc25c33b8acea4d9e41163edf90dcd0b19f9069f3"
dependencies = [
"windows-link",
"windows_aarch64_gnullvm",
"windows_aarch64_msvc",
"windows_i686_gnu",
"windows_i686_gnullvm",
"windows_i686_msvc",
"windows_x86_64_gnu",
"windows_x86_64_gnullvm",
"windows_x86_64_msvc",
]
[[package]]
name = "windows_aarch64_gnullvm"
version = "0.53.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a9d8416fa8b42f5c947f8482c43e7d89e73a173cead56d044f6a56104a6d1b53"
[[package]]
name = "windows_aarch64_msvc"
version = "0.53.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b9d782e804c2f632e395708e99a94275910eb9100b2114651e04744e9b125006"
[[package]]
name = "windows_i686_gnu"
version = "0.53.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "960e6da069d81e09becb0ca57a65220ddff016ff2d6af6a223cf372a506593a3"
[[package]]
name = "windows_i686_gnullvm"
version = "0.53.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "fa7359d10048f68ab8b09fa71c3daccfb0e9b559aed648a8f95469c27057180c"
[[package]]
name = "windows_i686_msvc"
version = "0.53.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1e7ac75179f18232fe9c285163565a57ef8d3c89254a30685b57d83a38d326c2"
[[package]]
name = "windows_x86_64_gnu"
version = "0.53.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9c3842cdd74a865a8066ab39c8a7a473c0778a3f29370b5fd6b4b9aa7df4a499"
[[package]]
name = "windows_x86_64_gnullvm"
version = "0.53.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "0ffa179e2d07eee8ad8f57493436566c7cc30ac536a3379fdf008f47f6bb7ae1"
[[package]]
name = "windows_x86_64_msvc"
version = "0.53.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d6bbff5f0aada427a1e5a6da5f1f98158182f26556f345ac9e04d36d0ebed650"
+71
View File
@@ -0,0 +1,71 @@
# THIS FILE IS AUTOMATICALLY GENERATED BY CARGO
#
# When uploading crates to the registry Cargo will automatically
# "normalize" Cargo.toml files for maximal compatibility
# with all versions of Cargo and also rewrite `path` dependencies
# to registry (e.g., crates.io) dependencies.
#
# If you are reading this file be aware that the original Cargo.toml
# will likely look very different (and much more reasonable).
# See Cargo.toml.orig for the original contents.
[package]
edition = "2024"
rust-version = "1.85.0"
name = "mea"
version = "0.6.5"
build = false
autolib = false
autobins = false
autoexamples = false
autotests = false
autobenches = false
description = "A runtime-agnostic library providing essential synchronization primitives for asynchronous Rust programming."
homepage = "https://github.com/fast/mea"
documentation = "https://docs.rs/mea"
readme = "README.md"
keywords = [
"async",
"concurrency",
"synchronization",
"waitgroup",
"mutex",
]
categories = [
"asynchronous",
"concurrency",
]
license = "Apache-2.0"
repository = "https://github.com/fast/mea"
resolver = "2"
[package.metadata.docs.rs]
all-features = true
rustdoc-args = [
"--cfg",
"docsrs",
]
[lib]
name = "mea"
path = "src/lib.rs"
[dependencies.slab]
version = "0.4.11"
[dev-dependencies.pollster]
version = "0.4.0"
features = ["macro"]
[dev-dependencies.tokio]
version = "1.41.0"
features = ["full"]
[dev-dependencies.tokio-test]
version = "0.4.4"
[lints.clippy]
dbg_macro = "deny"
[lints.rust]
unknown_lints = "deny"
+44
View File
@@ -0,0 +1,44 @@
# Copyright 2024 tison <wander4096@gmail.com>
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
[package]
name = "mea"
version = "0.6.5"
categories = ["asynchronous", "concurrency"]
description = "A runtime-agnostic library providing essential synchronization primitives for asynchronous Rust programming."
documentation = "https://docs.rs/mea"
keywords = ["async", "concurrency", "synchronization", "waitgroup", "mutex"]
edition.workspace = true
homepage.workspace = true
license.workspace = true
readme.workspace = true
repository.workspace = true
rust-version.workspace = true
[package.metadata.docs.rs]
all-features = true
rustdoc-args = ["--cfg", "docsrs"]
[dependencies]
slab = { version = "0.4.11" }
[dev-dependencies]
pollster = { version = "0.4.0", features = ["macro"] }
tokio = { version = "1.41.0", features = ["full"] }
tokio-test = { version = "0.4.4" }
[lints]
workspace = true
+201
View File
@@ -0,0 +1,201 @@
Apache License
Version 2.0, January 2004
http://www.apache.org/licenses/
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
1. Definitions.
"License" shall mean the terms and conditions for use, reproduction,
and distribution as defined by Sections 1 through 9 of this document.
"Licensor" shall mean the copyright owner or entity authorized by
the copyright owner that is granting the License.
"Legal Entity" shall mean the union of the acting entity and all
other entities that control, are controlled by, or are under common
control with that entity. For the purposes of this definition,
"control" means (i) the power, direct or indirect, to cause the
direction or management of such entity, whether by contract or
otherwise, or (ii) ownership of fifty percent (50%) or more of the
outstanding shares, or (iii) beneficial ownership of such entity.
"You" (or "Your") shall mean an individual or Legal Entity
exercising permissions granted by this License.
"Source" form shall mean the preferred form for making modifications,
including but not limited to software source code, documentation
source, and configuration files.
"Object" form shall mean any form resulting from mechanical
transformation or translation of a Source form, including but
not limited to compiled object code, generated documentation,
and conversions to other media types.
"Work" shall mean the work of authorship, whether in Source or
Object form, made available under the License, as indicated by a
copyright notice that is included in or attached to the work
(an example is provided in the Appendix below).
"Derivative Works" shall mean any work, whether in Source or Object
form, that is based on (or derived from) the Work and for which the
editorial revisions, annotations, elaborations, or other modifications
represent, as a whole, an original work of authorship. For the purposes
of this License, Derivative Works shall not include works that remain
separable from, or merely link (or bind by name) to the interfaces of,
the Work and Derivative Works thereof.
"Contribution" shall mean any work of authorship, including
the original version of the Work and any modifications or additions
to that Work or Derivative Works thereof, that is intentionally
submitted to Licensor for inclusion in the Work by the copyright owner
or by an individual or Legal Entity authorized to submit on behalf of
the copyright owner. For the purposes of this definition, "submitted"
means any form of electronic, verbal, or written communication sent
to the Licensor or its representatives, including but not limited to
communication on electronic mailing lists, source code control systems,
and issue tracking systems that are managed by, or on behalf of, the
Licensor for the purpose of discussing and improving the Work, but
excluding communication that is conspicuously marked or otherwise
designated in writing by the copyright owner as "Not a Contribution."
"Contributor" shall mean Licensor and any individual or Legal Entity
on behalf of whom a Contribution has been received by Licensor and
subsequently incorporated within the Work.
2. Grant of Copyright License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
copyright license to reproduce, prepare Derivative Works of,
publicly display, publicly perform, sublicense, and distribute the
Work and such Derivative Works in Source or Object form.
3. Grant of Patent License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
(except as stated in this section) patent license to make, have made,
use, offer to sell, sell, import, and otherwise transfer the Work,
where such license applies only to those patent claims licensable
by such Contributor that are necessarily infringed by their
Contribution(s) alone or by combination of their Contribution(s)
with the Work to which such Contribution(s) was submitted. If You
institute patent litigation against any entity (including a
cross-claim or counterclaim in a lawsuit) alleging that the Work
or a Contribution incorporated within the Work constitutes direct
or contributory patent infringement, then any patent licenses
granted to You under this License for that Work shall terminate
as of the date such litigation is filed.
4. Redistribution. You may reproduce and distribute copies of the
Work or Derivative Works thereof in any medium, with or without
modifications, and in Source or Object form, provided that You
meet the following conditions:
(a) You must give any other recipients of the Work or
Derivative Works a copy of this License; and
(b) You must cause any modified files to carry prominent notices
stating that You changed the files; and
(c) You must retain, in the Source form of any Derivative Works
that You distribute, all copyright, patent, trademark, and
attribution notices from the Source form of the Work,
excluding those notices that do not pertain to any part of
the Derivative Works; and
(d) If the Work includes a "NOTICE" text file as part of its
distribution, then any Derivative Works that You distribute must
include a readable copy of the attribution notices contained
within such NOTICE file, excluding those notices that do not
pertain to any part of the Derivative Works, in at least one
of the following places: within a NOTICE text file distributed
as part of the Derivative Works; within the Source form or
documentation, if provided along with the Derivative Works; or,
within a display generated by the Derivative Works, if and
wherever such third-party notices normally appear. The contents
of the NOTICE file are for informational purposes only and
do not modify the License. You may add Your own attribution
notices within Derivative Works that You distribute, alongside
or as an addendum to the NOTICE text from the Work, provided
that such additional attribution notices cannot be construed
as modifying the License.
You may add Your own copyright statement to Your modifications and
may provide additional or different license terms and conditions
for use, reproduction, or distribution of Your modifications, or
for any such Derivative Works as a whole, provided Your use,
reproduction, and distribution of the Work otherwise complies with
the conditions stated in this License.
5. Submission of Contributions. Unless You explicitly state otherwise,
any Contribution intentionally submitted for inclusion in the Work
by You to the Licensor shall be under the terms and conditions of
this License, without any additional terms or conditions.
Notwithstanding the above, nothing herein shall supersede or modify
the terms of any separate license agreement you may have executed
with Licensor regarding such Contributions.
6. Trademarks. This License does not grant permission to use the trade
names, trademarks, service marks, or product names of the Licensor,
except as required for reasonable and customary use in describing the
origin of the Work and reproducing the content of the NOTICE file.
7. Disclaimer of Warranty. Unless required by applicable law or
agreed to in writing, Licensor provides the Work (and each
Contributor provides its Contributions) on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
implied, including, without limitation, any warranties or conditions
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
PARTICULAR PURPOSE. You are solely responsible for determining the
appropriateness of using or redistributing the Work and assume any
risks associated with Your exercise of permissions under this License.
8. Limitation of Liability. In no event and under no legal theory,
whether in tort (including negligence), contract, or otherwise,
unless required by applicable law (such as deliberate and grossly
negligent acts) or agreed to in writing, shall any Contributor be
liable to You for damages, including any direct, indirect, special,
incidental, or consequential damages of any character arising as a
result of this License or out of the use or inability to use the
Work (including but not limited to damages for loss of goodwill,
work stoppage, computer failure or malfunction, or any and all
other commercial damages or losses), even if such Contributor
has been advised of the possibility of such damages.
9. Accepting Warranty or Additional Liability. While redistributing
the Work or Derivative Works thereof, You may choose to offer,
and charge a fee for, acceptance of support, warranty, indemnity,
or other liability obligations and/or rights consistent with this
License. However, in accepting such obligations, You may act only
on Your own behalf and on Your sole responsibility, not on behalf
of any other Contributor, and only if You agree to indemnify,
defend, and hold each Contributor harmless for any liability
incurred by, or claims asserted against, such Contributor by reason
of your accepting any such warranty or additional liability.
END OF TERMS AND CONDITIONS
APPENDIX: How to apply the Apache License to your work.
To apply the Apache License to your work, attach the following
boilerplate notice, with the fields enclosed by brackets "[]"
replaced with your own identifying information. (Don't include
the brackets!) The text should be enclosed in the appropriate
comment syntax for the file format. We also recommend that a
file or class name and description of purpose be included on the
same "printed page" as the copyright notice for easier
identification within third-party archives.
Copyright [yyyy] [name of copyright owner]
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
+90
View File
@@ -0,0 +1,90 @@
# MEA (Make Easy Async)
[![Crates.io][crates-badge]][crates-url]
[![Documentation][docs-badge]][docs-url]
[![MSRV 1.85][msrv-badge]](https://www.whatrustisit.com)
[![Apache 2.0 licensed][license-badge]][license-url]
[![Build Status][actions-badge]][actions-url]
[crates-badge]: https://img.shields.io/crates/v/mea.svg
[crates-url]: https://crates.io/crates/mea
[docs-badge]: https://docs.rs/mea/badge.svg
[docs-url]: https://docs.rs/mea
[msrv-badge]: https://img.shields.io/badge/MSRV-1.85-green?logo=rust
[license-badge]: https://img.shields.io/crates/l/mea
[license-url]: LICENSE
[actions-badge]: https://github.com/fast/mea/actions/workflows/ci.yml/badge.svg
[actions-url]: https://github.com/fast/mea/actions/workflows/ci.yml
## Overview
MEA is a runtime-agnostic library providing essential synchronization primitives for asynchronous Rust programming. The library offers a collection of well-tested, efficient synchronization tools that work with any async runtime.
## Features
* [**admission::FairShare**](https://docs.rs/mea/*/mea/admission/struct.FairShare.html): A work-conserving admission policy that fairly shares bounded concurrency across keys.
* [**Barrier**](https://docs.rs/mea/*/mea/barrier/struct.Barrier.html): A synchronization primitive that enables tasks to wait until all participants arrive.
* [**Condvar**](https://docs.rs/mea/*/mea/condvar/struct.Condvar.html): A condition variable that allows tasks to wait for a notification.
* [**Latch**](https://docs.rs/mea/*/mea/latch/struct.Latch.html): A synchronization primitive that allows one or more tasks to wait until a set of operations completes.
* [**Mutex**](https://docs.rs/mea/*/mea/mutex/struct.Mutex.html): A mutual exclusion primitive for protecting shared data.
* [**Once**](https://docs.rs/mea/*/mea/once/struct.Once.html): A primitive that ensures a one-time asynchronous operation runs at most once, even when called concurrently.
* [**OnceCell**](https://docs.rs/mea/*/mea/once/struct.OnceCell.html): A cell that can be written to at most once, providing safe, lazy initialization.
* [**OnceMap**](https://docs.rs/mea/*/mea/once/struct.OnceMap.html): A hash map that runs computation only once for each key and stores the result.
* [**RwLock**](https://docs.rs/mea/*/mea/rwlock/struct.RwLock.html): A reader-writer lock that allows multiple readers or a single writer at a time.
* [**Semaphore**](https://docs.rs/mea/*/mea/semaphore/struct.Semaphore.html): A synchronization primitive that controls access to a shared resource.
* [**ShutdownSend, ShutdownRecv & ShutdownWatch**](https://docs.rs/mea/*/mea/shutdown/): A composite synchronization primitive for managing shutdown signals.
* [**WaitGroup**](https://docs.rs/mea/*/mea/waitgroup/struct.WaitGroup.html): A synchronization primitive that allows waiting for multiple tasks to complete.
* [**atomicbox**](https://docs.rs/mea/*/mea/atomicbox/): A safe, owning version of AtomicPtr for heap-allocated data.
* [**broadcast**](https://docs.rs/mea/*/mea/broadcast/): A multi-producer, multi-consumer broadcast channel.
* [**mpsc::bounded**](https://docs.rs/mea/*/mea/mpsc/fn.bounded.html): A multi-producer, single-consumer bounded queue for sending values between asynchronous tasks.
* [**mpsc::unbounded**](https://docs.rs/mea/*/mea/mpsc/fn.unbounded.html): A multi-producer, single-consumer unbounded queue for sending values between asynchronous tasks.
* [**oneshot::channel**](https://docs.rs/mea/*/mea/oneshot/): A one-shot channel for sending a single value between tasks.
* [**singleflight::Group**](https://docs.rs/mea/*/mea/singleflight/): A duplicate function call suppression mechanism.
## Installation
Add the dependency to your `Cargo.toml` via:
```shell
cargo add mea
```
## Runtime Agnostic
All synchronization primitives in this library are runtime-agnostic, meaning they can be used with any async runtime like Tokio, async-std, or others. This makes the library highly versatile and portable.
## Thread Safety
All types in this library implement `Send` and `Sync`, making them safe to share across thread boundaries. This is essential for concurrent programming where data needs to be accessed from multiple threads.
## Minimum Supported Rust Version (MSRV)
This crate is built against the latest stable release, and its minimum supported rustc version is 1.85.0.
The policy is that the minimum Rust version required to use this crate can be increased in minor version updates. For example, if MEA 1.0 requires Rust 1.20.0, then MEA 1.0.z for all values of z will also require Rust 1.20.0 or newer. However, Mea 1.y for y > 0 may require a newer minimum version of Rust.
## License
This project is licensed under [Apache License, Version 2.0](LICENSE).
## History
This crate collects runtime-agnostic synchronization primitives from spare parts:
* **admission::FairShare** is written from scratch to bound global concurrency while balancing held permits across contending keys.
* **Barrier** is inspired by `std::sync::Barrier` and `tokio::sync::Barrier`, with a different implementation based on the internal `WaitSet` primitive.
* **Condvar** is inspired by `std::sync::Condvar` and `async_std::sync::Condvar`, with a different implementation based on the internal `Semaphore` primitive. Different from the async_std implementation, this condvar is fair.
* **Latch** is inspired by [`latches`](https://github.com/mirromutth/latches), with a different implementation based on the internal `CountdownState` primitive. No `wait` or `watch` method is provided, since it can be easily implemented by [composing delay futures](https://docs.rs/fastimer/*/fastimer/fn.timeout.html). No sync variant is provided, since it can be easily implemented with block_on of any runtime.
* **Mutex** is derived from `tokio::sync::Mutex`. No blocking method is provided, since it can be easily implemented with block_on of any runtime.
* **OnceCell** is derived from `tokio::sync::OnceCell`, but using our own semaphore implementation.
* **OnceMap** is inspired by `uv-once-map` but the interface and implementation are redesigned.
* **RwLock** is derived from `tokio::sync::RwLock`, but the `max_readers` can be any `NonZeroUsize` (effectively any positive `usize`) instead of `[0, u32::MAX >> 3]`. No blocking method is provided, since it can be easily implemented with block_on of any runtime.
* **Semaphore** is derived from `tokio::sync::Semaphore`, without `close` method since it is quite tricky to use. And thus, this semaphore doesn't have the limitation of max permits. Besides, new methods like `forget_exact` are added to fit the specific use case.
* **WaitGroup** is inspired by [`waitgroup-rs`](https://github.com/laizy/waitgroup-rs), providing different API flavor with a different implementation based on the internal `CountdownState` primitive.
* **atomicbox** is forked from [`atomicbox`](https://github.com/jorendorff/atomicbox/) at commit 07756444.
* **broadcast::channel** is derived from `tokio::sync::broadcast::channel`, with a different implementation based on the internal `WaitSet` primitive.
* **oneshot::channel** is derived from [`oneshot`](https://github.com/faern/oneshot), with significant simplifications since we need not support synchronized receiving functions.
Other parts are written from scratch.
NB. The optimization considerations are different when implementing a sync primitive for sync code and async code. Generally speaking, once you have an async + runtime-agnostic implementation, you can immediately have a sync implementation by block_on any async runtime ([`pollster`](https://github.com/zesterer/pollster) is the most lightweight runtime that park the current thread). However, a sync-oriented implementation may leverage some platform-specific features to achieve better performance. This library is designed for async code, so it doesn't consider sync-oriented optimization. I often find libraries that try to provide both sync and async implementations end up with a clumsy API design. So I prefer to keep them separate.
+538
View File
@@ -0,0 +1,538 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::collections::HashMap;
use std::collections::VecDeque;
use std::future::Future;
use std::hash::BuildHasher;
use std::hash::Hash;
use std::hash::RandomState;
use std::pin::Pin;
use std::sync::Arc;
use std::task::Context;
use std::task::Poll;
use std::task::Waker;
use slab::Slab;
use crate::internal::Mutex;
/// An admission controller that fairly shares a fixed number of permits across keys.
///
/// Each acquisition belongs to a key. When a permit becomes available,
/// [`FairShare`] admits a queued acquisition for the key with the fewest
/// permits currently held. Ties are resolved by queue order.
///
/// See the [module-level documentation](super) for details about the fairness
/// guarantee.
#[derive(Debug)]
pub struct FairShare<K, S = RandomState>
where
K: Eq + Hash,
S: BuildHasher,
{
state: Mutex<State<K, S>>,
}
impl<K> FairShare<K, RandomState>
where
K: Eq + Hash,
{
/// Creates a fair-share admission controller with the given number of permits.
///
/// # Panics
///
/// Panics if `permits` is zero.
///
/// # Examples
///
/// ```
/// use mea::admission::FairShare;
///
/// let admission = FairShare::<String>::new(3);
/// assert_eq!(admission.available_permits(), 3);
/// ```
pub fn new(permits: usize) -> Self {
Self::with_hasher(permits, RandomState::new())
}
}
impl<K, S> FairShare<K, S>
where
K: Eq + Hash,
S: BuildHasher,
{
/// Creates a fair-share admission controller with the given number of
/// permits and hash builder.
///
/// # Panics
///
/// Panics if `permits` is zero.
pub fn with_hasher(permits: usize, hash_builder: S) -> Self {
assert!(permits > 0, "FairShare requires at least one permit");
Self {
state: Mutex::new(State::new(permits, hash_builder)),
}
}
/// Returns the current number of permits available for immediate admission.
///
/// A permit already assigned to a queued acquisition counts as held by its
/// key, even if that acquisition has not yet been polled again.
pub fn available_permits(&self) -> usize {
self.state.lock().available_permits
}
/// Returns the number of acquisitions currently waiting for a permit.
///
/// An acquisition is no longer counted once it has been assigned a permit,
/// even if its future has not yet been polled again.
pub fn num_waiters(&self) -> usize {
self.state.lock().num_waiters
}
/// Attempts to acquire one permit for `key` without waiting.
///
/// This method does not bypass queued acquisitions.
pub fn try_acquire(&self, key: K) -> Option<FairSharePermit<'_, K, S>> {
let key = Arc::new(key);
let admitted = self.state.lock().try_admit(key.clone());
admitted.then(|| FairSharePermit {
admission: self,
key,
})
}
/// Acquires one permit for `key`.
///
/// # Cancel safety
///
/// Cancelling this method loses the acquisition's place in the queue. If
/// a permit has already been assigned, cancellation releases it for another
/// queued acquisition.
pub async fn acquire(&self, key: K) -> FairSharePermit<'_, K, S> {
let key = Arc::new(key);
Acquire::new(self, key.clone()).await;
FairSharePermit {
admission: self,
key,
}
}
/// Attempts to acquire one owned permit for `key` without waiting.
///
/// The admission controller must be wrapped in an [`Arc`] to call this
/// method.
pub fn try_acquire_owned(self: Arc<Self>, key: K) -> Option<OwnedFairSharePermit<K, S>> {
let key = Arc::new(key);
let admitted = self.state.lock().try_admit(key.clone());
admitted.then(|| OwnedFairSharePermit {
admission: self,
key,
})
}
/// Acquires one owned permit for `key`.
///
/// The admission controller must be wrapped in an [`Arc`] to call this
/// method.
///
/// # Cancel safety
///
/// This method has the same cancellation behavior as [`Self::acquire`].
pub async fn acquire_owned(self: Arc<Self>, key: K) -> OwnedFairSharePermit<K, S> {
let key = Arc::new(key);
Acquire::new(&self, key.clone()).await;
OwnedFairSharePermit {
admission: self,
key,
}
}
fn release(&self, key: &K) {
let mut wakers = Vec::new();
{
let mut state = self.state.lock();
state.release(key);
state.admit_waiters(&mut wakers);
}
wake_all(wakers);
}
}
#[derive(Debug)]
struct State<K, S>
where
K: Eq + Hash,
S: BuildHasher,
{
total_permits: usize,
available_permits: usize,
num_waiters: usize,
next_sequence: u64,
groups: HashMap<Arc<K>, GroupState, S>,
waiters: Slab<Waiter>,
}
impl<K, S> State<K, S>
where
K: Eq + Hash,
S: BuildHasher,
{
fn new(permits: usize, hash_builder: S) -> Self {
Self {
total_permits: permits,
available_permits: permits,
num_waiters: 0,
next_sequence: 0,
groups: HashMap::with_hasher(hash_builder),
waiters: Slab::new(),
}
}
fn try_admit(&mut self, key: Arc<K>) -> bool {
if self.available_permits == 0 || self.num_waiters != 0 {
return false;
}
self.available_permits -= 1;
self.groups.entry(key).or_default().held_permits += 1;
true
}
fn enqueue(&mut self, key: Arc<K>, waker: &Waker) -> usize {
let sequence = self.next_sequence;
self.next_sequence += 1;
let waiter = self.waiters.insert(Waiter {
sequence,
waker: Some(waker.clone()),
admitted: false,
});
self.groups.entry(key).or_default().queue.push_back(waiter);
self.num_waiters += 1;
waiter
}
fn poll_waiter(&mut self, waiter: usize, waker: &Waker) -> Poll<()> {
let state = self
.waiters
.get_mut(waiter)
.expect("FairShare waiter is missing");
if state.admitted {
self.waiters.remove(waiter);
Poll::Ready(())
} else {
if state
.waker
.as_ref()
.is_none_or(|current| !current.will_wake(waker))
{
state.waker = Some(waker.clone());
}
Poll::Pending
}
}
fn cancel(&mut self, waiter_id: usize, key: &K) {
let waiter = self.waiters.remove(waiter_id);
if waiter.admitted {
self.release(key);
return;
}
let remove_group = {
let group = self
.groups
.get_mut(key)
.expect("FairShare waiter group is missing");
let position = group
.queue
.iter()
.position(|candidate| *candidate == waiter_id)
.expect("FairShare waiter is missing from its group");
group.queue.remove(position);
group.held_permits == 0 && group.queue.is_empty()
};
self.num_waiters -= 1;
if remove_group {
self.groups.remove(key);
}
}
fn admit_waiters(&mut self, wakers: &mut Vec<Waker>) {
while self.available_permits > 0 && self.num_waiters > 0 {
let key = self
.next_group()
.expect("FairShare has pending acquisitions without a group");
let waiter = self.groups[&key]
.queue
.front()
.copied()
.expect("FairShare pending group has no waiters");
{
let group = self
.groups
.get_mut(&key)
.expect("FairShare pending group is missing");
let popped = group.queue.pop_front();
debug_assert_eq!(popped, Some(waiter));
group.held_permits += 1;
}
self.available_permits -= 1;
self.num_waiters -= 1;
let waiter = &mut self.waiters[waiter];
waiter.admitted = true;
if let Some(waker) = waiter.waker.take() {
wakers.push(waker);
}
}
}
fn next_group(&self) -> Option<Arc<K>> {
self.groups
.iter()
.filter_map(|(key, group)| {
let waiter = *group.queue.front()?;
let sequence = self.waiters[waiter].sequence;
Some((group.held_permits, sequence, key))
})
.min_by_key(|(held_permits, sequence, _)| (*held_permits, *sequence))
.map(|(_, _, key)| key.clone())
}
fn release(&mut self, key: &K) {
let remove_group = {
let group = self
.groups
.get_mut(key)
.expect("FairShare released a permit for an unknown key");
debug_assert!(group.held_permits > 0);
group.held_permits -= 1;
group.held_permits == 0 && group.queue.is_empty()
};
if remove_group {
self.groups.remove(key);
}
self.available_permits += 1;
debug_assert!(self.available_permits <= self.total_permits);
}
}
#[derive(Debug, Default)]
struct GroupState {
held_permits: usize,
queue: VecDeque<usize>,
}
#[derive(Debug)]
struct Waiter {
sequence: u64,
waker: Option<Waker>,
admitted: bool,
}
#[derive(Debug)]
struct Acquire<'a, K, S>
where
K: Eq + Hash,
S: BuildHasher,
{
admission: &'a FairShare<K, S>,
key: Arc<K>,
waiter: Option<usize>,
completed: bool,
}
impl<'a, K, S> Acquire<'a, K, S>
where
K: Eq + Hash,
S: BuildHasher,
{
fn new(admission: &'a FairShare<K, S>, key: Arc<K>) -> Self {
Self {
admission,
key,
waiter: None,
completed: false,
}
}
}
impl<K, S> Drop for Acquire<'_, K, S>
where
K: Eq + Hash,
S: BuildHasher,
{
fn drop(&mut self) {
let Some(waiter) = self.waiter.take() else {
return;
};
let mut wakers = Vec::new();
{
let mut state = self.admission.state.lock();
state.cancel(waiter, &self.key);
state.admit_waiters(&mut wakers);
}
wake_all(wakers);
}
}
impl<K, S> Future for Acquire<'_, K, S>
where
K: Eq + Hash,
S: BuildHasher,
{
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.get_mut();
if this.completed {
return Poll::Ready(());
}
if let Some(waiter) = this.waiter {
if this
.admission
.state
.lock()
.poll_waiter(waiter, cx.waker())
.is_ready()
{
this.waiter = None;
this.completed = true;
return Poll::Ready(());
}
return Poll::Pending;
}
let mut wakers = Vec::new();
let ready = {
let mut state = this.admission.state.lock();
if state.try_admit(this.key.clone()) {
this.completed = true;
return Poll::Ready(());
}
let waiter = state.enqueue(this.key.clone(), cx.waker());
this.waiter = Some(waiter);
state.admit_waiters(&mut wakers);
state.poll_waiter(waiter, cx.waker()).is_ready()
};
wake_all(wakers);
if ready {
this.waiter = None;
this.completed = true;
Poll::Ready(())
} else {
Poll::Pending
}
}
}
/// A permit from a [`FairShare`] admission controller.
///
/// This type is created by the [`acquire`] and [`try_acquire`] methods on
/// [`FairShare`]. It represents one admitted operation associated with a key.
/// Dropping it returns the permit and may admit another queued acquisition.
///
/// [`acquire`]: FairShare::acquire
/// [`try_acquire`]: FairShare::try_acquire
#[must_use = "permits are released immediately when dropped"]
#[derive(Debug)]
pub struct FairSharePermit<'a, K, S = RandomState>
where
K: Eq + Hash,
S: BuildHasher,
{
admission: &'a FairShare<K, S>,
key: Arc<K>,
}
impl<K, S> FairSharePermit<'_, K, S>
where
K: Eq + Hash,
S: BuildHasher,
{
/// Returns the key associated with this permit.
pub fn key(&self) -> &K {
&self.key
}
}
impl<K, S> Drop for FairSharePermit<'_, K, S>
where
K: Eq + Hash,
S: BuildHasher,
{
fn drop(&mut self) {
self.admission.release(&self.key);
}
}
/// An owned permit from a [`FairShare`] admission controller.
///
/// This type is created by the [`acquire_owned`] and [`try_acquire_owned`]
/// methods on [`FairShare`]. Unlike [`FairSharePermit`], it owns an [`Arc`] to
/// the admission controller and has no lifetime parameter. Dropping it returns
/// the permit and may admit another queued acquisition.
///
/// [`acquire_owned`]: FairShare::acquire_owned
/// [`try_acquire_owned`]: FairShare::try_acquire_owned
#[must_use = "permits are released immediately when dropped"]
#[derive(Debug)]
pub struct OwnedFairSharePermit<K, S = RandomState>
where
K: Eq + Hash,
S: BuildHasher,
{
admission: Arc<FairShare<K, S>>,
key: Arc<K>,
}
impl<K, S> OwnedFairSharePermit<K, S>
where
K: Eq + Hash,
S: BuildHasher,
{
/// Returns the key associated with this permit.
pub fn key(&self) -> &K {
&self.key
}
}
impl<K, S> Drop for OwnedFairSharePermit<K, S>
where
K: Eq + Hash,
S: BuildHasher,
{
fn drop(&mut self) {
self.admission.release(&self.key);
}
}
fn wake_all(wakers: Vec<Waker>) {
for waker in wakers {
waker.wake();
}
}
+32
View File
@@ -0,0 +1,32 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//! Admission control policies for bounded asynchronous work.
//!
//! This module provides [`FairShare`], a work-conserving admission policy for
//! workloads partitioned by key. It maintains a fixed number of permits and,
//! when contended, admits work for the key with the fewest permits currently
//! held. Ties are resolved by queue order.
//!
//! Fairness applies to the number of permits held by contending keys. It does
//! not reserve permits for idle keys or account for differences in execution
//! time or work cost.
mod fair_share;
#[cfg(test)]
mod tests;
pub use fair_share::FairShare;
pub use fair_share::FairSharePermit;
pub use fair_share::OwnedFairSharePermit;
+334
View File
@@ -0,0 +1,334 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::collections::hash_map::DefaultHasher;
use std::future::Future;
use std::hash::BuildHasherDefault;
use std::pin::Pin;
use std::pin::pin;
use std::sync::Arc;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::task::Context;
use std::task::Poll;
use std::task::Waker;
use super::FairShare;
fn poll_once<F>(future: Pin<&mut F>) -> Poll<F::Output>
where
F: Future,
{
future.poll(&mut Context::from_waker(Waker::noop()))
}
#[test]
#[should_panic(expected = "FairShare requires at least one permit")]
fn zero_permits_panics() {
FairShare::<usize>::new(0);
}
#[test]
fn tracks_available_permits() {
let admission = FairShare::new(2);
assert_eq!(admission.available_permits(), 2);
let permit_a0 = admission.try_acquire("a").unwrap();
let permit_a1 = admission.try_acquire("a").unwrap();
assert_eq!(permit_a0.key(), &"a");
assert_eq!(admission.available_permits(), 0);
assert!(admission.try_acquire("b").is_none());
drop(permit_a0);
assert_eq!(admission.available_permits(), 1);
drop(permit_a1);
assert_eq!(admission.available_permits(), 2);
}
#[test]
fn uses_all_permits_without_reservations() {
let admission = FairShare::new(3);
let permits = [
admission.try_acquire("a").unwrap(),
admission.try_acquire("a").unwrap(),
admission.try_acquire("a").unwrap(),
];
assert_eq!(admission.available_permits(), 0);
drop(permits);
assert_eq!(admission.available_permits(), 3);
}
#[test]
fn admits_the_key_with_the_smallest_share() {
let admission = FairShare::new(2);
let permit_a0 = admission.try_acquire("a").unwrap();
let permit_a1 = admission.try_acquire("a").unwrap();
let acquire_a = admission.acquire("a");
let mut acquire_a = pin!(acquire_a);
assert!(poll_once(acquire_a.as_mut()).is_pending());
let acquire_b = admission.acquire("b");
let mut acquire_b = pin!(acquire_b);
assert!(poll_once(acquire_b.as_mut()).is_pending());
drop(permit_a0);
assert!(poll_once(acquire_a.as_mut()).is_pending());
let permit_b = match poll_once(acquire_b.as_mut()) {
Poll::Ready(permit) => permit,
Poll::Pending => panic!("key b should receive the released permit"),
};
assert_eq!(permit_b.key(), &"b");
drop(permit_a1);
let permit_a = match poll_once(acquire_a.as_mut()) {
Poll::Ready(permit) => permit,
Poll::Pending => panic!("key a should receive the next permit"),
};
assert_eq!(permit_a.key(), &"a");
}
#[test]
fn shares_permits_across_contending_keys() {
let admission = FairShare::new(3);
let mut held_by_a = vec![
admission.try_acquire("a").unwrap(),
admission.try_acquire("a").unwrap(),
admission.try_acquire("a").unwrap(),
];
let acquire_a = admission.acquire("a");
let mut acquire_a = pin!(acquire_a);
assert!(poll_once(acquire_a.as_mut()).is_pending());
let acquire_b = admission.acquire("b");
let mut acquire_b = pin!(acquire_b);
assert!(poll_once(acquire_b.as_mut()).is_pending());
let acquire_c = admission.acquire("c");
let mut acquire_c = pin!(acquire_c);
assert!(poll_once(acquire_c.as_mut()).is_pending());
drop(held_by_a.pop().unwrap());
let permit_b = match poll_once(acquire_b.as_mut()) {
Poll::Ready(permit) => permit,
Poll::Pending => panic!("key b should receive the first released permit"),
};
drop(held_by_a.pop().unwrap());
let permit_c = match poll_once(acquire_c.as_mut()) {
Poll::Ready(permit) => permit,
Poll::Pending => panic!("key c should receive the second released permit"),
};
drop(held_by_a.pop().unwrap());
let permit_a = match poll_once(acquire_a.as_mut()) {
Poll::Ready(permit) => permit,
Poll::Pending => panic!("key a should receive the third released permit"),
};
assert_eq!(permit_a.key(), &"a");
assert_eq!(permit_b.key(), &"b");
assert_eq!(permit_c.key(), &"c");
assert_eq!(admission.available_permits(), 0);
drop((permit_a, permit_b, permit_c));
assert_eq!(admission.available_permits(), 3);
}
#[test]
fn breaks_equal_share_ties_by_queue_order() {
let admission = FairShare::new(1);
let held = admission.try_acquire("held").unwrap();
let acquire_b = admission.acquire("b");
let mut acquire_b = pin!(acquire_b);
assert!(poll_once(acquire_b.as_mut()).is_pending());
let acquire_a = admission.acquire("a");
let mut acquire_a = pin!(acquire_a);
assert!(poll_once(acquire_a.as_mut()).is_pending());
drop(held);
let permit_b = match poll_once(acquire_b.as_mut()) {
Poll::Ready(permit) => permit,
Poll::Pending => panic!("the first queued acquisition should win an equal-share tie"),
};
assert!(poll_once(acquire_a.as_mut()).is_pending());
drop(permit_b);
let permit_a = match poll_once(acquire_a.as_mut()) {
Poll::Ready(permit) => permit,
Poll::Pending => panic!("the second queued acquisition should be admitted next"),
};
assert_eq!(permit_a.key(), &"a");
}
#[test]
fn preserves_queue_order_within_a_key() {
let admission = FairShare::new(1);
let held = admission.try_acquire(7usize).unwrap();
let first = admission.acquire(7usize);
let mut first = pin!(first);
assert!(poll_once(first.as_mut()).is_pending());
let second = admission.acquire(7usize);
let mut second = pin!(second);
assert!(poll_once(second.as_mut()).is_pending());
drop(held);
let first_permit = match poll_once(first.as_mut()) {
Poll::Ready(permit) => permit,
Poll::Pending => panic!("the first acquisition should be admitted first"),
};
assert!(poll_once(second.as_mut()).is_pending());
drop(first_permit);
let second_permit = match poll_once(second.as_mut()) {
Poll::Ready(permit) => permit,
Poll::Pending => panic!("the second acquisition should be admitted second"),
};
assert_eq!(second_permit.key(), &7);
}
#[test]
fn cancelling_a_pending_acquire_removes_it() {
let admission = FairShare::new(1);
let held = admission.try_acquire(1usize).unwrap();
{
let acquire = admission.acquire(2usize);
let mut acquire = pin!(acquire);
assert!(poll_once(acquire.as_mut()).is_pending());
}
drop(held);
assert_eq!(admission.available_permits(), 1);
let permit = admission.try_acquire(3usize).unwrap();
assert_eq!(permit.key(), &3);
}
#[test]
fn cancelling_an_admitted_acquire_reassigns_its_permit() {
let admission = FairShare::new(1);
let held = admission.try_acquire("held").unwrap();
let mut first = Box::pin(admission.acquire("first"));
assert!(poll_once(first.as_mut()).is_pending());
let mut second = Box::pin(admission.acquire("second"));
assert!(poll_once(second.as_mut()).is_pending());
drop(held);
assert_eq!(admission.available_permits(), 0);
drop(first);
assert_eq!(admission.available_permits(), 0);
let permit = match poll_once(second.as_mut()) {
Poll::Ready(permit) => permit,
Poll::Pending => panic!("cancellation should reassign the granted permit"),
};
assert_eq!(permit.key(), &"second");
}
#[test]
fn cancelling_within_a_key_preserves_its_queue() {
let admission = FairShare::new(1);
let held = admission.try_acquire("held").unwrap();
let mut first = Box::pin(admission.acquire("tenant"));
assert!(poll_once(first.as_mut()).is_pending());
let mut second = Box::pin(admission.acquire("tenant"));
assert!(poll_once(second.as_mut()).is_pending());
drop(first);
drop(held);
let permit = match poll_once(second.as_mut()) {
Poll::Ready(permit) => permit,
Poll::Pending => panic!("cancelling one acquisition must not detach the next"),
};
assert_eq!(permit.key(), &"tenant");
}
#[test]
fn supports_a_custom_hash_builder() {
let admission = FairShare::<String, BuildHasherDefault<DefaultHasher>>::with_hasher(
1,
BuildHasherDefault::default(),
);
let permit = admission.try_acquire("tenant".to_owned()).unwrap();
assert_eq!(permit.key(), "tenant");
}
#[test]
fn owned_permit_keeps_the_admission_controller_alive() {
let admission = Arc::new(FairShare::new(1));
let permit = admission
.clone()
.try_acquire_owned("tenant")
.expect("a permit should be available");
drop(admission);
assert_eq!(permit.key(), &"tenant");
drop(permit);
}
#[test]
fn acquire_futures_are_send() {
fn assert_send<T: Send>(_: T) {}
let admission = FairShare::<String>::new(1);
assert_send(admission.acquire("tenant".to_owned()));
let admission = Arc::new(FairShare::<String>::new(1));
assert_send(admission.acquire_owned("tenant".to_owned()));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn stress_test_preserves_permit_limit() {
let admission = Arc::new(FairShare::new(3));
let active = Arc::new(AtomicUsize::new(0));
let max_active = Arc::new(AtomicUsize::new(0));
let mut handles = Vec::new();
for key in 0..5usize {
for _ in 0..32usize {
let admission = admission.clone();
let active = active.clone();
let max_active = max_active.clone();
handles.push(tokio::spawn(async move {
let _permit = admission.acquire_owned(key).await;
let now = active.fetch_add(1, Ordering::SeqCst) + 1;
max_active.fetch_max(now, Ordering::SeqCst);
tokio::task::yield_now().await;
active.fetch_sub(1, Ordering::SeqCst);
}));
}
}
for handle in handles {
handle.await.unwrap();
}
assert_eq!(active.load(Ordering::SeqCst), 0);
assert!(max_active.load(Ordering::SeqCst) <= 3);
assert_eq!(admission.available_permits(), 3);
}
+298
View File
@@ -0,0 +1,298 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
// This is derived from https://github.com/jorendorff/atomicbox/blob/07756444/src/atomic_box.rs.
use std::fmt;
use std::marker::PhantomData;
use std::mem;
use std::ptr;
use std::sync::atomic::AtomicPtr;
use std::sync::atomic::Ordering;
/// A type that holds a single `Box<T>` value and can be safely shared between threads.
pub struct AtomicBox<T> {
ptr: AtomicPtr<T>,
/// This effectively makes `AtomicBox<T>` non-`Send` and non-`Sync` if `T`
/// is non-`Send`.
phantom: PhantomData<Box<T>>,
}
/// Mark `AtomicBox<T>` as safe to share across threads.
///
/// This is safe because shared access to an `AtomicBox<T>` does not provide
/// shared access to any `T` value. However, it does provide the ability to get
/// a `Box<T>` from another thread, so `T: Send` is required.
unsafe impl<T> Sync for AtomicBox<T> where T: Send {}
impl<T> AtomicBox<T> {
/// Creates a new `AtomicBox` with the given value.
///
/// # Examples
///
/// ```rust
/// use mea::atomicbox::AtomicBox;
///
/// let atomic_box = AtomicBox::new(Box::new(0));
/// ```
pub fn new(value: Box<T>) -> AtomicBox<T> {
AtomicBox {
ptr: AtomicPtr::new(Box::into_raw(value)),
phantom: PhantomData,
}
}
/// Atomically set this `AtomicBox` to `other` and return the previous value.
///
/// This does not allocate or free memory, and it neither clones nor drops
/// any values. `other` is moved into `self`; the value previously in
/// `self` is returned.
///
/// # Examples
///
/// ```rust
/// use mea::atomicbox::AtomicBox;
///
/// let atom = AtomicBox::new(Box::new("one"));
/// let prev_value = atom.swap(Box::new("two"));
/// assert_eq!(*prev_value, "one");
/// ```
pub fn swap(&self, other: Box<T>) -> Box<T> {
let mut result = other;
self.swap_mut(&mut result);
result
}
/// Atomically set this `AtomicBox` to `other` and drop its previous value.
///
/// The `AtomicBox` takes ownership of `other`.
///
/// # Examples
///
/// ```rust
/// use mea::atomicbox::AtomicBox;
///
/// let atom = AtomicBox::new(Box::new("one"));
/// atom.store(Box::new("two"));
/// assert_eq!(atom.into_inner(), Box::new("two"));
/// ```
pub fn store(&self, other: Box<T>) {
self.swap(other);
}
/// Atomically swaps the contents of this `AtomicBox` with the contents of `other`.
///
/// This does not allocate or free memory, and it neither clones nor drops
/// any values. The pointers in `*other` and `self` are simply exchanged.
///
/// # Examples
///
/// ```rust
/// use mea::atomicbox::AtomicBox;
///
/// let atom = AtomicBox::new(Box::new("one"));
/// let mut boxed = Box::new("two");
/// atom.swap_mut(&mut boxed);
/// assert_eq!(*boxed, "one");
/// ```
pub fn swap_mut(&self, other: &mut Box<T>) {
let other_ptr = Box::into_raw(unsafe { ptr::read(other) });
let ptr = self.ptr.swap(other_ptr, Ordering::AcqRel);
unsafe { ptr::write(other, Box::from_raw(ptr)) };
}
/// Consume this `AtomicBox`, returning the last box value it contained.
///
/// # Examples
///
/// ```rust
/// use mea::atomicbox::AtomicBox;
///
/// let atom = AtomicBox::new(Box::new("hello"));
/// assert_eq!(atom.into_inner(), Box::new("hello"));
/// ```
pub fn into_inner(mut self) -> Box<T> {
let result = unsafe { Box::from_raw(*self.ptr.get_mut()) };
mem::forget(self);
result
}
/// Returns a mutable reference to the contained value.
///
/// This is safe because it borrows the `AtomicBox` mutably, which ensures
/// that no other threads can concurrently access either the atomic pointer field
/// or the boxed data it points to.
pub fn get_mut(&mut self) -> &mut T {
unsafe { &mut **self.ptr.get_mut() }
}
}
impl<T> Drop for AtomicBox<T> {
/// Dropping an `AtomicBox<T>` drops the final `Box<T>` value stored in it.
fn drop(&mut self) {
let ptr = *self.ptr.get_mut();
unsafe { drop(Box::from_raw(ptr)) }
}
}
impl<T> Default for AtomicBox<T>
where
Box<T>: Default,
{
/// The default `AtomicBox<T>` value boxes the default `T` value.
fn default() -> AtomicBox<T> {
AtomicBox::new(Default::default())
}
}
impl<T> fmt::Debug for AtomicBox<T> {
/// The `{:?}` format of an `AtomicBox<T>` looks like `"AtomicBox(0x12341234)"`.
/// The address is the address of the `Box` allocation, not the address of
/// the `AtomicBox`.
fn fmt(&self, f: &mut fmt::Formatter) -> Result<(), fmt::Error> {
let p = self.ptr.load(Ordering::Relaxed);
f.write_str("AtomicBox(")?;
fmt::Pointer::fmt(&p, f)?;
f.write_str(")")?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::sync::Barrier;
use super::*;
#[test]
fn atomic_box_swap_works() {
let b = AtomicBox::new(Box::new("hello world"));
let bis = Box::new("bis");
assert_eq!(b.swap(bis), Box::new("hello world"));
assert_eq!(b.swap(Box::new("")), Box::new("bis"));
}
#[test]
fn atomic_box_store_works() {
let b = AtomicBox::new(Box::new("hello world"));
let bis = Box::new("bis");
b.store(bis);
assert_eq!(b.into_inner(), Box::new("bis"));
}
#[test]
fn atomic_box_swap_mut_works() {
let b = AtomicBox::new(Box::new("hello world"));
let mut bis = Box::new("bis");
b.swap_mut(&mut bis);
assert_eq!(bis, Box::new("hello world"));
b.swap_mut(&mut bis);
assert_eq!(bis, Box::new("bis"));
}
#[test]
fn atomic_box_pointer_identity() {
let box1 = Box::new(1);
let p1 = format!("{box1:p}");
let atom = AtomicBox::new(box1);
let box2 = Box::new(2);
let p2 = format!("{box2:p}");
assert_ne!(p2, p1);
let box3 = atom.swap(box2); // box1 out, box2 in
let p3 = format!("{box3:p}");
assert_eq!(p3, p1); // box3 is box1
let box4 = atom.swap(Box::new(5)); // box2 out, throwaway value in
let p4 = format!("{box4:p}");
assert_eq!(p4, p2); // box4 is box2
}
#[test]
fn atomic_box_drops() {
use std::sync::Arc;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
struct K(Arc<AtomicUsize>, usize);
impl Drop for K {
fn drop(&mut self) {
self.0.fetch_add(self.1, Ordering::Relaxed);
}
}
let n = Arc::new(AtomicUsize::new(0));
{
let ab = AtomicBox::new(Box::new(K(n.clone(), 5)));
assert_eq!(n.load(Ordering::Relaxed), 0);
let first = ab.swap(Box::new(K(n.clone(), 13)));
assert_eq!(n.load(Ordering::Relaxed), 0);
drop(first);
assert_eq!(n.load(Ordering::Relaxed), 5);
}
assert_eq!(n.load(Ordering::Relaxed), 5 + 13);
}
#[test]
fn atomic_threads() {
const NTHREADS: usize = 9;
let gate = Arc::new(Barrier::new(NTHREADS));
let abox: Arc<AtomicBox<Vec<u8>>> = Arc::new(Default::default());
let handles: Vec<_> = (0..NTHREADS as u8)
.map(|t| {
let my_gate = gate.clone();
let my_box = abox.clone();
std::thread::spawn(move || {
my_gate.wait();
let mut my_vec = Box::new(vec![]);
for _ in 0..100 {
my_vec = my_box.swap(my_vec);
my_vec.push(t);
}
my_vec
})
})
.collect();
let mut counts = [0usize; NTHREADS];
for h in handles {
for val in *h.join().unwrap() {
counts[val as usize] += 1;
}
}
// Don't forget the data still in `abox`!
// There are NTHREADS+1 vectors in all.
for val in *abox.swap(Box::new(vec![])) {
counts[val as usize] += 1;
}
println!("{counts:?}");
for count in counts {
assert_eq!(count, 100);
}
}
#[test]
fn debug_fmt() {
let my_box = Box::new(32);
let expected = format!("AtomicBox({my_box:p})");
assert_eq!(format!("{:?}", AtomicBox::new(my_box)), expected);
}
}
@@ -0,0 +1,333 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
// This is derived from https://github.com/jorendorff/atomicbox/blob/07756444/src/atomic_option_box.rs.
use std::fmt;
use std::marker::PhantomData;
use std::mem;
use std::ptr;
use std::sync::atomic::AtomicPtr;
use std::sync::atomic::Ordering;
/// A type that holds a single `Option<Box<T>>` value and can be safely shared
/// between threads.
pub struct AtomicOptionBox<T> {
/// Pointer to a `T` value in the heap, representing `Some(t)`;
/// or a null pointer for `None`.
ptr: AtomicPtr<T>,
/// This effectively makes `AtomicOptionBox<T>` non-`Send` and non-`Sync`
/// if `T` is non-`Send`.
phantom: PhantomData<Box<T>>,
}
/// Mark `AtomicOptionBox<T>` as safe to share across threads.
///
/// This is safe because shared access to an `AtomicOptionBox<T>` does not
/// provide shared access to any `T` value. However, it does provide the
/// ability to get a `Box<T>` from another thread, so `T: Send` is required.
unsafe impl<T> Sync for AtomicOptionBox<T> where T: Send {}
fn into_ptr<T>(value: Option<Box<T>>) -> *mut T {
match value {
Some(box_value) => Box::into_raw(box_value),
None => ptr::null_mut(),
}
}
// SAFETY: The caller must ensure that `ptr` was obtained from `Box::into_raw` or is null.
unsafe fn from_ptr<T>(ptr: *mut T) -> Option<Box<T>> {
if ptr.is_null() {
None
} else {
Some(unsafe { Box::from_raw(ptr) })
}
}
impl<T> AtomicOptionBox<T> {
/// Creates a new `AtomicOptionBox` with the given value.
///
/// # Examples
///
/// ```rust
/// use mea::atomicbox::AtomicOptionBox;
///
/// let atomic_box = AtomicOptionBox::new(Some(Box::new(0)));
/// ```
pub fn new(value: Option<Box<T>>) -> AtomicOptionBox<T> {
AtomicOptionBox {
ptr: AtomicPtr::new(into_ptr(value)),
phantom: PhantomData,
}
}
/// Creates a new `AtomicOptionBox` with no value.
///
/// Equivalent to `AtomicOptionBox::new(None)`, but can be used in `const` context.
///
/// # Examples
///
/// ```rust
/// use mea::atomicbox::AtomicOptionBox;
///
/// static GLOBAL_BOX: AtomicOptionBox<u32> = AtomicOptionBox::none();
/// ```
pub const fn none() -> Self {
Self {
ptr: AtomicPtr::new(ptr::null_mut()),
phantom: PhantomData,
}
}
/// Atomically set this `AtomicOptionBox` to `other` and return the previous value.
///
/// This does not allocate or free memory, and it neither clones nor drops any values. `other`
/// is moved into `self`; the value previously in `self` is returned.
///
/// # Examples
///
/// ```rust
/// use mea::atomicbox::AtomicOptionBox;
///
/// let atom = AtomicOptionBox::new(None);
/// let prev_value = atom.swap(Some(Box::new("ok")));
/// assert_eq!(prev_value, None);
/// ```
pub fn swap(&self, other: Option<Box<T>>) -> Option<Box<T>> {
let order = match other {
Some(_) => Ordering::AcqRel,
None => Ordering::Acquire,
};
let new_ptr = into_ptr(other);
let old_ptr = self.ptr.swap(new_ptr, order);
unsafe { from_ptr(old_ptr) }
}
/// Atomically set this `AtomicOptionBox` to `other` and drop the previous value.
///
/// # Examples
///
/// ```rust
/// use mea::atomicbox::AtomicOptionBox;
///
/// let atom = AtomicOptionBox::new(None);
/// atom.store(Some(Box::new("ok")));
/// assert_eq!(atom.into_inner(), Some(Box::new("ok")));
/// ```
pub fn store(&self, other: Option<Box<T>>) {
self.swap(other);
}
/// Atomically set this `AtomicOptionBox` to `None` and return the previous value.
///
/// This does not allocate or free memory, and it neither clones nor drops any values. It is
/// equivalent to calling `self.swap(None)`
///
/// # Examples
///
/// ```rust
/// use mea::atomicbox::AtomicOptionBox;
///
/// let atom = AtomicOptionBox::new(Some(Box::new("ok")));
/// let prev_value = atom.take();
/// assert!(prev_value.is_some());
/// let prev_value = atom.take();
/// assert!(prev_value.is_none());
/// ```
pub fn take(&self) -> Option<Box<T>> {
self.swap(None)
}
/// Atomically swaps the contents of this `AtomicOptionBox` with the contents of `other`.
///
/// This does not allocate or free memory, and it neither clones nor drops any values. The
/// pointers in `*other` and `self` are simply exchanged.
///
/// # Examples
///
/// ```rust
/// use mea::atomicbox::AtomicOptionBox;
///
/// let atom = AtomicOptionBox::new(None);
/// let mut boxed = Some(Box::new("ok"));
/// let prev_value = atom.swap_mut(&mut boxed);
/// assert_eq!(boxed, None);
/// ```
pub fn swap_mut(&self, other: &mut Option<Box<T>>) {
let previous = self.swap(other.take());
*other = previous;
}
/// Consume this `AtomicOptionBox`, returning the last option value it
/// contained.
///
/// # Examples
///
/// ```rust
/// use mea::atomicbox::AtomicOptionBox;
///
/// let atom = AtomicOptionBox::new(Some(Box::new("hello")));
/// assert_eq!(atom.into_inner(), Some(Box::new("hello")));
/// ```
pub fn into_inner(mut self) -> Option<Box<T>> {
let result = unsafe { from_ptr(*self.ptr.get_mut()) };
mem::forget(self);
result
}
/// Returns a mutable reference to the contained value.
///
/// This is safe because it borrows the `AtomicOptionBox` mutably, which
/// ensures that no other threads can concurrently access either the atomic
/// pointer field or the boxed data it points to.
pub fn get_mut(&mut self) -> Option<&mut T> {
unsafe { self.ptr.get_mut().as_mut() }
}
}
impl<T> Drop for AtomicOptionBox<T> {
/// Dropping an `AtomicOptionBox<T>` drops the final `Box<T>` value (if any) stored in it.
fn drop(&mut self) {
let ptr = *self.ptr.get_mut();
unsafe { drop(from_ptr(ptr)) }
}
}
impl<T> Default for AtomicOptionBox<T> {
/// The default `AtomicOptionBox<T>` value is `AtomicBox::new(None)`.
fn default() -> AtomicOptionBox<T> {
AtomicOptionBox::new(None)
}
}
impl<T> fmt::Debug for AtomicOptionBox<T> {
/// The `{:?}` format of an `AtomicOptionBox<T>` looks like
/// `"AtomicOptionBox(0x12341234)"` or `"AtomicOptionBox(None)"`.
///
/// The address is the address of the `Box` allocation, if any, not the
/// address of the `AtomicOptionBox`.
fn fmt(&self, f: &mut fmt::Formatter) -> Result<(), fmt::Error> {
let p = self.ptr.load(Ordering::Relaxed);
f.write_str("AtomicOptionBox(")?;
if p.is_null() {
f.write_str("None")?;
} else {
fmt::Pointer::fmt(&p, f)?;
}
f.write_str(")")?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use core::sync::atomic::Ordering;
use std::sync::Arc;
use std::sync::atomic::AtomicUsize;
use super::*;
#[test]
fn atomic_option_box_swap_works() {
let b = AtomicOptionBox::new(Some(Box::new("hello world")));
let bis = Box::new("bis");
assert_eq!(b.swap(None), Some(Box::new("hello world")));
assert_eq!(b.swap(Some(bis)), None);
assert_eq!(b.swap(None), Some(Box::new("bis")));
}
#[test]
fn atomic_option_box_store_works() {
let b = AtomicOptionBox::new(Some(Box::new("hello world")));
b.store(None);
assert_eq!(b.into_inner(), None);
let b = AtomicOptionBox::new(Some(Box::new("hello world")));
let bis = Box::new("bis");
b.store(Some(bis));
assert_eq!(b.into_inner(), Some(Box::new("bis")));
}
#[test]
fn atomic_option_box_swap_mut_works() {
let b = AtomicOptionBox::new(Some(Box::new("hello world")));
let mut bis = None;
b.swap_mut(&mut bis);
assert_eq!(bis, Some(Box::new("hello world")));
bis = Some(Box::new("bis"));
b.swap_mut(&mut bis);
assert_eq!(bis, None);
b.swap_mut(&mut bis);
assert_eq!(bis, Some(Box::new("bis")));
}
#[test]
fn atomic_option_box_pointer_identity() {
let box1 = Box::new(1);
let p1 = &*box1 as *const i32;
let atom = AtomicOptionBox::new(Some(box1));
let box2 = Box::new(2);
let p2 = &*box2 as *const i32;
assert_ne!(p2, p1);
let box3 = atom.swap(Some(box2)).unwrap(); // box1 out, box2 in
let p3 = &*box3 as *const i32;
assert_eq!(p3, p1); // box3 is box1
let box4 = atom.swap(None).unwrap(); // box2 out, None in
let p4 = &*box4 as *const i32;
assert_eq!(p4, p2); // box4 is box2
}
#[test]
fn atomic_box_drops() {
struct K(Arc<AtomicUsize>, usize);
impl Drop for K {
fn drop(&mut self) {
self.0.fetch_add(self.1, Ordering::Relaxed);
}
}
let n = Arc::new(AtomicUsize::new(0));
{
let ab = AtomicOptionBox::new(Some(Box::new(K(n.clone(), 5))));
assert_eq!(n.load(Ordering::Relaxed), 0);
let first = ab.swap(None);
assert_eq!(n.load(Ordering::Relaxed), 0);
drop(first);
assert_eq!(n.load(Ordering::Relaxed), 5);
let second = ab.swap(Some(Box::new(K(n.clone(), 13))));
assert!(second.is_none());
assert_eq!(n.load(Ordering::Relaxed), 5);
}
assert_eq!(n.load(Ordering::Relaxed), 5 + 13);
}
#[test]
fn debug_fmt() {
let my_box = Box::new(32);
let expected = format!("AtomicOptionBox({my_box:p})");
assert_eq!(
format!("{:?}", AtomicOptionBox::new(Some(my_box))),
expected
);
assert_eq!(
format!("{:?}", AtomicOptionBox::<String>::new(None)),
"AtomicOptionBox(None)"
);
}
}
+43
View File
@@ -0,0 +1,43 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
// This module is derived from https://github.com/jorendorff/atomicbox/.
//! [`AtomicBox`] and [`AtomicOptionBox`] are safe, owning versions of std's [`AtomicPtr`].
//!
//! This can be useful to avoid resource leaks when you forget to call `Box::from_raw` on
//! the pointer stored in `AtomicPtr`.
//!
//! Unfortunately, the only operations you can perform on an atomic box are swaps and stores: you
//! can not just use the box without taking ownership of it. Imagine a `Box` without `Deref` or
//! `DerefMut` implementations, and you can get the idea. Still, this is sufficient for some
//! lock-free data structures, so here it is.
//!
//! ## Why no `Deref`?
//!
//! It would not be safe. The point of an `AtomicBox` is that other threads can obtain the boxed
//! value, take ownership of it, even drop it, all without taking a lock. So there is no safe way to
//! borrow that value—except to swap it out of the `AtomicBox` yourself.
//!
//! This is pretty much the same reason you can not borrow a reference to the contents of any other
//! atomic type. It would invite data races. The only difference here is that those contents happen
//! to be on the heap.
//!
//! [`AtomicPtr`]: std::sync::atomic::AtomicPtr
mod atomic_box;
mod atomic_option_box;
pub use atomic_box::AtomicBox;
pub use atomic_option_box::AtomicOptionBox;
+284
View File
@@ -0,0 +1,284 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//! A synchronization primitive that enables multiple tasks to wait for each other.
//!
//! The barrier ensures that no task proceeds past a certain point until all tasks have reached it.
//! This is useful for scenarios where multiple tasks need to proceed together after reaching a
//! certain point in their execution.
//!
//! A barrier enables multiple tasks to synchronize the beginning of some computation.
//! When a barrier is created, it is initialized with a count of the number of tasks
//! that will synchronize on the barrier. Each task can then call [`wait()`] on the
//! barrier to indicate it is ready to proceed. The barrier ensures that no task
//! proceeds past the barrier point until all tasks have made the call.
//!
//! # Examples
//!
//! ```
//! # #[tokio::main]
//! # async fn main() {
//! use std::sync::Arc;
//!
//! use mea::barrier::Barrier;
//!
//! let barrier = Arc::new(Barrier::new(3));
//! let mut handles = Vec::new();
//!
//! for i in 0..3 {
//! let barrier = barrier.clone();
//! handles.push(tokio::spawn(async move {
//! println!("Task {} before barrier", i);
//! let result = barrier.wait().await;
//! println!("Task {} after barrier (leader: {})", i, result.is_leader());
//! }));
//! }
//!
//! for handle in std::mem::take(&mut handles) {
//! handle.await.unwrap();
//! }
//!
//! // The barrier can be reused (generation is increased).
//! for i in 0..3 {
//! let barrier = barrier.clone();
//! handles.push(tokio::spawn(async move {
//! println!("Task {} before barrier", i);
//! let result = barrier.wait().await;
//! println!("Task {} after barrier (leader: {})", i, result.is_leader());
//! }));
//! }
//!
//! for handle in handles {
//! handle.await.unwrap();
//! }
//! # }
//! ```
//!
//! [`wait()`]: Barrier::wait
use std::fmt;
use std::future::Future;
use std::pin::Pin;
use std::task::Context;
use std::task::Poll;
use crate::internal::Mutex;
use crate::internal::WaitSet;
#[cfg(test)]
mod tests;
/// A synchronization primitive for multiple tasks that need to wait for each other.
///
/// See the [module level documentation](self) for more.
#[derive(Debug)]
pub struct Barrier {
n: u32,
state: Mutex<BarrierState>,
}
struct BarrierState {
arrived: u32,
generation: usize,
waiters: WaitSet,
}
impl fmt::Debug for BarrierState {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("BarrierState")
.field("arrived", &self.arrived)
.field("generation", &self.generation)
.finish_non_exhaustive()
}
}
/// A `BarrierWaitResult` is returned by [`Barrier::wait()`] when all threads
/// in the [`Barrier`] have rendezvoused.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::barrier::Barrier;
///
/// let barrier = Barrier::new(1);
/// let barrier_wait_result = barrier.wait().await;
/// # }
/// ```
pub struct BarrierWaitResult(bool);
impl fmt::Debug for BarrierWaitResult {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("BarrierWaitResult")
.field("is_leader", &self.is_leader())
.finish()
}
}
impl BarrierWaitResult {
/// Returns `true` if this worker is the "leader" for the call to [`Barrier::wait()`].
///
/// Only one worker will have `true` returned from their result, all other
/// workers will have `false` returned.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::barrier::Barrier;
///
/// let barrier = Barrier::new(1);
/// let barrier_wait_result = barrier.wait().await;
/// println!("{:?}", barrier_wait_result.is_leader());
/// # }
/// ```
#[must_use]
pub fn is_leader(&self) -> bool {
self.0
}
}
impl Barrier {
/// Creates a new barrier that can block the specified number of tasks.
///
/// A barrier will block `n-1` tasks and release them all at once when the `n`th task arrives.
///
/// # Arguments
///
/// * `n`: The number of tasks to wait for. If `n` is 0, it will be treated as 1.
///
/// # Examples
///
/// ```
/// use mea::barrier::Barrier;
///
/// let barrier = Barrier::new(3); // Creates a barrier for 3 tasks
/// ```
pub fn new(n: u32) -> Self {
// If n is 0, it's not clear what behavior the user wants.
// std::sync::Barrier works with n = 0 the same as n = 1,
// where every .wait() immediately unblocks, so we adopt that here as well.
let n = if n > 0 { n } else { 1 };
Self {
n,
state: Mutex::new(BarrierState {
arrived: 0,
generation: 0,
waiters: WaitSet::with_capacity(n as usize),
}),
}
}
/// Waits for all tasks to reach this point.
///
/// The barrier will block the current task until all `n` tasks have called `wait()`.
/// The last task to call `wait()` will be designated as the leader and receive `true`
/// as the return value. All other tasks will receive `false`.
///
/// # Returns
///
/// Returns a `Future` that resolves to:
/// * `true` if this task is the last (leader) task to arrive at the barrier
/// * `false` for all other tasks
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::barrier::Barrier;
///
/// let barrier = Arc::new(Barrier::new(2));
/// let barrier2 = barrier.clone();
///
/// let handle = tokio::spawn(async move {
/// let result = barrier2.wait().await;
/// println!("Task 1: leader = {}", result.is_leader());
/// });
///
/// let result = barrier.wait().await;
/// println!("Task 2: leader = {}", result.is_leader());
/// handle.await.unwrap();
/// # }
/// ```
pub async fn wait(&self) -> BarrierWaitResult {
let generation = {
let mut state = self.state.lock();
let generation = state.generation;
state.arrived += 1;
// the last arriver is the leader;
// wake up other waiters, increment the generation, and return
if state.arrived == self.n {
state.arrived = 0;
state.generation += 1;
state.waiters.wake_all();
return BarrierWaitResult(true);
}
generation
};
let fut = BarrierWait {
idx: None,
generation,
barrier: self,
};
fut.await;
BarrierWaitResult(false)
}
}
/// A future returned by [`Barrier::wait()`].
///
/// This future will complete when all tasks have reached the barrier point.
#[must_use = "futures do nothing unless you `.await` or poll them"]
struct BarrierWait<'a> {
idx: Option<usize>,
generation: usize,
barrier: &'a Barrier,
}
impl fmt::Debug for BarrierWait<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("BarrierWait")
.field("generation", &self.generation)
.finish_non_exhaustive()
}
}
impl Future for BarrierWait<'_> {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let Self {
idx,
generation,
barrier,
} = self.get_mut();
let mut state = barrier.state.lock();
if *generation < state.generation {
Poll::Ready(())
} else {
state.waiters.register_waker(idx, cx);
Poll::Pending
}
}
}
+99
View File
@@ -0,0 +1,99 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use tokio_test::assert_pending;
use tokio_test::assert_ready;
use tokio_test::task::spawn;
use crate::barrier::Barrier;
#[test]
fn zero_does_not_block() {
let b = Barrier::new(0);
{
let mut f = spawn(b.wait());
let leader = assert_ready!(f.poll());
assert!(leader.is_leader());
}
{
let mut f = spawn(b.wait());
let leader = assert_ready!(f.poll());
assert!(leader.is_leader());
}
}
#[test]
fn single() {
let b = Barrier::new(1);
{
let mut f = spawn(b.wait());
let leader = assert_ready!(f.poll());
assert!(leader.is_leader());
}
{
let mut f = spawn(b.wait());
let leader = assert_ready!(f.poll());
assert!(leader.is_leader());
}
{
let mut f = spawn(b.wait());
let leader = assert_ready!(f.poll());
assert!(leader.is_leader());
}
}
#[test]
fn tango() {
let b = Barrier::new(2);
let mut f1 = spawn(b.wait());
assert_pending!(f1.poll());
let mut f2 = spawn(b.wait());
let f2_leader = assert_ready!(f2.poll()).is_leader();
let f1_leader = assert_ready!(f1.poll()).is_leader();
assert!(f1_leader || f2_leader);
assert!(!(f1_leader && f2_leader));
}
#[test]
fn lots() {
let b = Barrier::new(100);
for _ in 0..10 {
let mut wait = Vec::new();
for _ in 0..99 {
let mut f = spawn(b.wait());
assert_pending!(f.poll());
wait.push(f);
}
for f in &mut wait {
assert_pending!(f.poll());
}
// pass the barrier
let mut f = spawn(b.wait());
let mut found_leader = assert_ready!(f.poll()).is_leader();
for mut f in wait {
let leader = assert_ready!(f.poll());
if leader.is_leader() {
assert!(!found_leader);
found_leader = true;
}
}
assert!(found_leader);
}
}
+21
View File
@@ -0,0 +1,21 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//! A multi-producer multi-consumer broadcast channel.
//!
//! This module provides broadcast channels in one of the following policies:
//!
//! * [`overflow`]: when the channel is full, the oldest messages are overwritten.
pub mod overflow;
+479
View File
@@ -0,0 +1,479 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//! A multi-producer multi-consumer broadcast channel.
//!
//! This channel supports multiple senders and multiple receivers. Each message sent by any
//! sender is received by all receivers. If a receiver falls behind, it may miss messages,
//! which is reported via [`RecvError::Lagged`].
//!
//! # Examples
//!
//! Basic usage:
//!
//! ```
//! use mea::broadcast::overflow;
//!
//! # #[tokio::main]
//! # async fn main() {
//! let (tx, mut rx1) = overflow::channel(16);
//! let mut rx2 = tx.subscribe();
//!
//! tx.send(10);
//! tx.send(20);
//!
//! assert_eq!(rx1.recv().await, Ok(10));
//! assert_eq!(rx1.recv().await, Ok(20));
//! assert_eq!(rx2.recv().await, Ok(10));
//! assert_eq!(rx2.recv().await, Ok(20));
//! # }
//! ```
//!
//! Handling lag:
//!
//! ```
//! use mea::broadcast::overflow;
//! use mea::broadcast::overflow::RecvError;
//!
//! # #[tokio::main]
//! # async fn main() {
//! let (tx, mut rx) = overflow::channel(2);
//!
//! tx.send(1);
//! tx.send(2);
//! tx.send(3); // overwrites the oldest message (1)
//!
//! assert_eq!(rx.recv().await, Err(RecvError::Lagged(1)));
//! assert_eq!(rx.recv().await, Ok(2));
//! assert_eq!(rx.recv().await, Ok(3));
//! # }
//! ```
use std::fmt;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::AtomicU64;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::task::Context;
use std::task::Poll;
use crate::internal::Mutex;
use crate::internal::RwLock;
use crate::internal::WaitSet;
#[cfg(test)]
mod tests;
/// Creates a new broadcast channel with the given hint `capacity`. The actual capacity may be
/// greater than the provided `capacity`.
///
/// See [module-level documentation](self) for broadcast channel semantics.
///
/// # Panics
///
/// Panics if `capacity` is 0.
///
/// # Examples
///
/// ```
/// use mea::broadcast::overflow;
///
/// let (tx, mut rx) = overflow::channel(16);
/// tx.send(10);
/// assert_eq!(rx.try_recv(), Ok(10));
/// ```
pub fn channel<T: Clone>(capacity: usize) -> (Sender<T>, Receiver<T>) {
assert!(capacity > 0, "capacity must be greater than 0");
let capacity = capacity.next_power_of_two();
let mask = capacity - 1;
let mut buffer = Vec::with_capacity(capacity);
for _ in 0..capacity {
buffer.push(RwLock::new(Slot {
msg: None,
version: 0,
}));
}
let shared = Arc::new(Shared {
buffer: buffer.into_boxed_slice(),
capacity,
mask,
tail_cnt: AtomicU64::new(0),
senders: AtomicUsize::new(1),
waiters: Mutex::new(WaitSet::new()),
});
let sender = Sender {
shared: shared.clone(),
};
let receiver = Receiver { shared, head: 0 };
(sender, receiver)
}
/// Error returned by [`Receiver::recv`].
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RecvError {
/// The receiver lagged too far behind.
///
/// The count is the number of messages skipped. The receiver's internal cursor has been
/// advanced to the oldest available message.
Lagged(u64),
/// The sender has become disconnected, and there will never be any more data received on it.
Disconnected,
}
impl fmt::Display for RecvError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
RecvError::Lagged(n) => write!(f, "receiver has been lagged by {n}"),
RecvError::Disconnected => write!(f, "receiving on a closed channel"),
}
}
}
impl std::error::Error for RecvError {}
/// Error returned by [`Receiver::try_recv`].
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum TryRecvError {
/// This channel is currently empty, but the sender(s) have not yet disconnected, so data may
/// yet become available.
Empty,
/// The receiver lagged too far behind.
///
/// The count is the number of messages skipped. The receiver's internal cursor has been
/// advanced to the oldest available message.
Lagged(u64),
/// The sender has become disconnected, and there will never be any more data received on it.
Disconnected,
}
impl fmt::Display for TryRecvError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
TryRecvError::Empty => write!(f, "receiving on an empty channel"),
TryRecvError::Lagged(n) => write!(f, "receiver has been lagged by {n}"),
TryRecvError::Disconnected => write!(f, "receiving on a closed channel"),
}
}
}
impl std::error::Error for TryRecvError {}
#[derive(Debug)]
struct Slot<T> {
/// The message. `None` if the slot is empty (initial state only).
msg: Option<T>,
/// The absolute version of the message in this slot.
version: u64,
}
struct Shared<T> {
buffer: Box<[RwLock<Slot<T>>]>,
capacity: usize,
mask: usize,
/// The global tail cursor. Points to the next slot to write.
/// Strictly monotonically increasing.
tail_cnt: AtomicU64,
/// Number of active senders.
senders: AtomicUsize,
/// Waiters (receivers) waiting for new messages.
waiters: Mutex<WaitSet>,
}
/// A sender handle to the broadcast channel.
///
/// The sender can be cloned to create multiple producers. When all senders are dropped,
/// the channel is closed.
pub struct Sender<T> {
shared: Arc<Shared<T>>,
}
impl<T> Clone for Sender<T> {
fn clone(&self) -> Self {
self.shared.senders.fetch_add(1, Ordering::Release);
Self {
shared: self.shared.clone(),
}
}
}
impl<T> Drop for Sender<T> {
fn drop(&mut self) {
match self.shared.senders.fetch_sub(1, Ordering::AcqRel) {
1 => {
// If this is the last sender, we need to wake up the receiver so it can
// observe the disconnected state.
self.shared.waiters.lock().wake_all();
}
_ => {
// there are still other senders left, do nothing
}
}
}
}
impl<T> Sender<T> {
/// Broadcasts a value to all active receivers.
///
/// This operation is non-blocking. If the channel buffer is full, the oldest message
/// in the buffer is overwritten. Any receiver that was waiting for that overwritten
/// message will receive a [`RecvError::Lagged`] error on its next call to `recv`.
///
/// # Examples
///
/// ```
/// use mea::broadcast::overflow;
///
/// let (tx, mut rx) = overflow::channel(16);
/// tx.send(10);
/// assert_eq!(rx.try_recv(), Ok(10));
/// ```
pub fn send(&self, msg: T) {
let tail = self.shared.tail_cnt.fetch_add(1, Ordering::SeqCst);
let idx = (tail as usize) & self.shared.mask;
{
let mut slot = self.shared.buffer[idx].write();
slot.msg = Some(msg);
slot.version = tail;
}
// Notify all waiting receivers.
self.shared.waiters.lock().wake_all();
}
/// Creates a new receiver that starts receiving messages from the current tail of the channel.
///
/// # Examples
///
/// ```
/// use mea::broadcast::overflow;
/// use mea::broadcast::overflow::TryRecvError;
///
/// # #[tokio::main]
/// # async fn main() {
/// let (tx, _) = overflow::channel(16);
/// tx.send(10);
///
/// let mut rx = tx.subscribe();
/// assert_eq!(rx.try_recv(), Err(TryRecvError::Empty));
/// tx.send(20);
/// assert_eq!(rx.recv().await, Ok(20));
/// # }
/// ```
pub fn subscribe(&self) -> Receiver<T> {
// Receiver starts at the current tail.
let head = self.shared.tail_cnt.load(Ordering::SeqCst);
let shared = self.shared.clone();
Receiver { shared, head }
}
}
/// A receiver handle to the broadcast channel.
///
/// The receiver can be cloned to create multiple consumers. Each receiver sees every
/// message sent to the channel (unless it lags behind).
pub struct Receiver<T> {
shared: Arc<Shared<T>>,
head: u64,
}
impl<T> Clone for Receiver<T> {
fn clone(&self) -> Self {
Self {
shared: self.shared.clone(),
head: self.head,
}
}
}
impl<T: Clone> Receiver<T> {
/// Receives the next value for this receiver.
///
/// # Returns
///
/// * `Ok(T)`: The next message.
/// * `Err(RecvError::Lagged(u64))`: The receiver lagged behind. The internal cursor is advanced
/// to the oldest available message. The count indicates how many messages were skipped.
/// * `Err(RecvError::Disconnected)`: All senders have been dropped and no more messages are
/// available.
///
/// # Examples
///
/// ```
/// use mea::broadcast::overflow;
///
/// # #[tokio::main]
/// # async fn main() {
/// let (tx, mut rx) = overflow::channel(16);
/// tx.send(10);
/// assert_eq!(rx.recv().await, Ok(10));
/// # }
/// ```
pub async fn recv(&mut self) -> Result<T, RecvError> {
Recv {
receiver: self,
index: None,
}
.await
}
/// Attempts to receive the next value for this receiver without blocking.
///
/// # Returns
///
/// * `Ok(T)`: The next message.
/// * `Err(TryRecvError::Empty)`: No message is currently available.
/// * `Err(TryRecvError::Lagged(u64))`: The receiver lagged behind. The internal cursor is
/// advanced to the oldest available message. The count indicates how many messages were
/// skipped.
/// * `Err(TryRecvError::Disconnected)`: All senders have been dropped and no more messages are
/// available.
///
/// # Examples
///
/// ```
/// use mea::broadcast::overflow;
///
/// let (tx, mut rx) = overflow::channel(16);
/// tx.send(10);
/// assert_eq!(rx.try_recv(), Ok(10));
/// ```
pub fn try_recv(&mut self) -> Result<T, TryRecvError> {
let shared = &self.shared;
let cap = shared.capacity as u64;
let tail = shared.tail_cnt.load(Ordering::SeqCst);
let head = self.head;
// diff represents how far behind the head is from the tail.
let diff = tail.wrapping_sub(head);
// 1. Check for Lag
if diff > cap {
let missed = diff - cap;
self.head = tail.wrapping_sub(cap);
return Err(TryRecvError::Lagged(missed));
}
// 2. Check if a message is available
if diff > 0 {
let idx = (head as usize) & shared.mask;
let slot = shared.buffer[idx].read();
if slot.version == head {
return if let Some(msg) = &slot.msg {
self.head = head.wrapping_add(1);
Ok(msg.clone())
} else {
Err(TryRecvError::Empty)
};
}
drop(slot);
// If version != head, the slot was overwritten.
// This means we lagged, but the `diff > cap` check missed it (likely due to overflow
// wrapping). We treat this as a lag.
let missed = tail.wrapping_sub(self.head).wrapping_sub(cap);
self.head = tail.wrapping_sub(cap);
return Err(TryRecvError::Lagged(missed));
}
// 3. No message available (diff == 0). Check for Closed.
if shared.senders.load(Ordering::Acquire) == 0 {
return Err(TryRecvError::Disconnected);
}
Err(TryRecvError::Empty)
}
}
impl<T> Receiver<T> {
/// Re-subscribes to the channel, returning a new receiver that starts receiving messages
/// from the *current* tail of the channel.
///
/// This is useful if the receiver has lagged too far behind and wants to jump to the latest
/// message, skipping everything in between.
///
/// # Examples
///
/// ```
/// use mea::broadcast::overflow;
///
/// let (tx, mut rx) = overflow::channel(2);
/// tx.send(1);
/// tx.send(2);
///
/// let mut rx2 = rx.resubscribe();
/// tx.send(3);
///
/// assert_eq!(rx2.try_recv(), Ok(3));
/// ```
pub fn resubscribe(&self) -> Self {
// Resubscribe starts at the current tail.
let head = self.shared.tail_cnt.load(Ordering::SeqCst);
let shared = self.shared.clone();
Self { shared, head }
}
}
struct Recv<'a, T> {
receiver: &'a mut Receiver<T>,
index: Option<usize>,
}
impl<T: Clone> Future for Recv<'_, T> {
type Output = Result<T, RecvError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let Self { receiver, index } = self.get_mut();
loop {
match receiver.try_recv() {
Ok(val) => return Poll::Ready(Ok(val)),
Err(TryRecvError::Lagged(n)) => return Poll::Ready(Err(RecvError::Lagged(n))),
Err(TryRecvError::Disconnected) => {
return Poll::Ready(Err(RecvError::Disconnected));
}
Err(TryRecvError::Empty) => {}
}
let shared = &receiver.shared;
let mut waiters = shared.waiters.lock();
// Double check tail to avoid race conditions.
let tail_now = shared.tail_cnt.load(Ordering::SeqCst);
if tail_now != receiver.head {
// New message arrived while acquiring the lock. Retry.
drop(waiters);
continue;
}
// Check for Closed
// Use Acquire to ensure we see all writes before the sender dropped.
if shared.senders.load(Ordering::Acquire) == 0 {
return Poll::Ready(Err(RecvError::Disconnected));
}
// Register Waker
waiters.register_waker(index, cx);
return Poll::Pending;
}
}
}
+257
View File
@@ -0,0 +1,257 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use super::*;
#[tokio::test]
async fn test_broadcast_basic() {
let (tx, mut rx1) = channel(10);
let mut rx2 = rx1.clone();
tx.send(10);
tx.send(20);
assert_eq!(rx1.recv().await, Ok(10));
assert_eq!(rx1.recv().await, Ok(20));
assert_eq!(rx2.recv().await, Ok(10));
assert_eq!(rx2.recv().await, Ok(20));
}
#[tokio::test]
async fn test_broadcast_lagged() {
let (tx, mut rx) = channel(2);
tx.send(1);
tx.send(2);
tx.send(3);
// Overwrites 1. Rx lagged by 1 (missed msg '1').
// Rx should return Lagged(1) and catch up to 2 (oldest valid).
assert_eq!(rx.recv().await, Err(RecvError::Lagged(1)));
assert_eq!(rx.recv().await, Ok(2));
assert_eq!(rx.recv().await, Ok(3));
}
#[tokio::test]
async fn test_broadcast_lagged_multi() {
let (tx, mut rx) = channel(2);
tx.send(1);
tx.send(2);
tx.send(3);
tx.send(4);
// Overwrites 1 and 2. Missed 2 messages.
assert_eq!(rx.recv().await, Err(RecvError::Lagged(2)));
assert_eq!(rx.recv().await, Ok(3));
assert_eq!(rx.recv().await, Ok(4));
}
#[tokio::test]
async fn test_broadcast_closed() {
let (tx, mut rx) = channel::<()>(10);
drop(tx);
assert_eq!(rx.recv().await, Err(RecvError::Disconnected));
}
#[tokio::test]
async fn test_wait_mechanism() {
let (tx, mut rx) = channel(10);
let handle = tokio::spawn(async move { rx.recv().await });
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
tx.send(42);
assert_eq!(handle.await.unwrap(), Ok(42));
}
#[tokio::test]
async fn test_subscribe() {
let (tx, _rx) = channel(10);
let mut rx = tx.subscribe();
tx.send(100);
assert_eq!(rx.recv().await, Ok(100));
}
#[tokio::test]
async fn test_resubscribe() {
let (tx, mut rx) = channel(2);
tx.send(1);
tx.send(2);
let mut rx2 = rx.resubscribe();
// rx sees 1, 2
// rx2 sees nothing yet (starts at tail=2)
tx.send(3);
assert_eq!(rx.recv().await, Err(RecvError::Lagged(1)));
assert_eq!(rx.recv().await, Ok(2));
assert_eq!(rx2.recv().await, Ok(3));
}
#[tokio::test]
async fn test_overflow() {
let (tx, mut rx) = channel(4);
let mut rx2 = rx.clone();
let boundary = u64::MAX - 2;
tx.shared.tail_cnt.store(boundary, Ordering::SeqCst);
rx.head = boundary;
tx.send(1);
assert_eq!(rx.recv().await, Ok(1));
tx.send(2);
tx.send(3);
tx.send(4);
tx.send(5);
tx.send(6);
tx.send(7);
tx.send(8);
assert_eq!(rx.recv().await, Err(RecvError::Lagged(3)));
assert_eq!(rx.recv().await, Ok(5));
assert_eq!(rx.recv().await, Ok(6));
assert_eq!(rx.recv().await, Ok(7));
assert_eq!(rx.recv().await, Ok(8));
assert_eq!(rx2.recv().await, Err(RecvError::Lagged(1)));
assert_eq!(rx2.recv().await, Ok(5));
assert_eq!(rx2.recv().await, Ok(6));
assert_eq!(rx2.recv().await, Ok(7));
assert_eq!(rx2.recv().await, Ok(8));
}
#[tokio::test]
async fn test_overflow_exactly_overwritten() {
let (tx, mut rx) = channel(4);
let mut rx2 = rx.clone();
let boundary = u64::MAX - 2;
tx.shared.tail_cnt.store(boundary, Ordering::SeqCst);
rx.head = boundary;
tx.send(1);
assert_eq!(rx.recv().await, Ok(1));
tx.send(2);
tx.send(3);
tx.send(4);
tx.send(5);
assert_eq!(rx.recv().await, Ok(2));
// Note: wrapping just hit the head.
// This requires the tail to wrap around the entire u64 space (approx 584 years at 10^9 msg/s),
// which effectively creates an ABA problem where version 0 (wrapped) looks like version 0
// (start). This is a known limitation of the wrapping arithmetic logic, accepted for
// performance reasons as it is practically impossible to trigger without manually setting
// the tail.
assert_eq!(rx2.recv().await, Ok(4));
}
#[tokio::test]
async fn test_capacity_rounding() {
let (tx, _) = channel::<()>(3);
assert_eq!(tx.shared.capacity, 4);
assert_eq!(tx.shared.mask, 3);
let (tx, _) = channel::<()>(4);
assert_eq!(tx.shared.capacity, 4);
assert_eq!(tx.shared.mask, 3);
let (tx, _) = channel::<()>(5);
assert_eq!(tx.shared.capacity, 8);
assert_eq!(tx.shared.mask, 7);
}
#[tokio::test]
async fn test_try_recv() {
let (tx, mut rx) = channel(16);
// Empty
assert_eq!(rx.try_recv(), Err(TryRecvError::Empty));
// Success
tx.send(10);
assert_eq!(rx.try_recv(), Ok(10));
assert_eq!(rx.try_recv(), Err(TryRecvError::Empty));
// Closed
drop(tx);
assert_eq!(rx.try_recv(), Err(TryRecvError::Disconnected));
}
#[tokio::test]
async fn test_try_recv_lagged() {
let (tx, mut rx) = channel(2);
tx.send(1);
tx.send(2);
tx.send(3);
assert_eq!(rx.try_recv(), Err(TryRecvError::Lagged(1)));
assert_eq!(rx.try_recv(), Ok(2));
assert_eq!(rx.try_recv(), Ok(3));
assert_eq!(rx.try_recv(), Err(TryRecvError::Empty));
}
#[tokio::test]
async fn test_try_recv_unwritten_slot_is_empty() {
let (tx, mut rx) = channel::<u64>(2);
drop(tx);
// Simulate tail advanced but slot not written yet
rx.shared.tail_cnt.store(1, Ordering::SeqCst);
assert_eq!(rx.try_recv(), Err(TryRecvError::Empty));
assert_eq!(rx.head, 0);
}
#[tokio::test]
async fn test_multi_senders_concurrent() {
let (tx, mut rx) = channel(100);
let tx1 = tx.clone();
let tx2 = tx.clone();
tokio::spawn(async move {
for i in 0..10 {
tx1.send(i);
}
});
tokio::spawn(async move {
for i in 10..20 {
tx2.send(i);
}
});
// Main tx can also send
for i in 20..30 {
tx.send(i);
}
drop(tx);
let mut received = Vec::new();
while let Ok(n) = rx.recv().await {
received.push(n);
}
received.sort();
let expected = (0..30).collect::<Vec<_>>();
assert_eq!(received, expected);
}
+229
View File
@@ -0,0 +1,229 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//! A condition variable that allows tasks to wait for a notification.
//!
//! # Examples
//!
//! ```
//! # #[tokio::main]
//! # async fn main() {
//! use std::sync::Arc;
//!
//! use mea::condvar::Condvar;
//! use mea::mutex::Mutex;
//!
//! let pair = Arc::new((Mutex::new(false), Condvar::new()));
//! let pair_clone = pair.clone();
//!
//! // Inside our lock, spawn a new thread, and then wait for it to start.
//! tokio::spawn(async move {
//! let (lock, cvar) = &*pair_clone;
//! let mut started = lock.lock().await;
//! *started = true;
//! // We notify the condvar that the value has changed.
//! cvar.notify_one();
//! });
//!
//! // Wait for the thread to start up.
//! let (lock, cvar) = &*pair;
//! let mut started = lock.lock().await;
//! while !*started {
//! started = cvar.wait(started).await;
//! }
//! # }
//! ```
use std::fmt;
use std::task::Waker;
use crate::internal;
use crate::mutex;
use crate::mutex::MutexGuard;
use crate::mutex::OwnedMutexGuard;
#[cfg(test)]
mod tests;
/// A condition variable that allows tasks to wait for a notification.
///
/// See the [module level documentation](self) for more.
pub struct Condvar {
s: internal::Semaphore,
}
impl fmt::Debug for Condvar {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Condvar").finish_non_exhaustive()
}
}
impl Default for Condvar {
fn default() -> Self {
Self::new()
}
}
impl Condvar {
/// Creates a new condition variable
///
/// # Examples
///
/// ```
/// use mea::condvar::Condvar;
///
/// let cvar = Condvar::new();
/// ```
pub const fn new() -> Condvar {
Condvar {
s: internal::Semaphore::new(0),
}
}
/// Wakes up one blocked task on this condvar.
pub fn notify_one(&self) {
self.s.release(1);
}
/// Wakes up all blocked tasks on this condvar.
pub fn notify_all(&self) {
self.s.notify_all();
}
/// Yields the current task until this condition variable receives a notification.
///
/// Unlike the std equivalent, this does not check that a single mutex is used at runtime.
/// However, as a best practice avoid using with multiple mutexes.
pub async fn wait<'a, T>(&self, guard: MutexGuard<'a, T>) -> MutexGuard<'a, T> {
let mutex = mutex::guard_lock(&guard);
// register waiter while holding lock
let mut acquire = self.s.poll_acquire(1);
let _ = acquire.poll_once(Waker::noop());
drop(guard);
// await for notification, and then reacquire the lock
acquire.await;
mutex.lock().await
}
/// Yields the current task until this condition variable receives a notification.
///
/// Unlike the std equivalent, this does not check that a single mutex is used at runtime.
/// However, as a best practice avoid using with multiple mutexes.
pub async fn wait_owned<T>(&self, guard: OwnedMutexGuard<T>) -> OwnedMutexGuard<T> {
let mutex = mutex::owned_guard_lock(&guard);
// register waiter while holding lock
let mut acquire = self.s.poll_acquire(1);
let _ = acquire.poll_once(Waker::noop());
drop(guard);
// await for notification, and then reacquire the lock
acquire.await;
mutex.lock_owned().await
}
/// Yields the current task until this condition variable receives a notification and the
/// provided condition becomes false. Spurious wake-ups are ignored and this function will only
/// return once the condition has been met.
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::condvar::Condvar;
/// use mea::mutex::Mutex;
///
/// let pair = Arc::new((Mutex::new(false), Condvar::new()));
/// let pair_clone = pair.clone();
///
/// tokio::spawn(async move {
/// let (lock, cvar) = &*pair_clone;
/// let mut started = lock.lock().await;
/// *started = true;
/// // We notify the condvar that the value has changed.
/// cvar.notify_one();
/// });
///
/// // Wait for the thread to start up.
/// let (lock, cvar) = &*pair;
/// // As long as the value inside the `Mutex<bool>` is `false`, we wait.
/// let guard = cvar
/// .wait_while(lock.lock().await, |started| !*started)
/// .await;
/// assert!(*guard);
/// # }
/// ```
pub async fn wait_while<'a, T, F>(
&self,
mut guard: MutexGuard<'a, T>,
mut condition: F,
) -> MutexGuard<'a, T>
where
F: FnMut(&mut T) -> bool,
{
while condition(&mut *guard) {
guard = self.wait(guard).await;
}
guard
}
/// Yields the current task until this condition variable receives a notification and the
/// provided condition becomes false. Spurious wake-ups are ignored and this function will only
/// return once the condition has been met.
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::condvar::Condvar;
/// use mea::mutex::Mutex;
///
/// let pair = (Arc::new(Mutex::new(false)), Arc::new(Condvar::new()));
/// let pair_clone = pair.clone();
///
/// tokio::spawn(async move {
/// let (lock, cvar) = pair_clone;
/// let mut started = lock.lock_owned().await;
/// *started = true;
/// // We notify the condvar that the value has changed.
/// cvar.notify_one();
/// });
///
/// // Wait for the thread to start up.
/// let (lock, cvar) = pair;
/// // As long as the value inside the `Mutex<bool>` is `false`, we wait.
/// let guard = cvar
/// .wait_while_owned(lock.lock_owned().await, |started| !*started)
/// .await;
/// assert!(*guard);
/// # }
/// ```
pub async fn wait_while_owned<T, F>(
&self,
mut guard: OwnedMutexGuard<T>,
mut condition: F,
) -> OwnedMutexGuard<T>
where
F: FnMut(&mut T) -> bool,
{
while condition(&mut *guard) {
guard = self.wait_owned(guard).await;
}
guard
}
}
+58
View File
@@ -0,0 +1,58 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::sync::Arc;
use std::time::Duration;
use tokio::task::JoinHandle;
use crate::condvar::Condvar;
use crate::mutex::Mutex;
use crate::test_runtime;
#[test]
fn notify_all() {
test_runtime().block_on(async {
let mut tasks: Vec<JoinHandle<()>> = Vec::new();
let pair = Arc::new((Mutex::new(0u32), Condvar::new()));
for _ in 0..10 {
let pair = pair.clone();
tasks.push(tokio::spawn(async move {
let (m, c) = &*pair;
let mut count = m.lock().await;
while *count == 0 {
count = c.wait(count).await;
}
*count += 1;
}));
}
// Give some time for tasks to start up
tokio::time::sleep(Duration::from_millis(50)).await;
let (m, c) = &*pair;
{
let mut count = m.lock().await;
*count += 1;
c.notify_all();
}
for t in tasks {
t.await.unwrap();
}
let count = m.lock().await;
assert_eq!(11, *count);
});
}
+105
View File
@@ -0,0 +1,105 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::sync::atomic::AtomicU32;
use std::sync::atomic::Ordering;
use std::task::Context;
use crate::internal::Mutex;
use crate::internal::WaitSet;
#[derive(Debug)]
pub(crate) struct CountdownState {
state: AtomicU32,
waiters: Mutex<WaitSet>,
}
impl CountdownState {
pub(crate) const fn new(count: u32) -> Self {
Self {
state: AtomicU32::new(count),
waiters: Mutex::new(WaitSet::new()),
}
}
/// Performs volatile read on `state`.
///
/// All other writes to `state` should be at least [`Ordering::Release`].
pub(crate) fn state(&self) -> u32 {
self.state.load(Ordering::Acquire)
}
/// Performs volatile CAS on `state`.
///
/// If the comparison succeeds, performs read-modify-write operation with [`Ordering::Relaxed`]
/// for read, and [`Ordering::Release`] for write; if the comparison fails, performs load
/// operation with [`Ordering::Relaxed`].
///
/// @see https://doc.rust-lang.org/std/sync/atomic/struct.AtomicU32.html#method.compare_exchange_weak
/// @see https://en.cppreference.com/w/cpp/atomic/atomic_compare_exchange
pub(crate) fn cas_state(&self, current: u32, new: u32) -> Result<(), u32> {
self.state
.compare_exchange_weak(current, new, Ordering::Release, Ordering::Relaxed)
.map(|_| ())
}
/// Drain and wake up all waiters.
pub(crate) fn wake_all(&self) {
let mut waiters = self.waiters.lock();
waiters.wake_all();
}
/// Registers a waker to be woken up when the countdown reaches zero.
///
/// `idx` must be `None` when the waker is not registered, or `Some(key)` where `key` is
/// a value previously returned by this method.
pub(crate) fn register_waker(&self, idx: &mut Option<usize>, cx: &mut Context<'_>) {
let mut waiters = self.waiters.lock();
waiters.register_waker(idx, cx);
}
/// Returns `Ok(())` if the counter is zero, otherwise returns `Err(s)` where `s` is the current
/// counter value.
pub(crate) fn spin_wait(&self, n: usize) -> Result<(), u32> {
for _ in 0..n {
if self.state() == 0 {
return Ok(());
}
std::hint::spin_loop();
}
match self.state() {
0 => Ok(()),
s => Err(s),
}
}
/// Decrements the counter, and returns whether the caller should wake up all waiters.
pub(crate) fn decrement(&self, n: u32) -> bool {
let mut cnt = self.state();
loop {
if cnt == 0 {
// the one who decrements the counter to zero should wake up all waiters, not this
// one
return false;
}
let new_cnt = cnt.saturating_sub(n);
match self.cas_state(cnt, new_cnt) {
Ok(_) => return new_cnt == 0,
Err(x) => cnt = x,
}
}
}
}
+31
View File
@@ -0,0 +1,31 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
mod countdown;
pub(crate) use countdown::*;
mod mutex;
pub(crate) use mutex::*;
mod rwlock;
pub(crate) use rwlock::*;
mod semaphore;
pub(crate) use semaphore::*;
mod waitlist;
pub(crate) use waitlist::*;
mod waitset;
pub(crate) use waitset::*;
+54
View File
@@ -0,0 +1,54 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::fmt;
use std::sync::PoisonError;
pub(crate) struct Mutex<T: ?Sized>(std::sync::Mutex<T>);
impl<T: ?Sized + fmt::Debug> fmt::Debug for Mutex<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
self.0.fmt(f)
}
}
impl<T> Mutex<T> {
pub(crate) const fn new(t: T) -> Self {
Self(std::sync::Mutex::new(t))
}
pub(crate) fn lock(&self) -> std::sync::MutexGuard<'_, T> {
self.0.lock().unwrap_or_else(PoisonError::into_inner)
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use crate::internal::Mutex;
#[test]
fn test_poison_mutex() {
let mutex = Arc::new(Mutex::new(42));
let m = mutex.clone();
let handle = std::thread::spawn(move || {
let _guard = m.lock();
panic!("poison");
});
let _ = handle.join();
let guard = mutex.lock();
assert_eq!(*guard, 42);
}
}
+53
View File
@@ -0,0 +1,53 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::sync::PoisonError;
pub(crate) struct RwLock<T: ?Sized>(std::sync::RwLock<T>);
impl<T> RwLock<T> {
pub(crate) const fn new(t: T) -> Self {
Self(std::sync::RwLock::new(t))
}
}
impl<T: ?Sized> RwLock<T> {
pub(crate) fn read(&self) -> std::sync::RwLockReadGuard<'_, T> {
self.0.read().unwrap_or_else(PoisonError::into_inner)
}
pub(crate) fn write(&self) -> std::sync::RwLockWriteGuard<'_, T> {
self.0.write().unwrap_or_else(PoisonError::into_inner)
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use crate::internal::RwLock;
#[test]
fn test_poison_rwlock() {
let rwlock = Arc::new(RwLock::new(42));
let r = rwlock.clone();
let handle = std::thread::spawn(move || {
let _guard = r.write();
panic!("poison");
});
let _ = handle.join();
assert_eq!(*rwlock.read(), 42);
assert_eq!(*rwlock.write(), 42);
}
}
+369
View File
@@ -0,0 +1,369 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::future::Future;
use std::pin::Pin;
use std::sync::MutexGuard;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::task::Context;
use std::task::Poll;
use std::task::Waker;
use crate::internal::Mutex;
use crate::internal::WaitList;
/// The internal semaphore that provides low-level async primitives.
#[derive(Debug)]
pub(crate) struct Semaphore {
/// The current number of available permits in the semaphore.
permits: AtomicUsize,
waiters: Mutex<WaitList<WaitNode>>,
}
#[derive(Debug)]
struct WaitNode {
permits: usize,
waker: Option<Waker>,
}
impl Semaphore {
pub(crate) const fn new(permits: usize) -> Self {
Self {
permits: AtomicUsize::new(permits),
waiters: Mutex::new(WaitList::new()),
}
}
/// Returns the current number of available permits.
pub(crate) fn available_permits(&self) -> usize {
self.permits.load(Ordering::Acquire)
}
/// Tries to acquire `n` permits from the semaphore.
///
/// Returns `true` if the permits were acquired, `false` otherwise.
pub(crate) fn try_acquire(&self, n: usize) -> bool {
let mut current = self.permits.load(Ordering::Acquire);
loop {
if current < n {
return false;
}
let next = current - n;
match self
.permits
.compare_exchange(current, next, Ordering::AcqRel, Ordering::Acquire)
{
Ok(_) => return true,
Err(actual) => current = actual,
}
}
}
/// Decrease the semaphore's permits by a maximum of `n`.
///
/// Return the number of permits that were actually reduced.
pub(crate) fn forget(&self, n: usize) -> usize {
if n == 0 {
return 0;
}
let mut current = self.permits.load(Ordering::Acquire);
loop {
let new = current.saturating_sub(n);
match self.permits.compare_exchange_weak(
current,
new,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => return n.min(current),
Err(actual) => current = actual,
}
}
}
/// Decrease the semaphore's permits by `n`.
///
/// If the semaphore has not enough permits, enqueue front an empty waiter to consume the
/// permits.
pub(crate) fn forget_exact(&self, n: usize) {
acquired_or_enqueue(self, n, &mut None, None, false);
}
/// Acquires `n` permits from the semaphore.
pub(crate) async fn acquire(&self, n: usize) {
let fut = Acquire {
permits: n,
index: None,
semaphore: self,
done: false,
};
fut.await
}
/// Returns a future that is resolved when acquired `n` permits from the semaphore.
pub(crate) fn poll_acquire(&self, n: usize) -> Acquire<'_> {
Acquire {
permits: n,
index: None,
semaphore: self,
done: false,
}
}
/// Adds `n` permits to the semaphore.
pub(crate) fn release(&self, n: usize) {
if n != 0 {
self.insert_permits_with_lock(n, self.waiters.lock());
}
}
/// Adds `n` permits to the semaphore if there is any waiter.
pub(crate) fn release_if_nonempty(&self, n: usize) {
let waiters = self.waiters.lock();
if !waiters.is_empty() {
self.insert_permits_with_lock(n, waiters);
}
}
/// Adds as many permits until there is no waiter.
pub(crate) fn notify_all(&self) {
let mut waiters = self.waiters.lock();
let mut wakers = Vec::new();
loop {
match waiters.remove_first_waiter(|node| {
node.permits = 0;
true
}) {
None => break,
Some(waiter) => {
if let Some(waker) = waiter.waker.take() {
wakers.push(waker);
}
}
}
}
drop(waiters);
for w in wakers.drain(..) {
w.wake();
}
}
fn insert_permits_with_lock(
&self,
mut rem: usize,
waiters: MutexGuard<'_, WaitList<WaitNode>>,
) {
const NUM_WAKER: usize = 32;
let mut wakers = Vec::with_capacity(NUM_WAKER);
let mut lock = Some(waiters);
while rem > 0 {
let mut waiters = lock.take().unwrap_or_else(|| self.waiters.lock());
while wakers.len() < NUM_WAKER {
match waiters.remove_first_waiter(|node| {
if node.permits <= rem {
rem -= node.permits;
node.permits = 0;
true
} else {
node.permits -= rem;
rem = 0;
false
}
}) {
None => break,
Some(waiter) => {
if let Some(waker) = waiter.waker.take() {
wakers.push(waker);
}
}
}
}
if rem > 0 && waiters.is_empty() {
let permits = rem;
let prev = self.permits.fetch_add(permits, Ordering::Release);
assert!(
prev.checked_add(permits).is_some(),
"number of added permits ({permits}) would overflow usize::MAX (prev: {prev})"
);
rem = 0;
}
drop(waiters);
for w in wakers.drain(..) {
w.wake();
}
}
}
}
#[derive(Debug)]
pub(crate) struct Acquire<'a> {
permits: usize,
index: Option<usize>,
semaphore: &'a Semaphore,
done: bool,
}
impl Drop for Acquire<'_> {
fn drop(&mut self) {
if let Some(index) = self.index {
let mut waiters = self.semaphore.waiters.lock();
let mut acquired = 0;
waiters.remove_waiter(index, |node| {
acquired = self.permits - node.permits;
node.permits = 0;
true
});
waiters.with_mut(index, |_| true); // drop
if acquired > 0 {
self.semaphore.insert_permits_with_lock(acquired, waiters);
}
}
}
}
impl Acquire<'_> {
pub(crate) fn poll_once(&mut self, waker: &Waker) -> Poll<()> {
let Self {
permits,
index,
semaphore,
done,
} = self;
if *done {
return Poll::Ready(());
}
match index {
Some(idx) => {
let mut waiters = semaphore.waiters.lock();
let mut ready = false;
waiters.with_mut(*idx, |node| {
if node.permits > 0 {
let update_waker = node.waker.as_ref().is_none_or(|w| !w.will_wake(waker));
if update_waker {
node.waker = Some(waker.clone());
}
false
} else {
ready = true;
true
}
});
if ready {
*index = None;
*done = true;
return Poll::Ready(());
}
}
None => {
// not yet enqueued
let needed = *permits;
if acquired_or_enqueue(semaphore, needed, index, Some(waker), true) {
*done = true;
return Poll::Ready(());
}
}
};
Poll::Pending
}
}
impl Future for Acquire<'_> {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.get_mut();
this.poll_once(cx.waker())
}
}
/// Returns `true` if successfully acquired the semaphore; `false` otherwise.
fn acquired_or_enqueue(
sem: &Semaphore,
needed: usize,
idx: &mut Option<usize>,
waker: Option<&Waker>,
enqueue_last: bool,
) -> bool {
let mut current = sem.permits.load(Ordering::Acquire);
let mut lock = None;
loop {
let (remaining, next) = if current >= needed {
(0, current - needed)
} else {
(needed - current, 0)
};
if remaining > 0 && lock.is_none() {
// No permits were immediately available, so this permit will
// (probably) need to wait. We'll need to acquire a lock on the
// wait queue before continuing. We need to do this _before_ the
// CAS that sets the new value of the semaphore's `permits`
// counter. Otherwise, if we subtract the permits and then
// acquire the lock, we might miss additional permits being
// added while waiting for the lock.
lock = Some(sem.waiters.lock());
}
if let Err(actual) =
sem.permits
.compare_exchange(current, next, Ordering::AcqRel, Ordering::Acquire)
{
// other thread changed the permits; retry
current = actual;
continue;
}
// all needed permits were acquired
if remaining == 0 {
return true;
}
// all available permits were acquired, but more are needed;
// enqueue a waiter with the remaining needed permits
let mut waiters = lock.take().unwrap_or_else(|| {
unreachable!("lock must be acquired when remaining {remaining} > 0");
});
if enqueue_last {
waiters.register_waiter_to_tail(idx, || {
Some(WaitNode {
permits: remaining,
waker: waker.cloned(),
})
});
} else {
waiters.register_waiter_to_head(idx, || {
Some(WaitNode {
permits: remaining,
waker: waker.cloned(),
})
});
}
return false;
}
}
+165
View File
@@ -0,0 +1,165 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use slab::Slab;
/// A guarded linked list.
///
/// * `guard`'s `next` points to the first node (regular head).
/// * `guard`'s `prev` points to the last node (regular tail).
#[derive(Debug)]
pub(crate) struct WaitList<T> {
// if None, the list is uninitialized and empty
guard: Option<usize>,
nodes: Slab<Node<T>>,
}
#[derive(Debug)]
struct Node<T> {
prev: usize,
next: usize,
stat: Option<T>,
}
impl<T> WaitList<T> {
/// Ensures the wait list is initialized, returning the guard index.
fn ensure_init(&mut self) -> usize {
if let Some(guard) = self.guard {
return guard;
}
let first = self.nodes.vacant_entry();
let guard = first.key();
first.insert(Node {
prev: guard,
next: guard,
stat: None,
});
self.guard = Some(guard);
guard
}
pub(crate) const fn new() -> Self {
Self {
guard: None,
nodes: Slab::new(),
}
}
/// Registers a waiter to the head of the wait list.
///
/// # Panic
///
/// Panics if `idx` is `Some`.
pub(crate) fn register_waiter_to_head(
&mut self,
idx: &mut Option<usize>,
f: impl FnOnce() -> Option<T>,
) {
assert!(idx.is_none());
let guard = self.ensure_init();
let stat = f();
let prev_head = self.nodes[guard].next;
let new_node = Node {
prev: guard,
next: prev_head,
stat,
};
let new_key = self.nodes.insert(new_node);
self.nodes[guard].next = new_key;
self.nodes[prev_head].prev = new_key;
*idx = Some(new_key);
}
/// Registers a waiter to the tail of the wait list.
///
/// # Panic
///
/// Panics if `idx` is `Some`.
pub(crate) fn register_waiter_to_tail(
&mut self,
idx: &mut Option<usize>,
f: impl FnOnce() -> Option<T>,
) {
assert!(idx.is_none());
let guard = self.ensure_init();
let stat = f();
let prev_tail = self.nodes[guard].prev;
let new_node = Node {
prev: prev_tail,
next: guard,
stat,
};
let new_key = self.nodes.insert(new_node);
self.nodes[guard].prev = new_key;
self.nodes[prev_tail].next = new_key;
*idx = Some(new_key);
}
/// Removes a previously registered waker from the wait list, if the predicate `f` returns
/// `true`.
pub(crate) fn remove_waiter(
&mut self,
idx: usize,
f: impl FnOnce(&mut T) -> bool,
) -> Option<&mut T> {
// SAFETY: the wait list must be initialized before any waiter can be registered
let guard = self.guard.expect("wait list must be uninitialized");
assert_ne!(idx, guard);
fn retrieve_stat<T>(node: &mut Node<T>) -> &mut T {
// SAFETY: `idx` is a valid key + non-guard node always has `Some(stat)`
node.stat.as_mut().unwrap()
}
if f(retrieve_stat(&mut self.nodes[idx])) {
let prev = self.nodes[idx].prev;
let next = self.nodes[idx].next;
self.nodes[prev].next = next;
self.nodes[next].prev = prev;
self.nodes[idx].prev = idx;
self.nodes[idx].next = idx;
Some(retrieve_stat(&mut self.nodes[idx]))
} else {
None
}
}
/// Removes the first waiter from the wait list, if the predicate `f` returns `true`.
pub(crate) fn remove_first_waiter(&mut self, f: impl FnOnce(&mut T) -> bool) -> Option<&mut T> {
let guard = self.guard?;
let first = self.nodes[guard].next;
if first != guard {
self.remove_waiter(first, f)
} else {
None
}
}
/// Returns `true` if the wait list is empty.
pub(crate) fn is_empty(&self) -> bool {
self.guard
.is_none_or(|guard| self.nodes[guard].next == guard)
}
pub(crate) fn with_mut(&mut self, idx: usize, drop: impl FnOnce(&mut T) -> bool) {
let node = &mut self.nodes[idx];
if drop(node.stat.as_mut().unwrap()) {
self.nodes.remove(idx);
}
}
}
+80
View File
@@ -0,0 +1,80 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::task::Context;
use std::task::Waker;
use slab::Slab;
#[derive(Debug)]
pub(crate) struct WaitSet {
waiters: Slab<Waker>,
}
impl WaitSet {
/// Construct a new, empty wait set.
pub const fn new() -> Self {
Self {
waiters: Slab::new(),
}
}
/// Construct a new, empty wait set with the specified capacity.
pub fn with_capacity(capacity: usize) -> Self {
Self {
waiters: Slab::with_capacity(capacity),
}
}
/// Drain and wake up all waiters.
pub(crate) fn wake_all(&mut self) {
for w in self.waiters.drain() {
w.wake();
}
}
/// Registers a waker to the wait set.
///
/// `idx` must be `None` when the waker is not registered, or `Some(key)` where `key` is
/// a value previously returned by this method.
pub(crate) fn register_waker(&mut self, idx: &mut Option<usize>, cx: &mut Context<'_>) {
match *idx {
None => {
let key = self.waiters.insert(cx.waker().clone());
*idx = Some(key);
}
Some(key) => {
if self.waiters.contains(key) {
if !self.waiters[key].will_wake(cx.waker()) {
self.waiters[key] = cx.waker().clone();
}
} else {
// DEFENSIVE NOTE:
//
// This is possible if latch/waitgroup is fired between the first and second
// state check.
//
// In this case, it does not harm to re-register the waker. Because
// the second state check will finish the future and the WaitSet gets
// dropped.
//
// Barrier holds the lock during check and register, so the race condition
// above won't happen.
let key = self.waiters.insert(cx.waker().clone());
*idx = Some(key);
}
}
}
}
}
+315
View File
@@ -0,0 +1,315 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//! A countdown latch that allows one or more tasks to wait until a set of operations completes.
//!
//! Unlike a barrier, a latch's count can only decrease and cannot be reused once it reaches zero.
//! This makes it ideal for scenarios where you need to wait for a specific number of events or
//! operations to complete.
//!
//! A latch starts with an initial count and tasks can wait for this count to reach zero.
//! The count can be decremented by calling [`count_down()`] or [`arrive()`]. Once the count
//! reaches zero, all waiting tasks are unblocked.
//!
//! # Examples
//!
//! ```
//! # #[tokio::main]
//! # async fn main() {
//! use std::sync::Arc;
//!
//! use mea::latch::Latch;
//!
//! let latch = Arc::new(Latch::new(3));
//! let mut handles = Vec::new();
//!
//! for i in 0..3 {
//! let latch = latch.clone();
//! handles.push(tokio::spawn(async move {
//! println!("Task {} starting", i);
//! // Simulate some work
//! latch.count_down(); // Signal completion
//! }));
//! }
//!
//! // Wait for all tasks to complete
//! latch.wait().await;
//! println!("All tasks completed");
//! # }
//! ```
//!
//! [`count_down()`]: Latch::count_down
//! [`arrive()`]: Latch::arrive
use std::fmt;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::task::Context;
use std::task::Poll;
use crate::internal::CountdownState;
#[cfg(test)]
mod tests;
/// A synchronization primitive that can be used to coordinate multiple tasks.
///
/// See the [module level documentation](self) for more.
#[derive(Debug)]
pub struct Latch {
state: CountdownState,
}
impl Latch {
/// Creates a new latch initialized with the given count.
///
/// # Arguments
///
/// * `count` - The initial count value. Tasks will wait until this count reaches zero.
///
/// # Examples
///
/// ```
/// use mea::latch::Latch;
///
/// let latch = Latch::new(3); // Creates a latch with count of 3
/// ```
pub fn new(count: u32) -> Self {
Self {
state: CountdownState::new(count),
}
}
/// Returns the current count.
///
/// This method is typically used for debugging and testing purposes.
///
/// # Examples
///
/// ```
/// use mea::latch::Latch;
///
/// let latch = Latch::new(5);
/// assert_eq!(latch.count(), 5);
/// ```
pub fn count(&self) -> u32 {
self.state.state()
}
/// Decrements the latch count by one, waking up all pending tasks if the counter reaches zero.
///
/// If the current count is zero, this method has no effect.
///
/// # Examples
///
/// ```
/// use mea::latch::Latch;
///
/// let latch = Latch::new(2);
/// latch.count_down(); // Count is now 1
/// latch.count_down(); // Count is now 0, all waiting tasks are woken
/// ```
pub fn count_down(&self) {
if self.state.decrement(1) {
self.state.wake_all();
}
}
/// Decrements the latch count by `n`, waking up all waiting tasks if the counter reaches zero.
///
/// This method provides a way to decrement the counter by more than one at a time.
/// It will not cause an overflow when decrementing the counter.
///
/// # Arguments
///
/// * `n` - The amount to decrement the counter by
///
/// # Behavior
///
/// * If `n` is zero or the counter has already reached zero, nothing happens
/// * If the current count is greater than `n`, it is decremented by `n`
/// * If the current count is greater than 0 but less than or equal to `n`, the count becomes
/// zero and all waiting tasks are woken
///
/// # Examples
///
/// ```
/// use mea::latch::Latch;
///
/// let latch = Latch::new(5);
/// latch.arrive(3); // Count is now 2
/// latch.arrive(2); // Count is now 0, all waiting tasks are woken
/// ```
pub fn arrive(&self, n: u32) {
if n != 0 && self.state.decrement(n) {
self.state.wake_all();
}
}
/// Attempts to wait for the latch count to reach zero without blocking.
///
/// # Returns
///
/// * `Ok(())` if the count is zero
/// * `Err(count)` if the count is not zero, where `count` is the current count
///
/// # Examples
///
/// ```
/// use mea::latch::Latch;
///
/// let latch = Latch::new(2);
/// assert_eq!(latch.try_wait(), Err(2));
/// latch.count_down();
/// assert_eq!(latch.try_wait(), Err(1));
/// latch.count_down();
/// assert_eq!(latch.try_wait(), Ok(()));
/// ```
pub fn try_wait(&self) -> Result<(), u32> {
self.state.spin_wait(0)
}
/// Returns a future that will complete when the latch count reaches zero.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::latch::Latch;
///
/// let latch = Arc::new(Latch::new(1));
/// let latch2 = latch.clone();
///
/// // Spawn a task that will wait for the latch
/// let handle = tokio::spawn(async move {
/// latch2.wait().await;
/// println!("Latch reached zero!");
/// });
///
/// // Count down the latch
/// latch.count_down();
/// handle.await.unwrap();
/// # }
/// ```
pub async fn wait(&self) {
let fut = LatchWait {
idx: None,
latch: self,
};
fut.await
}
/// Returns a future that will complete when the latch count reaches zero.
///
/// The latch must be wrapped in an [`Arc`] to call this method. Thus, the returned future has
/// no lifetime constraints.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::latch::Latch;
///
/// let latch = Arc::new(Latch::new(1));
/// let latch2 = latch.clone();
///
/// // Spawn a task that will wait for the latch
/// let handle = tokio::spawn(async move {
/// latch2.wait_owned().await;
/// println!("Latch reached zero!");
/// });
///
/// // Count down the latch
/// latch.count_down();
/// handle.await.unwrap();
/// # }
/// ```
pub async fn wait_owned(self: Arc<Self>) {
let fut = OwnedLatchWait {
idx: None,
latch: self,
};
fut.await
}
}
impl Latch {
fn intern_poll(&self, idx: &mut Option<usize>, cx: &mut Context<'_>) -> Poll<()> {
// register waker if the counter is not zero
if self.state.spin_wait(16).is_err() {
self.state.register_waker(idx, cx);
// double check after register waker, to catch the update between two steps
if self.state.spin_wait(0).is_err() {
return Poll::Pending;
}
}
Poll::Ready(())
}
}
/// A wait future returned by [`Latch::wait()`].
///
/// This future will complete when the latch count reaches zero.
#[must_use = "futures do nothing unless you `.await` or poll them"]
pub struct LatchWait<'a> {
idx: Option<usize>,
latch: &'a Latch,
}
impl fmt::Debug for LatchWait<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("LatchWait").finish_non_exhaustive()
}
}
impl Future for LatchWait<'_> {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let Self { idx, latch } = self.get_mut();
latch.intern_poll(idx, cx)
}
}
/// An owned wait future returned by [`Latch::wait()`].
///
/// This future will complete when the latch count reaches zero.
#[must_use = "futures do nothing unless you `.await` or poll them"]
pub struct OwnedLatchWait {
idx: Option<usize>,
latch: Arc<Latch>,
}
impl fmt::Debug for OwnedLatchWait {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("OwnedLatchWait").finish_non_exhaustive()
}
}
impl Future for OwnedLatchWait {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let Self { idx, latch } = self.get_mut();
latch.intern_poll(idx, cx)
}
}
+194
View File
@@ -0,0 +1,194 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::prelude::rust_2015::Vec;
use std::sync::Arc;
use std::sync::atomic::AtomicU32;
use std::sync::atomic::Ordering;
use std::time::Duration;
use std::time::Instant;
use super::*;
macro_rules! assert_time {
($time:expr, $mills:literal $(,)?) => {
assert!(
(Duration::from_millis($mills - 1)
..Duration::from_millis($mills + std::cmp::max($mills >> 1, 50)))
.contains(&$time)
)
};
}
#[test]
fn test_count_down() {
let latch = Latch::new(3);
latch.count_down();
latch.count_down();
latch.count_down();
assert_eq!(latch.count(), 0);
}
#[test]
fn test_try_wait() {
let latch = Latch::new(0);
assert_eq!(latch.try_wait(), Ok(()));
}
#[test]
fn test_try_wait_err() {
let latch = Latch::new(3);
assert_eq!(latch.try_wait(), Err(3));
}
#[test]
fn test_arrive_zero() {
let latch = Latch::new(2);
latch.arrive(0);
assert_eq!(latch.count(), 2);
}
#[test]
fn test_more_arrive() {
let latch = Latch::new(10);
for _ in 0..4 {
latch.arrive(3);
}
assert_eq!(latch.count(), 0);
}
#[tokio::test]
async fn test_arrive() {
let latch = Latch::new(3);
latch.arrive(3);
latch.wait().await;
assert_eq!(latch.count(), 0);
}
#[tokio::test]
async fn test_last_one_signal() {
let latch = Arc::new(Latch::new(3));
let l1 = latch.clone();
latch.count_down();
latch.count_down();
let start = Instant::now();
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(32)).await;
l1.count_down()
});
latch.wait().await;
assert_time!(start.elapsed(), 32);
assert_eq!(latch.count(), 0);
}
#[tokio::test]
async fn test_gate_wait() {
let latch = Arc::new(Latch::new(1));
let tasks: Vec<_> = (0..4)
.map(|_| {
let latch = latch.clone();
let start = Instant::now();
tokio::spawn(async move {
latch.wait().await;
start.elapsed()
})
})
.collect();
tokio::time::sleep(Duration::from_millis(20)).await;
latch.count_down();
for t in tasks {
assert_time!(t.await.unwrap(), 20);
}
}
#[tokio::test]
async fn test_multi_tasks() {
const SIZE: u32 = 16;
let latch = Arc::new(Latch::new(SIZE));
let counter = Arc::new(AtomicU32::new(0));
for _ in 0..SIZE {
let latch = latch.clone();
let counter = counter.clone();
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(20)).await;
counter.fetch_add(1, Ordering::Relaxed);
latch.count_down()
});
}
let start = Instant::now();
latch.wait().await;
assert_time!(start.elapsed(), 20);
assert_eq!(counter.load(Ordering::Relaxed), SIZE);
assert_eq!(latch.count(), 0);
}
#[tokio::test]
async fn test_more_count_down() {
const SIZE: u32 = 16;
let latch = Arc::new(Latch::new(SIZE));
let counter = Arc::new(AtomicU32::new(0));
for _ in 0..(SIZE + (SIZE >> 1)) {
let latch = latch.clone();
let counter = counter.clone();
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(10)).await;
counter.fetch_add(1, Ordering::Relaxed);
latch.count_down()
});
}
latch.wait().await;
assert!(counter.load(Ordering::Relaxed) >= SIZE);
assert_eq!(latch.count(), 0);
latch.count_down();
assert_eq!(latch.count(), 0);
}
#[tokio::test]
async fn test_select_two_wait() {
let latch1 = Arc::new(Latch::new(1));
let latch2 = Arc::new(Latch::new(1));
let l1 = latch1.clone();
let l2 = latch2.clone();
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(50)).await;
l1.count_down();
});
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(10)).await;
l2.count_down();
});
assert!(tokio::select! {
_ = latch1.wait() => false,
_ = latch2.wait() => true,
});
latch1.wait().await;
}
+210
View File
@@ -0,0 +1,210 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
#![cfg_attr(docsrs, feature(doc_cfg))]
#![deny(missing_docs)]
//! `mea` is a runtime-agnostic library providing essential synchronization primitives for
//! asynchronous Rust programming. The library offers a collection of well-tested, efficient
//! synchronization tools that work with any async runtime.
//!
//! # Features
//!
//! * [`admission::FairShare`]: A work-conserving admission policy that fairly shares bounded
//! concurrency across keys
//! * [`Barrier`]: A synchronization point where multiple tasks can wait until all participants
//! arrive
//! * [`Condvar`]: A condition variable that allows tasks to wait for a notification
//! * [`Latch`]: A single-use barrier that allows one or more tasks to wait until a signal is given
//! * [`Mutex`]: A mutual exclusion primitive for protecting shared data
//! * [`Once`]: A primitive that ensures a one-time asynchronous operation runs at most once, even
//! when called concurrently
//! * [`OnceCell`]: A cell that can be written to at most once and provides safe concurrent access
//! * [`OnceMap`]: A hash map that runs computation only once for each key and stores the result.
//! * [`RwLock`]: A reader-writer lock that allows multiple readers or a single writer at a time
//! * [`Semaphore`]: A synchronization primitive that controls access to a shared resource
//! * [`ShutdownSend`], [`ShutdownRecv`] & [`ShutdownWatch`]: A composite synchronization primitive
//! for managing shutdown signals
//! * [`WaitGroup`]: A synchronization primitive that allows waiting for multiple tasks to complete
//! * [`atomicbox`]: A safe, owning version of `AtomicPtr` for heap-allocated data.
//! * [`broadcast`]: A multi-producer, multi-consumer broadcast channel.
//! * [`mpsc::bounded`]: A multi-producer, single-consumer bounded queue for sending values between
//! asynchronous tasks.
//! * [`mpsc::unbounded`]: A multi-producer, single-consumer unbounded queue for sending values
//! between asynchronous tasks.
//! * [`oneshot::channel`]: A one-shot channel for sending a single value between tasks.
//! * [`singleflight::Group`]: A duplicate function call suppression mechanism.
//!
//! # Runtime Agnostic
//!
//! All synchronization primitives in this library are runtime-agnostic, meaning they can be used
//! with any async runtime like tokio, async-std, or others. This makes the library highly versatile
//! and portable.
//!
//! # Thread Safety
//!
//! All primitive types in this library implement `Send` and `Sync`, making them safe to share
//! across thread boundaries. This is essential for concurrent programming where data needs to be
//! accessed from multiple threads.
//!
//! [`Barrier`]: barrier::Barrier
//! [`Condvar`]: condvar::Condvar
//! [`Latch`]: latch::Latch
//! [`Mutex`]: mutex::Mutex
//! [`Once`]: once::Once
//! [`OnceCell`]: once::OnceCell
//! [`OnceMap`]: once::OnceMap
//! [`RwLock`]: rwlock::RwLock
//! [`Semaphore`]: semaphore::Semaphore
//! [`ShutdownSend`]: shutdown::ShutdownSend
//! [`ShutdownRecv`]: shutdown::ShutdownRecv
//! [`ShutdownWatch`]: shutdown::ShutdownWatch
//! [`WaitGroup`]: waitgroup::WaitGroup
mod internal;
pub mod admission;
pub mod atomicbox;
pub mod barrier;
pub mod broadcast;
pub mod condvar;
pub mod latch;
pub mod mpsc;
pub mod mutex;
pub mod once;
pub mod oneshot;
pub mod rwlock;
pub mod semaphore;
pub mod shutdown;
pub mod singleflight;
pub mod waitgroup;
#[cfg(test)]
fn test_runtime() -> &'static tokio::runtime::Runtime {
use std::sync::OnceLock;
use tokio::runtime::Runtime;
static RT: OnceLock<Runtime> = OnceLock::new();
RT.get_or_init(|| Runtime::new().unwrap())
}
#[cfg(test)]
mod tests {
use crate::admission::FairShare;
use crate::admission::FairSharePermit;
use crate::admission::OwnedFairSharePermit;
use crate::barrier::Barrier;
use crate::broadcast;
use crate::condvar::Condvar;
use crate::latch::Latch;
use crate::mpsc;
use crate::mutex::Mutex;
use crate::mutex::MutexGuard;
use crate::once::Once;
use crate::once::OnceCell;
use crate::once::OnceMap;
use crate::oneshot;
use crate::rwlock::RwLock;
use crate::rwlock::RwLockReadGuard;
use crate::rwlock::RwLockWriteGuard;
use crate::semaphore::Semaphore;
use crate::shutdown::ShutdownRecv;
use crate::shutdown::ShutdownSend;
use crate::shutdown::ShutdownWatch;
use crate::singleflight;
use crate::waitgroup::Wait;
use crate::waitgroup::WaitGroup;
#[test]
fn assert_send_and_sync() {
fn do_assert_send_and_sync<T: Send + Sync>() {}
do_assert_send_and_sync::<FairShare<String>>();
do_assert_send_and_sync::<FairSharePermit<'_, String>>();
do_assert_send_and_sync::<OwnedFairSharePermit<String>>();
do_assert_send_and_sync::<Barrier>();
do_assert_send_and_sync::<Condvar>();
do_assert_send_and_sync::<Once>();
do_assert_send_and_sync::<OnceCell<u32>>();
do_assert_send_and_sync::<OnceMap<String, u32>>();
do_assert_send_and_sync::<singleflight::Group<String, u32>>();
do_assert_send_and_sync::<Latch>();
do_assert_send_and_sync::<Semaphore>();
do_assert_send_and_sync::<ShutdownSend>();
do_assert_send_and_sync::<ShutdownRecv>();
do_assert_send_and_sync::<ShutdownWatch>();
do_assert_send_and_sync::<WaitGroup>();
do_assert_send_and_sync::<Mutex<i64>>();
do_assert_send_and_sync::<MutexGuard<'_, i64>>();
do_assert_send_and_sync::<RwLock<i64>>();
do_assert_send_and_sync::<RwLockReadGuard<'_, i64>>();
do_assert_send_and_sync::<RwLockWriteGuard<'_, i64>>();
do_assert_send_and_sync::<broadcast::overflow::Sender<i64>>();
do_assert_send_and_sync::<broadcast::overflow::Receiver<i64>>();
do_assert_send_and_sync::<broadcast::overflow::RecvError>();
do_assert_send_and_sync::<broadcast::overflow::TryRecvError>();
do_assert_send_and_sync::<oneshot::SendError<i64>>();
do_assert_send_and_sync::<oneshot::Sender<i64>>();
do_assert_send_and_sync::<mpsc::SendError<i64>>();
do_assert_send_and_sync::<mpsc::UnboundedSender<i64>>();
do_assert_send_and_sync::<mpsc::UnboundedReceiver<i64>>();
do_assert_send_and_sync::<mpsc::BoundedSender<i64>>();
do_assert_send_and_sync::<mpsc::BoundedReceiver<i64>>();
}
#[test]
fn assert_send() {
fn do_assert_send<T: Send>() {}
do_assert_send::<oneshot::Receiver<i64>>();
do_assert_send::<oneshot::Recv<i64>>();
}
#[test]
fn assert_unpin() {
fn do_assert_unpin<T: Unpin>() {}
do_assert_unpin::<FairShare<String>>();
do_assert_unpin::<FairSharePermit<'_, String>>();
do_assert_unpin::<OwnedFairSharePermit<String>>();
do_assert_unpin::<Barrier>();
do_assert_unpin::<Condvar>();
do_assert_unpin::<Latch>();
do_assert_unpin::<Once>();
do_assert_unpin::<OnceCell<u32>>();
do_assert_unpin::<OnceMap<String, u32>>();
do_assert_unpin::<singleflight::Group<String, u32>>();
do_assert_unpin::<Semaphore>();
do_assert_unpin::<ShutdownSend>();
do_assert_unpin::<ShutdownRecv>();
do_assert_unpin::<ShutdownWatch>();
do_assert_unpin::<WaitGroup>();
do_assert_unpin::<Wait>();
do_assert_unpin::<Mutex<i64>>();
do_assert_unpin::<MutexGuard<'_, i64>>();
do_assert_unpin::<RwLock<i64>>();
do_assert_unpin::<RwLockReadGuard<'_, i64>>();
do_assert_unpin::<RwLockWriteGuard<'_, i64>>();
do_assert_unpin::<broadcast::overflow::Sender<i64>>();
do_assert_unpin::<broadcast::overflow::Receiver<i64>>();
do_assert_unpin::<broadcast::overflow::RecvError>();
do_assert_unpin::<broadcast::overflow::TryRecvError>();
do_assert_unpin::<oneshot::Sender<i64>>();
do_assert_unpin::<oneshot::SendError<i64>>();
do_assert_unpin::<oneshot::Receiver<i64>>();
do_assert_unpin::<oneshot::Recv<i64>>();
do_assert_unpin::<mpsc::SendError<i64>>();
do_assert_unpin::<mpsc::UnboundedSender<i64>>();
do_assert_unpin::<mpsc::UnboundedReceiver<i64>>();
do_assert_unpin::<mpsc::BoundedSender<i64>>();
do_assert_unpin::<mpsc::BoundedReceiver<i64>>();
}
}
+358
View File
@@ -0,0 +1,358 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//! A bounded multi-producer, single-consumer queue for sending values between asynchronous
//! tasks with backpressure control.
use std::fmt;
use std::future::Future;
use std::future::poll_fn;
use std::pin::pin;
use std::sync::Arc;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::task::Context;
use std::task::Poll;
use std::task::Waker;
use crate::atomicbox::AtomicOptionBox;
use crate::internal::Acquire;
use crate::internal::Semaphore;
use crate::mpsc::RecvError;
use crate::mpsc::SendError;
use crate::mpsc::TryRecvError;
use crate::mpsc::error::TrySendError;
/// Creates a bounded mpsc channel for communicating between asynchronous
/// tasks with backpressure.
///
/// A `send` on this channel will wait if the buffer of the channel is full until a
/// `recv` is called on the receiver, which will consume the message and
/// free up space in the buffer.
#[track_caller]
pub fn bounded<T>(buffer: usize) -> (BoundedSender<T>, BoundedReceiver<T>) {
assert!(buffer > 0, "mpsc bounded channel requires buffer > 0");
let state = Arc::new(BoundedState {
senders: AtomicUsize::new(1),
tx_permits: Semaphore::new(0),
rx_task: AtomicOptionBox::none(),
});
let (sender, receiver) = std::sync::mpsc::sync_channel(buffer);
let sender = BoundedSender {
state: state.clone(),
sender: Some(sender),
};
let receiver = BoundedReceiver {
state: state.clone(),
receiver: Some(receiver),
};
(sender, receiver)
}
struct BoundedState {
senders: AtomicUsize,
tx_permits: Semaphore,
rx_task: AtomicOptionBox<Waker>,
}
/// Send values to the associated [`BoundedReceiver`].
///
/// Instances are created by the [`bounded`] function.
pub struct BoundedSender<T> {
state: Arc<BoundedState>,
sender: Option<std::sync::mpsc::SyncSender<T>>,
}
impl<T> Clone for BoundedSender<T> {
fn clone(&self) -> Self {
self.state.senders.fetch_add(1, Ordering::Release);
BoundedSender {
state: self.state.clone(),
sender: self.sender.clone(),
}
}
}
impl<T> fmt::Debug for BoundedSender<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("BoundedSender").finish_non_exhaustive()
}
}
impl<T> Drop for BoundedSender<T> {
fn drop(&mut self) {
// drop the sender; this closes the channel if it is the last sender
drop(self.sender.take());
match self.state.senders.fetch_sub(1, Ordering::AcqRel) {
1 => {
// If this is the last sender, we need to wake up the receiver so it can
// observe the disconnected state.
if let Some(waker) = self.state.rx_task.take() {
waker.wake();
}
}
_ => {
// there are still other senders left, do nothing
}
}
}
}
impl<T> BoundedSender<T> {
/// Attempts to send a message to the associated receiver.
///
/// This method will wait if the buffer of the channel is full until a `recv` is called on the
/// receiver, which will consume the message and free up space in the buffer.
///
/// If the receiver has been dropped, this function returns an error. The error includes
/// the value passed to `send`.
pub async fn send(&self, value: T) -> Result<(), SendError<T>> {
let value = match self.try_send(value) {
Ok(()) => return Ok(()),
Err(TrySendError::Disconnected(value)) => return Err(SendError::new(value)),
Err(TrySendError::Full(value)) => value,
};
struct SendState<'a, T> {
sender: &'a BoundedSender<T>,
value: Option<T>,
acquire: Acquire<'a>,
}
impl<T> SendState<'_, T> {
fn poll_send(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), SendError<T>>> {
let mut value = match self.value.take() {
Some(value) => value,
None => return Poll::Ready(Ok(())),
};
loop {
let poll = pin!(&mut self.acquire).poll(cx);
value = match self.sender.try_send(value) {
Ok(()) => return Poll::Ready(Ok(())),
Err(TrySendError::Disconnected(value)) => {
return Poll::Ready(Err(SendError::new(value)));
}
Err(TrySendError::Full(value)) => value,
};
if poll.is_ready() {
self.acquire = self.sender.state.tx_permits.poll_acquire(1);
} else {
self.value = Some(value);
return Poll::Pending;
}
}
}
}
let acquire = self.state.tx_permits.poll_acquire(1);
let mut send = SendState {
sender: self,
value: Some(value),
acquire,
};
poll_fn(|cx| send.poll_send(cx)).await
}
/// Attempts to send a message to the associated receiver without waiting.
///
/// This method returns the [`Full`] error if the buffer of the channel is full.
///
/// This method returns the [`Disconnected`] error if the channel is currently empty, and there
/// are no outstanding [receivers].
///
/// [`Full`]: TrySendError::Full
/// [`Disconnected`]: TrySendError::Disconnected
/// [receivers]: BoundedReceiver
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::mpsc::TrySendError;
/// use mea::mpsc::bounded;
/// let (tx, mut rx) = bounded::<i32>(1);
///
/// tx.try_send(1).unwrap();
/// assert_eq!(tx.try_send(2), Err(TrySendError::Full(2)));
///
/// drop(rx);
/// assert_eq!(tx.try_send(3), Err(TrySendError::Disconnected(3)));
/// # }
/// ```
pub fn try_send(&self, value: T) -> Result<(), TrySendError<T>> {
// SAFETY: The sender is guaranteed to be non-null before dropped.
let sender = self.sender.as_ref().unwrap();
match sender.try_send(value) {
Ok(()) => {
if let Some(waker) = self.state.rx_task.take() {
waker.wake();
}
Ok(())
}
Err(std::sync::mpsc::TrySendError::Full(value)) => Err(TrySendError::Full(value)),
Err(std::sync::mpsc::TrySendError::Disconnected(value)) => {
Err(TrySendError::Disconnected(value))
}
}
}
}
/// Receives values from the associated [`BoundedSender`].
///
/// Instances are created by the [`bounded`] function.
pub struct BoundedReceiver<T> {
state: Arc<BoundedState>,
receiver: Option<std::sync::mpsc::Receiver<T>>,
}
/// The only `!Sync` field `receiver` is protected by `&mut self` in `recv` and `try_recv`.
/// That is, `BoundedReceiver` can only be accessed by one thread at a time.
unsafe impl<T: Send> Sync for BoundedReceiver<T> {}
impl<T> fmt::Debug for BoundedReceiver<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("BoundedReceiver").finish_non_exhaustive()
}
}
impl<T> Drop for BoundedReceiver<T> {
fn drop(&mut self) {
drop(self.receiver.take());
self.state.tx_permits.notify_all();
}
}
impl<T> BoundedReceiver<T> {
/// Tries to receive the next value for this receiver and frees up a space in the buffer if
/// successful.
///
/// This method returns the [`Empty`] error if the channel is currently
/// empty, but there are still outstanding [senders].
///
/// This method returns the [`Disconnected`] error if the channel is
/// currently empty, and there are no outstanding [senders].
///
/// [`Empty`]: TryRecvError::Empty
/// [`Disconnected`]: TryRecvError::Disconnected
/// [senders]: BoundedSender
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::mpsc;
/// use mea::mpsc::TryRecvError;
/// let (tx, mut rx) = mpsc::bounded(2);
///
/// tx.send("hello").await.unwrap();
///
/// assert_eq!(Ok("hello"), rx.try_recv());
/// assert_eq!(Err(TryRecvError::Empty), rx.try_recv());
///
/// tx.send("hello").await.unwrap();
/// drop(tx);
///
/// assert_eq!(Ok("hello"), rx.try_recv());
/// assert_eq!(Err(TryRecvError::Disconnected), rx.try_recv());
/// # }
/// ```
pub fn try_recv(&mut self) -> Result<T, TryRecvError> {
// SAFETY: The receiver is guaranteed to be non-null before dropped.
let receiver = self.receiver.as_ref().unwrap();
match receiver.try_recv() {
Ok(v) => {
self.state.tx_permits.release_if_nonempty(1);
Ok(v)
}
Err(std::sync::mpsc::TryRecvError::Disconnected) => Err(TryRecvError::Disconnected),
Err(std::sync::mpsc::TryRecvError::Empty) => Err(TryRecvError::Empty),
}
}
/// Receives the next value for this receiver and frees up a space in the buffer if successful.
///
/// This method returns `Err(RecvError::Disconnected)` if the channel has been closed and there
/// are no remaining messages in the channel's buffer. This indicates that no further values
/// can ever be received from this `Receiver`. The channel is closed when all senders have been
/// dropped.
///
/// If there are no messages in the channel's buffer, but the channel has not yet been closed,
/// this method will sleep until a message is sent or the channel is closed.
///
/// # Cancel safety
///
/// This method is cancel safe. If `recv` is used as the event in a `select` statement
/// and some other branch completes first, it is guaranteed that no messages were received
/// on this channel.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::mpsc;
/// let (tx, mut rx) = mpsc::bounded(1);
///
/// tokio::spawn(async move {
/// tx.send("hello").await.unwrap();
/// });
///
/// assert_eq!(Ok("hello"), rx.recv().await);
/// assert_eq!(Err(mpsc::RecvError::Disconnected), rx.recv().await);
/// # }
/// ```
///
/// Values are buffered if the channel has enough capacity:
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::mpsc;
/// let (tx, mut rx) = mpsc::bounded(2);
///
/// tx.send("hello").await.unwrap();
/// tx.send("world").await.unwrap();
///
/// assert_eq!(Ok("hello"), rx.recv().await);
/// assert_eq!(Ok("world"), rx.recv().await);
/// # }
/// ```
pub async fn recv(&mut self) -> Result<T, RecvError> {
poll_fn(|cx| self.poll_recv(cx)).await
}
fn poll_recv(&mut self, cx: &mut Context<'_>) -> Poll<Result<T, RecvError>> {
match self.try_recv() {
Ok(v) => Poll::Ready(Ok(v)),
Err(TryRecvError::Disconnected) => Poll::Ready(Err(RecvError::Disconnected)),
Err(TryRecvError::Empty) => {
let waker = Some(Box::new(cx.waker().clone()));
self.state.rx_task.store(waker);
match self.try_recv() {
Ok(v) => Poll::Ready(Ok(v)),
Err(TryRecvError::Disconnected) => Poll::Ready(Err(RecvError::Disconnected)),
Err(TryRecvError::Empty) => Poll::Pending,
}
}
}
}
}
+146
View File
@@ -0,0 +1,146 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::any::type_name;
use std::fmt;
/// An error returned when trying to send on a closed channel.
///
/// Returned from [`UnboundedSender::send`] or [`BoundedSender::send`] if the
/// corresponding [`UnboundedReceiver`] or [`BoundedReceiver`] has already been
/// dropped.
///
/// The message that could not be sent can be retrieved again with
/// [`SendError::into_inner`].
///
/// [`UnboundedSender::send`]: crate::mpsc::UnboundedSender::send
/// [`BoundedSender::send`]: crate::mpsc::BoundedSender::send
/// [`UnboundedReceiver`]: crate::mpsc::UnboundedReceiver
/// [`BoundedReceiver`]: crate::mpsc::BoundedReceiver
#[derive(Clone, PartialEq, Eq)]
pub struct SendError<T>(T);
impl<T> SendError<T> {
/// Get a reference to the message that failed to be sent.
pub fn as_inner(&self) -> &T {
&self.0
}
/// Consumes the error and returns the message that failed to be sent.
pub fn into_inner(self) -> T {
self.0
}
/// Creates a new `SendError` with the given message.
pub(super) fn new(msg: T) -> SendError<T> {
SendError(msg)
}
}
impl<T> fmt::Display for SendError<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("sending on a closed channel")
}
}
impl<T> fmt::Debug for SendError<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "SendError<{}>(..)", type_name::<T>())
}
}
impl<T> std::error::Error for SendError<T> {}
/// Error returned by `try_send`.
#[derive(Clone, PartialEq, Eq)]
pub enum TrySendError<T> {
/// The channel is full, so data may not be sent at this time, but the receiver has not yet
/// disconnected.
Full(T),
/// The receiver has become disconnected, and there will never be any more data sent on it.
Disconnected(T),
}
impl<T> TrySendError<T> {
/// Gets a reference to the message that failed to be sent.
pub fn as_inner(&self) -> &T {
match self {
TrySendError::Full(msg) | TrySendError::Disconnected(msg) => msg,
}
}
/// Consumes the error and returns the message that failed to be sent.
pub fn into_inner(self) -> T {
match self {
TrySendError::Full(msg) | TrySendError::Disconnected(msg) => msg,
}
}
}
impl<T> fmt::Display for TrySendError<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(match self {
TrySendError::Full(_) => "sending on a full channel",
TrySendError::Disconnected(_) => "sending on a closed channel",
})
}
}
impl<T> fmt::Debug for TrySendError<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let ty = type_name::<T>();
match self {
TrySendError::Full(_) => write!(f, "TrySendError<{ty}>::Full(..)"),
TrySendError::Disconnected(_) => write!(f, "TrySendError<{ty}>::Disconnected(..)"),
}
}
}
impl<T> std::error::Error for TrySendError<T> {}
/// Error returned by `recv`.
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RecvError {
/// The sender has become disconnected, and there will never be any more data received on it.
Disconnected,
}
impl fmt::Display for RecvError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("receiving on a closed channel")
}
}
impl std::error::Error for RecvError {}
/// Error returned by `try_recv`.
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum TryRecvError {
/// This channel is currently empty, but the sender(s) have not yet disconnected, so data may
/// yet become available.
Empty,
/// The sender has become disconnected, and there will never be any more data received on it.
Disconnected,
}
impl fmt::Display for TryRecvError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(match self {
TryRecvError::Empty => "receiving on an empty channel",
TryRecvError::Disconnected => "receiving on a closed channel",
})
}
}
impl std::error::Error for TryRecvError {}
+32
View File
@@ -0,0 +1,32 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//! A multi-producer, single-consumer queue for sending values between asynchronous tasks.
mod bounded;
mod error;
#[cfg(test)]
mod tests;
mod unbounded;
pub use bounded::BoundedReceiver;
pub use bounded::BoundedSender;
pub use bounded::bounded;
pub use error::RecvError;
pub use error::SendError;
pub use error::TryRecvError;
pub use error::TrySendError;
pub use unbounded::UnboundedReceiver;
pub use unbounded::UnboundedSender;
pub use unbounded::unbounded;
+287
View File
@@ -0,0 +1,287 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::time::Instant;
use tokio_test::assert_ok;
use crate::mpsc;
use crate::mpsc::RecvError;
use crate::mpsc::TryRecvError;
use crate::mpsc::TrySendError;
use crate::test_runtime;
#[test]
fn test_unbounded_pressure() {
let n = 1024 * 1024;
let (tx, mut rx) = mpsc::unbounded();
test_runtime().block_on(async move {
let start = Instant::now();
tokio::spawn(async move {
for i in 0..n {
tx.send(i).unwrap();
}
});
for i in 0..n {
assert_eq!(rx.recv().await, Ok(i));
}
println!("Elapsed: {:?}", start.elapsed());
});
}
#[test]
fn test_unbounded_sum() {
let (tx, mut rx) = mpsc::unbounded();
test_runtime().block_on(async move {
for i in 0..100 {
let tx = tx.clone();
tokio::spawn(async move {
tx.send(i).unwrap();
});
}
drop(tx);
let mut sum = 0;
while let Ok(i) = rx.recv().await {
sum += i;
}
assert_eq!(sum, 4950);
});
}
#[tokio::test]
async fn select_streams() {
let (tx1, mut rx1) = mpsc::unbounded::<i32>();
let (tx2, mut rx2) = mpsc::unbounded::<i32>();
let (tx3, mut rx3) = mpsc::bounded(1);
let (tx4, mut rx4) = mpsc::bounded(1);
tokio::spawn(async move {
assert_ok!(tx2.send(1));
tokio::task::yield_now().await;
assert_ok!(tx1.send(2));
tokio::task::yield_now().await;
assert_ok!(tx2.send(3));
tokio::task::yield_now().await;
assert_ok!(tx3.send(4).await);
tokio::task::yield_now().await;
assert_ok!(tx4.send(5).await);
tokio::task::yield_now().await;
assert_ok!(tx3.send(6).await);
tokio::task::yield_now().await;
drop((tx1, tx2));
});
let mut rem = true;
let mut msgs = vec![];
let mut rx1_closed = false;
let mut rx2_closed = false;
let mut rx3_closed = false;
let mut rx4_closed = false;
while rem {
rem = !(rx1_closed && rx2_closed && rx3_closed && rx4_closed);
tokio::select! {
result = rx1.recv(), if !rx1_closed => {
match result {
Ok(x) => msgs.push(x),
Err(RecvError::Disconnected) => rx1_closed = true,
}
}
result = rx2.recv(), if !rx2_closed => {
match result {
Ok(y) => msgs.push(y),
Err(RecvError::Disconnected) => rx2_closed = true,
}
}
result = rx3.recv(), if !rx3_closed => {
match result {
Ok(z) => msgs.push(z),
Err(RecvError::Disconnected) => rx3_closed = true,
}
}
result = rx4.recv(), if !rx4_closed => {
match result {
Ok(w) => msgs.push(w),
Err(RecvError::Disconnected) => rx4_closed = true,
}
}
else => {
rx1_closed = true;
rx2_closed = true;
rx3_closed = true;
rx4_closed = true;
}
}
}
msgs.sort_unstable();
assert_eq!(&msgs[..], &[1, 2, 3, 4, 5, 6]);
}
#[tokio::test]
async fn send_recv_unbounded() {
let (tx, mut rx) = mpsc::unbounded::<i32>();
// Using `try_send`
assert_ok!(tx.send(1));
assert_ok!(tx.send(2));
assert_eq!(rx.recv().await, Ok(1));
assert_eq!(rx.recv().await, Ok(2));
drop(tx);
assert_eq!(rx.recv().await, Err(RecvError::Disconnected));
}
#[tokio::test]
async fn async_send_recv_unbounded() {
let (tx, mut rx) = mpsc::unbounded();
tokio::spawn(async move {
assert_ok!(tx.send(1));
assert_ok!(tx.send(2));
});
assert_eq!(Ok(1), rx.recv().await);
assert_eq!(Ok(2), rx.recv().await);
assert_eq!(Err(RecvError::Disconnected), rx.recv().await);
}
#[test]
fn try_recv_unbounded() {
for num in 0..100 {
let (tx, mut rx) = mpsc::unbounded();
for i in 0..num {
tx.send(i).unwrap();
}
for i in 0..num {
assert_eq!(rx.try_recv(), Ok(i));
}
assert_eq!(rx.try_recv(), Err(TryRecvError::Empty));
drop(tx);
assert_eq!(rx.try_recv(), Err(TryRecvError::Disconnected));
}
}
#[test]
fn try_recv_close_while_empty_unbounded() {
let (tx, mut rx) = mpsc::unbounded::<()>();
assert_eq!(Err(TryRecvError::Empty), rx.try_recv());
drop(tx);
assert_eq!(Err(TryRecvError::Disconnected), rx.try_recv());
}
#[tokio::test]
async fn send_recv_bounded() {
let (tx, mut rx) = mpsc::bounded(1);
tx.send(1).await.unwrap();
assert_eq!(rx.recv().await, Ok(1));
drop(tx);
assert_eq!(rx.recv().await, Err(RecvError::Disconnected));
}
#[tokio::test]
async fn async_send_recv_bounded() {
let (tx, mut rx) = mpsc::bounded(1);
tx.send(1).await.unwrap();
// This will block until the receiver is ready to receive.
tokio::spawn(async move {
tx.send(2).await.unwrap();
});
assert_eq!(Ok(1), rx.recv().await);
assert_eq!(Ok(2), rx.recv().await);
assert_eq!(Err(RecvError::Disconnected), rx.recv().await);
}
#[test]
fn try_send_recv_bounded() {
for num in 1..101 {
let (tx, mut rx) = mpsc::bounded(num);
for i in 0..num {
tx.try_send(i).unwrap();
}
assert_eq!(tx.try_send(num), Err(TrySendError::Full(num)));
for i in 0..num {
assert_eq!(rx.try_recv(), Ok(i));
}
assert_eq!(rx.try_recv(), Err(TryRecvError::Empty));
drop(tx);
assert_eq!(rx.try_recv(), Err(TryRecvError::Disconnected));
}
}
#[tokio::test]
async fn try_send_after_close_bounded() {
let (tx, rx) = mpsc::bounded(1);
tx.try_send(1).unwrap();
drop(rx);
assert_eq!(tx.try_send(3), Err(TrySendError::Disconnected(3)));
}
#[tokio::test]
async fn send_after_close_bounded() {
let (tx, mut rx) = mpsc::bounded(1);
tx.send(1).await.unwrap();
assert_eq!(rx.recv().await, Ok(1));
drop(rx);
assert_eq!(tx.send(2).await, Err(mpsc::SendError::new(2)));
}
#[test]
fn test_bounded_pressure() {
let n = 1024 * 1024;
let (tx, mut rx) = mpsc::bounded(1024);
test_runtime().block_on(async move {
let start = Instant::now();
tokio::spawn(async move {
for i in 0..n {
tx.send(i).await.unwrap();
}
});
for i in 0..n {
assert_eq!(rx.recv().await, Ok(i));
}
println!("Elapsed: {:?}", start.elapsed());
});
}
+257
View File
@@ -0,0 +1,257 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//! An unbounded multi-producer, single-consumer queue for sending values between asynchronous
//! tasks.
use std::fmt;
use std::future::poll_fn;
use std::sync::Arc;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::task::Context;
use std::task::Poll;
use std::task::Waker;
use crate::atomicbox::AtomicOptionBox;
use crate::mpsc::RecvError;
use crate::mpsc::SendError;
use crate::mpsc::TryRecvError;
/// Creates an unbounded mpsc channel for communicating between asynchronous
/// tasks without backpressure.
///
/// A `send` on this channel will always succeed as long as the receiver is alive.
/// If the receiver falls behind, messages will be arbitrarily buffered.
///
/// Note that the amount of available system memory is an implicit bound to
/// the channel. Using an `unbounded` channel has the ability of causing the
/// process to run out of memory. In this case, the process will be aborted.
pub fn unbounded<T>() -> (UnboundedSender<T>, UnboundedReceiver<T>) {
let state = Arc::new(UnboundedState {
senders: AtomicUsize::new(1),
rx_task: AtomicOptionBox::none(),
});
let (sender, receiver) = std::sync::mpsc::channel();
let sender = UnboundedSender {
state: state.clone(),
sender: Some(sender),
};
let receiver = UnboundedReceiver {
state: state.clone(),
receiver,
};
(sender, receiver)
}
struct UnboundedState {
senders: AtomicUsize,
rx_task: AtomicOptionBox<Waker>,
}
/// Send values to the associated [`UnboundedReceiver`].
///
/// Instances are created by the [`unbounded`] function.
pub struct UnboundedSender<T> {
state: Arc<UnboundedState>,
sender: Option<std::sync::mpsc::Sender<T>>,
}
impl<T> Clone for UnboundedSender<T> {
fn clone(&self) -> Self {
self.state.senders.fetch_add(1, Ordering::Release);
UnboundedSender {
state: self.state.clone(),
sender: self.sender.clone(),
}
}
}
impl<T> fmt::Debug for UnboundedSender<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("UnboundedSender").finish_non_exhaustive()
}
}
impl<T> Drop for UnboundedSender<T> {
fn drop(&mut self) {
// drop the sender; this closes the channel if it is the last sender
drop(self.sender.take());
match self.state.senders.fetch_sub(1, Ordering::AcqRel) {
1 => {
// If this is the last sender, we need to wake up the receiver so it can
// observe the disconnected state.
if let Some(waker) = self.state.rx_task.take() {
waker.wake();
}
}
_ => {
// there are still other senders left, do nothing
}
}
}
}
impl<T> UnboundedSender<T> {
/// Attempts to send a message without blocking.
///
/// This method is not marked async because sending a message to an unbounded channel
/// never requires any form of waiting. Because of this, the `send` method can be
/// used in both synchronous and asynchronous code without problems.
///
/// If the receiver has been dropped, this function returns an error. The error includes
/// the value passed to `send`.
pub fn send(&self, value: T) -> Result<(), SendError<T>> {
// SAFETY: The sender is guaranteed to be non-null before dropped.
let sender = self.sender.as_ref().unwrap();
sender.send(value).map_err(|err| SendError::new(err.0))?;
if let Some(waker) = self.state.rx_task.take() {
waker.wake();
}
Ok(())
}
}
/// Receive values from the associated [`UnboundedSender`].
///
/// Instances are created by the [`unbounded`] function.
pub struct UnboundedReceiver<T> {
state: Arc<UnboundedState>,
receiver: std::sync::mpsc::Receiver<T>,
}
/// The only `!Sync` field `receiver` is protected by `&mut self` in `recv` and `try_recv`.
/// That is, `UnboundedReceiver` can only be accessed by one thread at a time.
unsafe impl<T: Send> Sync for UnboundedReceiver<T> {}
impl<T> fmt::Debug for UnboundedReceiver<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("UnboundedReceiver").finish_non_exhaustive()
}
}
impl<T> UnboundedReceiver<T> {
/// Tries to receive the next value for this receiver.
///
/// This method returns the [`Empty`] error if the channel is currently
/// empty, but there are still outstanding [senders].
///
/// This method returns the [`Disconnected`] error if the channel is
/// currently empty, and there are no outstanding [senders].
///
/// [`Empty`]: TryRecvError::Empty
/// [`Disconnected`]: TryRecvError::Disconnected
/// [senders]: UnboundedSender
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::mpsc;
/// use mea::mpsc::TryRecvError;
/// let (tx, mut rx) = mpsc::unbounded();
///
/// tx.send("hello").unwrap();
///
/// assert_eq!(Ok("hello"), rx.try_recv());
/// assert_eq!(Err(TryRecvError::Empty), rx.try_recv());
///
/// tx.send("hello").unwrap();
/// drop(tx);
///
/// assert_eq!(Ok("hello"), rx.try_recv());
/// assert_eq!(Err(TryRecvError::Disconnected), rx.try_recv());
/// # }
/// ```
pub fn try_recv(&mut self) -> Result<T, TryRecvError> {
match self.receiver.try_recv() {
Ok(v) => Ok(v),
Err(std::sync::mpsc::TryRecvError::Disconnected) => Err(TryRecvError::Disconnected),
Err(std::sync::mpsc::TryRecvError::Empty) => Err(TryRecvError::Empty),
}
}
/// Receives the next value for this receiver.
///
/// This method returns `Err(RecvError::Disconnected)` if the channel has been closed and there
/// are no remaining messages in the channel's buffer. This indicates that no further values
/// can ever be received from this `Receiver`. The channel is closed when all senders have been
/// dropped.
///
/// If there are no messages in the channel's buffer, but the channel has not yet been closed,
/// this method will sleep until a message is sent or the channel is closed.
///
/// # Cancel safety
///
/// This method is cancel safe. If `recv` is used as the event in a `select` statement
/// and some other branch completes first, it is guaranteed that no messages were received
/// on this channel.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::mpsc;
/// let (tx, mut rx) = mpsc::unbounded();
///
/// tokio::spawn(async move {
/// tx.send("hello").unwrap();
/// });
///
/// assert_eq!(Ok("hello"), rx.recv().await);
/// assert_eq!(Err(mpsc::RecvError::Disconnected), rx.recv().await);
/// # }
/// ```
///
/// Values are buffered:
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::mpsc;
/// let (tx, mut rx) = mpsc::unbounded();
///
/// tx.send("hello").unwrap();
/// tx.send("world").unwrap();
///
/// assert_eq!(Ok("hello"), rx.recv().await);
/// assert_eq!(Ok("world"), rx.recv().await);
/// # }
/// ```
pub async fn recv(&mut self) -> Result<T, RecvError> {
poll_fn(|cx| self.poll_recv(cx)).await
}
fn poll_recv(&mut self, cx: &mut Context<'_>) -> Poll<Result<T, RecvError>> {
match self.try_recv() {
Ok(v) => Poll::Ready(Ok(v)),
Err(TryRecvError::Disconnected) => Poll::Ready(Err(RecvError::Disconnected)),
Err(TryRecvError::Empty) => {
let waker = Some(Box::new(cx.waker().clone()));
self.state.rx_task.store(waker);
match self.try_recv() {
Ok(v) => Poll::Ready(Ok(v)),
Err(TryRecvError::Disconnected) => Poll::Ready(Err(RecvError::Disconnected)),
Err(TryRecvError::Empty) => Poll::Pending,
}
}
}
}
}
File diff suppressed because it is too large Load Diff
+366
View File
@@ -0,0 +1,366 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::sync::Arc;
use super::*;
#[test]
fn test_try_lock_never_blocks() {
// Test that try_lock and try_lock_owned never block, even under contention
let mutex = Arc::new(Mutex::new(9));
let _guard = mutex.try_lock().unwrap();
let result = mutex.try_lock();
assert!(result.is_none());
let result = mutex.clone().try_lock_owned();
assert!(result.is_none());
}
#[test]
fn test_get_mut_provides_exclusive_access() {
// Test that get_mut provides direct access when we have exclusive ownership
let mut mutex = Mutex::new(11);
let data = mutex.get_mut();
*data = 100;
assert_eq!(*mutex.get_mut(), 100);
let inner = mutex.into_inner();
assert_eq!(inner, 100);
}
#[tokio::test]
async fn test_guard_map_preserves_lock() {
let data = (99i32, vec![1, 2, 3]);
let mutex = Mutex::new(data);
let guard = mutex.lock().await;
let mut mapped_guard = MutexGuard::map(guard, |data| &mut data.0);
assert!(mutex.try_lock().is_none());
*mapped_guard = 100;
// After dropping, mutex should be available
drop(mapped_guard);
let guard = mutex.try_lock().unwrap();
assert_eq!(guard.0, 100);
}
#[tokio::test]
async fn test_mapped_guard_holds_lock() {
// Test that MappedMutexGuard properly holds the lock even after the original guard is moved
let mutex = Arc::new(Mutex::new((10, 20)));
let guard = mutex.lock().await;
let mapped_guard = MutexGuard::map(guard, |data| &mut data.0);
assert!(
mutex.try_lock().is_none(),
"Lock should be held by the mapped guard"
);
assert_eq!(*mapped_guard, 10);
drop(mapped_guard);
assert!(
mutex.try_lock().is_some(),
"Lock should be released after mapped guard is dropped"
);
}
#[tokio::test]
async fn test_owned_mapped_guard_holds_lock() {
// Test that mapped owned guard properly holds the lock
let mutex = Arc::new(Mutex::new((30, 40)));
let owned_guard = mutex.clone().lock_owned().await;
let mapped_owned_guard = OwnedMutexGuard::map(owned_guard, |data| &mut data.1);
assert!(
mutex.try_lock().is_none(),
"Lock should be held by the mapped owned guard"
);
assert_eq!(*mapped_owned_guard, 40);
// When mapped owned guard is dropped, lock should be released
drop(mapped_owned_guard);
assert!(
mutex.try_lock().is_some(),
"Lock should be released after mapped owned guard is dropped"
);
}
#[tokio::test]
async fn test_guard_filter_map_failure() {
let data: Vec<i32> = vec![];
let mutex = Mutex::new(data);
let guard = mutex.lock().await;
let result = MutexGuard::filter_map(guard, |vec| vec.get_mut(0));
assert!(result.is_err());
if let Err(mut original_guard) = result {
original_guard.push(100);
assert_eq!(*original_guard, vec![100]);
} else {
panic!("Expected Err, but got Ok");
}
}
#[tokio::test]
async fn test_owned_guard_filter_map_failure() {
let data: Vec<i32> = vec![];
let mutex = Arc::new(Mutex::new(data));
let guard = mutex.clone().lock_owned().await;
let result = OwnedMutexGuard::filter_map(guard, |vec| vec.get_mut(0));
assert!(result.is_err());
if let Err(mut original_guard) = result {
original_guard.push(200);
assert_eq!(*original_guard, vec![200]);
} else {
panic!("Expected Err, but got Ok");
}
}
#[tokio::test]
async fn test_multiple_map_operations() {
// Test multiple consecutive map operations
let data = vec![vec![1, 2], vec![3, 4]];
let mutex = Mutex::new(data);
let guard = mutex.lock().await;
let first_vec = MutexGuard::map(guard, |data| &mut data[0]);
let mut first_element = MappedMutexGuard::map(first_vec, |vec| &mut vec[0]);
*first_element = 100;
drop(first_element);
let guard = mutex.lock().await;
assert_eq!(guard[0][0], 100);
assert_eq!(guard[0][1], 2);
assert_eq!(guard[1][0], 3);
}
#[tokio::test]
async fn test_stress() {
let mutex = Arc::new(Mutex::new(0));
let mut handles = Vec::new();
// Create many concurrent tasks
for i in 0..1000 {
let mutex = mutex.clone();
handles.push(tokio::spawn(async move {
let mut guard = mutex.lock().await;
*guard += 1;
if i % 10 == 0 {
tokio::task::yield_now().await;
}
}));
}
for handle in handles {
handle.await.unwrap();
}
let final_value = *mutex.lock().await;
assert_eq!(final_value, 1000);
}
#[tokio::test]
async fn test_guard_prevents_concurrent_access() {
// Test that holding a guard prevents other tasks from acquiring the lock
let mutex = Arc::new(Mutex::new(0));
let mutex_clone = mutex.clone();
let guard = mutex.lock().await;
assert!(
mutex.try_lock().is_none(),
"Lock should be held by the first guard"
);
let handle = tokio::spawn(async move {
let _guard2 = mutex_clone.lock().await;
123
});
tokio::task::yield_now().await;
assert!(
mutex.try_lock().is_none(),
"Lock should still be held after yielding"
);
drop(guard);
let result = handle.await.unwrap();
assert_eq!(result, 123);
assert!(
mutex.try_lock().is_some(),
"Lock should be available after all guards are dropped"
);
}
#[test]
fn test_lock_panic_safety() {
use std::panic::AssertUnwindSafe;
let mutex = Arc::new(Mutex::new(0));
let mutex_clone = mutex.clone();
let result = std::panic::catch_unwind(AssertUnwindSafe(move || {
let _guard = mutex_clone.try_lock().unwrap();
panic!("test panic");
}));
assert!(result.is_err());
// Lock should be released after panic
assert!(mutex.try_lock().is_some());
}
#[tokio::test]
async fn test_async_lock_panic_safety() {
// Test panic safety with async locks
let mutex = Arc::new(Mutex::new(0));
let mutex_clone = mutex.clone();
let handle = tokio::spawn(async move {
let _guard = mutex_clone.lock().await;
panic!("async test panic");
});
// panic
assert!(handle.await.is_err());
let guard = mutex.try_lock();
assert!(guard.is_some());
}
#[tokio::test]
async fn test_owned_guard_panic_safety() {
let mutex = Arc::new(Mutex::new(0));
let mutex_clone = mutex.clone();
let handle = tokio::spawn(async move {
let _guard = mutex_clone.clone().lock_owned().await;
panic!("owned guard panic");
});
assert!(handle.await.is_err());
// Lock should be available after the panicked task
let guard = mutex.try_lock();
assert!(guard.is_some());
}
#[tokio::test]
async fn test_mapped_guard_panic_safety() {
// Test panic safety with mapped guards
let mutex = Arc::new(Mutex::new((66, vec![1, 2, 3])));
let mutex_clone = mutex.clone();
let handle = tokio::spawn(async move {
let guard = mutex_clone.lock().await;
let _mapped = MutexGuard::map(guard, |data| &mut data.0);
panic!("mapped guard panic");
});
assert!(handle.await.is_err());
let guard = mutex.try_lock();
assert!(guard.is_some());
}
#[tokio::test]
async fn test_memory_ordering_correctness() {
// Test that mutex provides proper memory ordering guarantees
// When one task modifies data under mutex protection,
// another task should see the modification after acquiring the lock
let mutex = Arc::new(Mutex::new(vec![1, 2, 3]));
let mutex_clone = mutex.clone();
let handle = tokio::spawn(async move {
let mut guard = mutex_clone.lock().await;
guard.push(4);
guard[0] = 100;
// Lock is released when guard is dropped
});
handle.await.unwrap();
let guard = mutex.lock().await;
assert_eq!(*guard, vec![100, 2, 3, 4]);
// This test relies on mutex's acquire-release semantics to ensure
// that modifications made in the critical section are visible
// to subsequent lock acquisitions
}
#[tokio::test]
async fn test_mutex_zst() {
// Test that Mutex works correctly with Zero-Sized Types
let mutex = Arc::new(Mutex::new(()));
let mutex_clone = mutex.clone();
let handle = tokio::spawn(async move {
let guard = mutex_clone.lock().await;
*guard;
});
handle.await.unwrap();
// try_lock and owned guard should also work with ZST
let _owned_guard = mutex.clone().lock_owned().await;
assert!(mutex.try_lock().is_none());
drop(_owned_guard);
let guard = mutex.try_lock().unwrap();
*guard;
}
#[tokio::test]
async fn test_mapped_mutex_guard_send() {
// Test that MappedMutexGuard can be sent across await points
#[derive(Debug)]
struct TestStruct {
field1: i32,
_field2: String,
}
let mutex = Arc::new(Mutex::new(TestStruct {
field1: 2,
_field2: "oh".to_owned(),
}));
let mutex_clone = mutex.clone();
let handle = tokio::spawn(async move {
let guard = mutex_clone.lock().await;
let mapped_guard = MutexGuard::map(guard, |data| &mut data.field1);
tokio::task::yield_now().await;
*mapped_guard
});
let result = handle.await.unwrap();
assert_eq!(result, 2);
}
+32
View File
@@ -0,0 +1,32 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//! Asynchronous primitives for one-time async coordination.
//!
//! The module currently provides:
//!
//! * [`Once`]: A primitive that ensures a one-time asynchronous operation runs at most once, even
//! when called concurrently.
//! * [`OnceCell`]: A cell that can be written to at most once, storing a value produced
//! asynchronously.
//! * [`OnceMap`]: A hash map that runs computation only once for each key and stores the result.
#[allow(clippy::module_inception)]
mod once;
mod once_cell;
mod once_map;
pub use self::once::Once;
pub use self::once_cell::OnceCell;
pub use self::once_map::OnceMap;
+245
View File
@@ -0,0 +1,245 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::fmt;
use std::pin::Pin;
use std::task::Context;
use std::task::Poll;
use crate::internal::CountdownState;
use crate::semaphore::Semaphore;
#[cfg(test)]
mod tests;
/// A synchronization primitive which can be used to run a one-time async initialization.
///
/// Unlike [`std::sync::Once`], this type never blocks a thread. The provided closure must
/// produce a future and the future is awaited inside the primitive. Coordination happens
/// with asynchronous [`Semaphore`], which keeps the implementation runtime-agnostic.
///
/// This type also intentionally omits "poisoning" semantics. If an initialization future is
/// cancelled or panics, the attempt is abandoned and other tasks may retry the operation.
/// Encode partial-initialization detection in the future itself (e.g. return a `Result`)
/// when needed.
///
/// See the [module level documentation](super) for additional context.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::atomic::AtomicUsize;
/// use std::sync::atomic::Ordering;
///
/// use mea::once::Once;
///
/// static ONCE: Once = Once::new();
/// static COUNTER: AtomicUsize = AtomicUsize::new(0);
///
/// let handle1 = tokio::spawn(async {
/// ONCE.call_once(async || {
/// COUNTER.fetch_add(1, Ordering::SeqCst);
/// })
/// .await;
/// });
///
/// let handle2 = tokio::spawn(async {
/// ONCE.call_once(async || {
/// COUNTER.fetch_add(1, Ordering::SeqCst);
/// })
/// .await;
/// });
///
/// handle1.await.unwrap();
/// handle2.await.unwrap();
///
/// // The counter is incremented only once, even though two tasks called `call_once`.
/// assert_eq!(COUNTER.load(Ordering::SeqCst), 1);
/// # }
/// ```
pub struct Once {
done: CountdownState,
semaphore: Semaphore,
}
impl Default for Once {
fn default() -> Self {
Self::new()
}
}
impl fmt::Debug for Once {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let done = self.is_completed();
f.debug_struct("Once").field("done", &done).finish()
}
}
impl Once {
/// Creates a new `Once` instance.
///
/// # Examples
///
/// ```
/// use mea::once::Once;
///
/// static ONCE: Once = Once::new();
/// ```
pub const fn new() -> Self {
Self {
done: CountdownState::new(1),
semaphore: Semaphore::new(1),
}
}
/// Returns `true` if some `call_once` has completed successfully.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::once::Once;
///
/// static ONCE: Once = Once::new();
///
/// assert!(!ONCE.is_completed());
///
/// ONCE.call_once(async || {}).await;
///
/// assert!(ONCE.is_completed());
/// # }
/// ```
pub fn is_completed(&self) -> bool {
self.done.spin_wait(0).is_ok()
}
/// Calls the given async closure if this is the first time `call_once` has been called
/// on this `Once` instance.
///
/// If another task is currently running the closure, this call will wait for that task
/// to complete.
///
/// If the provided operation is cancelled, the initialization attempt is cancelled. If there
/// are other tasks waiting, one of them will start another attempt.
///
/// Calling call_once recursively on the same Once from within the closure will deadlock,
/// because the closure holds the semaphore permit while trying to acquire it again.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::once::Once;
///
/// static ONCE: Once = Once::new();
///
/// ONCE.call_once(async || {
/// println!("Do some one-time async thing.");
/// })
/// .await;
/// # }
/// ```
pub async fn call_once<F>(&self, f: F)
where
F: AsyncFnOnce(),
{
if self.is_completed() {
return;
}
let _permit = self.semaphore.acquire(1).await;
if self.is_completed() {
// double-checked: another task completed the initialization while we waited.
return;
}
f().await;
if let Err(cnt) = self.done.cas_state(1, 0) {
unreachable!("[BUG] Once completed more than once: {}", cnt);
}
self.done.wake_all();
}
/// Waits asynchronously until some `call_once` has completed successfully.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::once::Once;
///
/// static ONCE: Once = Once::new();
///
/// let handle = tokio::spawn(async {
/// ONCE.wait().await;
/// });
///
/// ONCE.call_once(async || {
/// println!("initialized");
/// })
/// .await;
///
/// handle.await.unwrap();
/// # }
/// ```
pub async fn wait(&self) {
if self.is_completed() {
return;
}
OnceWait {
idx: None,
once: self,
}
.await
}
}
struct OnceWait<'a> {
idx: Option<usize>,
once: &'a Once,
}
impl fmt::Debug for OnceWait<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("OnceWait").finish_non_exhaustive()
}
}
impl Future for OnceWait<'_> {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let Self { idx, once } = self.get_mut();
// register waker if the counter is not zero
if once.done.spin_wait(16).is_err() {
once.done.register_waker(idx, cx);
// double check after register waker, to catch the update between two steps
if !once.is_completed() {
return Poll::Pending;
};
}
Poll::Ready(())
}
}
+179
View File
@@ -0,0 +1,179 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::sync::Arc;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::time::Duration;
use tokio_test::assert_ready;
use super::*;
use crate::latch::Latch;
use crate::test_runtime;
#[tokio::test]
async fn test_call_once_runs_only_once() {
static ONCE: Once = Once::new();
static COUNTER: AtomicUsize = AtomicUsize::new(0);
assert!(!ONCE.is_completed());
ONCE.call_once(async || {
COUNTER.fetch_add(1, Ordering::SeqCst);
})
.await;
assert!(ONCE.is_completed());
assert_eq!(COUNTER.load(Ordering::SeqCst), 1);
// Second call should not run the closure
ONCE.call_once(async || {
COUNTER.fetch_add(1, Ordering::SeqCst);
})
.await;
assert_eq!(COUNTER.load(Ordering::SeqCst), 1);
}
#[test]
fn test_once_multi_task() {
static ONCE: Once = Once::new();
static COUNTER: AtomicUsize = AtomicUsize::new(0);
test_runtime().block_on(async {
const N: usize = 100;
let latch = Arc::new(Latch::new(N as u32));
let mut handles = Vec::with_capacity(N);
for _ in 0..N {
let latch = latch.clone();
handles.push(tokio::spawn(async move {
ONCE.call_once(async || {
COUNTER.fetch_add(1, Ordering::SeqCst);
})
.await;
latch.count_down();
}));
}
latch.wait().await;
for handle in handles {
handle.await.unwrap();
}
// Only one task should have incremented the counter
assert_eq!(COUNTER.load(Ordering::SeqCst), 1);
assert!(ONCE.is_completed());
});
}
#[tokio::test]
async fn test_once_cancelled() {
static ONCE: Once = Once::new();
static COUNTER: AtomicUsize = AtomicUsize::new(0);
let handle1 = tokio::spawn(async {
let fut = ONCE.call_once(async || {
tokio::time::sleep(Duration::from_millis(1000)).await;
COUNTER.fetch_add(1, Ordering::SeqCst);
});
let timeout = tokio::time::timeout(Duration::from_millis(1), fut).await;
assert!(timeout.is_err());
});
let handle2 = tokio::spawn(async {
tokio::time::sleep(Duration::from_millis(100)).await;
ONCE.call_once(async || {
COUNTER.fetch_add(10, Ordering::SeqCst);
})
.await;
});
handle1.await.unwrap();
handle2.await.unwrap();
// The second task should have run since the first was cancelled
assert_eq!(COUNTER.load(Ordering::SeqCst), 10);
assert!(ONCE.is_completed());
}
#[tokio::test]
async fn test_once_debug() {
let once = Once::new();
let debug_str = format!("{:?}", once);
assert!(debug_str.contains("Once"));
assert!(debug_str.contains("done"));
assert!(debug_str.contains("false"));
once.call_once(async || {}).await;
let debug_str = format!("{:?}", once);
assert!(debug_str.contains("true"));
}
#[tokio::test]
async fn test_once_default() {
let once = Once::default();
assert!(!once.is_completed());
}
#[tokio::test]
async fn test_once_retry_after_panic() {
static ONCE: Once = Once::new();
static COUNTER: AtomicUsize = AtomicUsize::new(0);
let handle = tokio::spawn(async {
ONCE.call_once(async || {
COUNTER.fetch_add(1, Ordering::SeqCst);
panic!("boom");
})
.await;
});
let err = handle.await.expect_err("once init should panic");
assert!(err.is_panic());
ONCE.call_once(async || {
COUNTER.fetch_add(1, Ordering::SeqCst);
})
.await;
assert_eq!(COUNTER.load(Ordering::SeqCst), 2);
assert!(ONCE.is_completed());
}
#[tokio::test]
async fn test_once_wait() {
// wait after call_once completed
{
let once = Once::new();
once.call_once(async || {}).await;
assert_ready!(tokio_test::task::spawn(once.wait()).poll());
}
// wait before call_once completed
{
static ONCE: Once = Once::new();
let handle = tokio::spawn(async {
ONCE.wait().await;
});
tokio::time::sleep(Duration::from_millis(100)).await;
ONCE.call_once(async || {}).await;
handle.await.unwrap();
}
}
+474
View File
@@ -0,0 +1,474 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::cell::UnsafeCell;
use std::convert::Infallible;
use std::fmt;
use std::mem::MaybeUninit;
use std::sync::atomic::AtomicBool;
use std::sync::atomic::Ordering;
use crate::semaphore::Semaphore;
use crate::semaphore::SemaphorePermit;
#[cfg(test)]
mod tests;
/// A thread-safe cell which can nominally be written to only once.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::once::OnceCell;
///
/// static CELL: OnceCell<u8> = OnceCell::new();
///
/// let handle1 = tokio::spawn(async { CELL.get_or_init(move || async { 1 }).await });
/// let handle2 = tokio::spawn(async { CELL.get_or_init(move || async { 2 }).await });
/// let result1 = handle1.await.unwrap();
/// let result2 = handle2.await.unwrap();
/// println!("Results: {}, {}", result1, result2);
/// # }
/// ```
///
/// The outputs must be either `Results: 1, 1` or `Results: 2, 2`, i.e. once the value is set via
/// an asynchronous function, the value inside the `OnceCell` will be immutable.
pub struct OnceCell<T> {
value_set: AtomicBool,
value: UnsafeCell<MaybeUninit<T>>,
semaphore: Semaphore,
}
// SAFETY: OnceCell<T> can be shared between threads as long as T is Sync + Send.
unsafe impl<T: Sync + Send> Sync for OnceCell<T> {}
// SAFETY: OnceCell<T> can be sent between threads as long as T is Send.
unsafe impl<T: Send> Send for OnceCell<T> {}
impl<T> Default for OnceCell<T> {
fn default() -> Self {
Self::new()
}
}
impl<T> OnceCell<T> {
/// Creates a new empty `OnceCell`.
pub const fn new() -> Self {
Self {
value_set: AtomicBool::new(false),
value: UnsafeCell::new(MaybeUninit::uninit()),
semaphore: Semaphore::new(1),
}
}
/// Creates a new `OnceCell` initialized with the provided value.
pub const fn from_value(value: T) -> Self {
Self {
value_set: AtomicBool::new(true),
value: UnsafeCell::new(MaybeUninit::new(value)),
semaphore: Semaphore::new(1),
}
}
/// Returns whether the internal value is set.
fn initialized(&self) -> bool {
self.value_set.load(Ordering::Acquire)
}
/// Returns whether the internal value is set.
fn initialized_mut(&mut self) -> bool {
*self.value_set.get_mut()
}
/// Gets the reference to the underlying value.
///
/// Returns `None` if the cell is uninitialized, or being initialized.
///
/// This method never blocks.
pub fn get(&self) -> Option<&T> {
if self.initialized() {
Some(unsafe { self.get_unchecked() })
} else {
None
}
}
/// Gets the mutable reference to the underlying value.
///
/// Returns `None` if the cell is uninitialized.
///
/// This method never blocks. Since it borrows the `OnceCell` mutably, it is statically
/// guaranteed that no active borrows to the `OnceCell` exist, including from other threads.
pub fn get_mut(&mut self) -> Option<&mut T> {
if self.initialized_mut() {
Some(unsafe { self.get_unchecked_mut() })
} else {
None
}
}
/// Gets the reference to the internal value, initializing it with the provided asynchronous
/// function if it is not set yet.
///
/// If some other task is currently working on initializing the `OnceCell`, this call will wait
/// for that other task to finish, then return the value that the other task produced.
///
/// If the provided operation is cancelled, the initialization attempt is cancelled. If there
/// are other tasks waiting for the value to be initialized, one of them will start another
/// attempt at initializing the value.
///
/// This will deadlock if `init` tries to initialize the cell recursively.
pub async fn get_or_init<F>(&self, init: F) -> &T
where
F: AsyncFnOnce() -> T,
{
match self
.get_or_try_init(async || Ok::<T, Infallible>(init().await))
.await
{
Ok(val) => val,
}
}
/// Gets the reference to the internal value, initializing it with the provided asynchronous
/// function if it is not set yet.
///
/// If some other task is currently working on initializing the `OnceCell`, this call will wait
/// for that other task to finish, then return the value that the other task produced.
///
/// If the provided operation returns an error, is cancelled or panics, the initialization
/// attempt is cancelled. If there are other tasks waiting for the value to be initialized
/// one of them will start another attempt at initializing the value.
///
/// This will deadlock if `init` tries to initialize the cell recursively.
pub async fn get_or_try_init<E, F>(&self, init: F) -> Result<&T, E>
where
F: AsyncFnOnce() -> Result<T, E>,
{
if let Some(v) = self.get() {
return Ok(v);
}
let permit = self.semaphore.acquire(1).await;
if let Some(v) = self.get() {
// double-checked: another task initialized the value
// while we were waiting for the permit
return Ok(v);
}
let value = init().await?;
Ok(self.set_value(value, permit))
}
/// Gets a mutable reference to the internal value, initializing it with the provided
/// asynchronous function if it is not set yet.
///
/// This method never blocks other tasks because it takes `&mut self`, which guarantees
/// exclusive access to the `OnceCell` and thus no concurrent initialization can be in
/// progress.
///
/// If the cell is already initialized, it returns a mutable reference to the existing value.
/// Otherwise, it runs `init`, stores the result, and returns a mutable reference to the newly
/// initialized value.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::once::OnceCell;
///
/// let mut cell: OnceCell<u32> = OnceCell::new();
/// let v = cell.get_mut_or_init(|| async { 41 }).await;
/// *v += 1;
/// assert_eq!(*cell.get().unwrap(), 42);
/// # }
/// ```
pub async fn get_mut_or_init<F>(&mut self, init: F) -> &mut T
where
F: AsyncFnOnce() -> T,
{
match self
.get_mut_or_try_init(async || Ok::<T, Infallible>(init().await))
.await
{
Ok(val) => val,
}
}
/// Gets a mutable reference to the internal value, initializing it with the provided
/// asynchronous function that may fail if it is not set yet.
///
/// This method never blocks other tasks because it takes `&mut self`, which guarantees
/// exclusive access to the `OnceCell` and thus no concurrent initialization can be in
/// progress.
///
/// If the cell is already initialized, it returns a mutable reference to the existing value.
/// Otherwise, it runs `init`. On success, it stores the result and returns a mutable
/// reference to the newly initialized value. On error, it returns the error and leaves the
/// cell uninitialized.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::once::OnceCell;
///
/// let mut cell: OnceCell<u32> = OnceCell::new();
/// assert!(
/// cell.get_mut_or_try_init(|| async { Err(()) })
/// .await
/// .is_err()
/// );
/// let v = cell
/// .get_mut_or_try_init(|| async { Ok::<_, ()>(10) })
/// .await
/// .unwrap();
/// *v += 5;
/// assert_eq!(*cell.get().unwrap(), 15);
/// # }
/// ```
pub async fn get_mut_or_try_init<E, F>(&mut self, init: F) -> Result<&mut T, E>
where
F: AsyncFnOnce() -> Result<T, E>,
{
// Workaround if let Some(v) = self.get_mut() { return Ok(v); }
// @see https://github.com/rust-lang/rust/issues/51545
if self.initialized_mut() {
return Ok(unsafe { self.get_unchecked_mut() });
}
let value = init().await?;
Ok(self.set_value_mut(value))
}
/// Initializes the contents of the cell to `value` if the cell was uninitialized,
/// then returns a reference to it.
///
/// May wait if another thread is currently attempting to initialize the cell. The cell is
/// guaranteed to contain a value when `try_insert` returns, though not necessarily the
/// one provided.
///
/// Returns `Ok(&value)` if the cell was uninitialized and `Err((&current_value, value))`
/// if it was already initialized.
///
/// # Examples
///
/// ```
/// use mea::once::OnceCell;
///
/// static CELL: OnceCell<i32> = OnceCell::new();
///
/// # #[tokio::main]
/// # async fn main() {
/// assert!(CELL.get().is_none());
///
/// tokio::spawn(async {
/// assert_eq!(CELL.try_insert(92).await, Ok(&92));
/// })
/// .await
/// .unwrap();
///
/// assert_eq!(CELL.try_insert(62).await, Err((&92, 62)));
/// assert_eq!(CELL.get(), Some(&92));
/// # }
/// ```
pub async fn try_insert(&self, value: T) -> Result<&T, (&T, T)> {
let mut value = Some(value);
let res = self.get_or_init(async || value.take().unwrap()).await;
match value {
None => Ok(res),
Some(value) => Err((res, value)),
}
}
/// Initializes the contents of the cell to `value`.
///
/// May wait if another thread is currently attempting to initialize the cell. The cell is
/// guaranteed to contain a value when `set` returns, though not necessarily the one provided.
///
/// Returns `Ok(())` if the cell was uninitialized and `Err(value)` if the cell was already
/// initialized.
///
/// # Examples
///
/// ```
/// use mea::once::OnceCell;
///
/// static CELL: OnceCell<i32> = OnceCell::new();
///
/// # #[tokio::main]
/// # async fn main() {
/// assert!(CELL.get().is_none());
///
/// tokio::spawn(async {
/// assert_eq!(CELL.set(92).await, Ok(()));
/// })
/// .await
/// .unwrap();
///
/// assert_eq!(CELL.set(62).await, Err(62));
/// assert_eq!(CELL.get(), Some(&92));
/// # }
/// ```
pub async fn set(&self, value: T) -> Result<(), T> {
match self.try_insert(value).await {
Ok(_) => Ok(()),
Err((_, value)) => Err(value),
}
}
/// Consumes the `OnceCell`, returning the wrapped value. Returns `None` if the cell was
/// uninitialized.
///
/// # Examples
///
/// ```
/// use mea::once::OnceCell;
///
/// # #[tokio::main]
/// # async fn main() {
/// let cell: OnceCell<String> = OnceCell::new();
/// assert_eq!(cell.into_inner(), None);
///
/// let cell = OnceCell::new();
/// cell.set("hello".to_string()).await.unwrap();
/// assert_eq!(cell.into_inner(), Some("hello".to_string()));
/// # }
/// ```
pub fn into_inner(mut self) -> Option<T> {
if self.initialized_mut() {
// set to uninitialized for the destructor of `OnceCell` to work properly
*self.value_set.get_mut() = false;
Some(unsafe { self.value.get_mut().assume_init_read() })
} else {
None
}
}
/// Takes the value out of this `OnceCell`, moving it back to an uninitialized state.
///
/// Has no effect and returns `None` if the `OnceCell` was uninitialized.
///
/// Since this method borrows the `OnceCell` mutably, it is statically guaranteed that
/// no active borrows to the `OnceCell` exist, including from other threads.
///
/// # Examples
///
/// ```
/// use mea::once::OnceCell;
///
/// # #[tokio::main]
/// # async fn main() {
/// let mut cell: OnceCell<String> = OnceCell::new();
/// assert_eq!(cell.take(), None);
///
/// let mut cell = OnceCell::new();
/// cell.set("hello".to_string()).await.unwrap();
/// assert_eq!(cell.take(), Some("hello".to_string()));
/// assert_eq!(cell.get(), None);
/// # }
/// ```
pub fn take(&mut self) -> Option<T> {
std::mem::take(self).into_inner()
}
/// # Safety
///
/// The cell must be initialized
#[inline]
unsafe fn get_unchecked(&self) -> &T {
debug_assert!(self.initialized());
unsafe { (&*self.value.get()).assume_init_ref() }
}
/// # Safety
///
/// The cell must be initialized
#[inline]
unsafe fn get_unchecked_mut(&mut self) -> &mut T {
debug_assert!(self.initialized_mut());
unsafe { (&mut *self.value.get()).assume_init_mut() }
}
fn set_value(&self, value: T, permit: SemaphorePermit<'_>) -> &T {
// Hold the permit to ensure exclusive access.
let _permit = permit;
let value_ptr = self.value.get();
unsafe { value_ptr.write(MaybeUninit::new(value)) };
// Use `store` with `Release` ordering to ensure that when loading it with `Acquire`
// ordering, the initialized value is visible.
self.value_set.store(true, Ordering::Release);
// SAFETY: value initialized above
unsafe { self.get_unchecked() }
}
fn set_value_mut(&mut self, value: T) -> &mut T {
let value = self.value.get_mut().write(value);
*self.value_set.get_mut() = true;
value
}
}
impl<T> Drop for OnceCell<T> {
fn drop(&mut self) {
if self.initialized_mut() {
// SAFETY: The cell is initialized and being dropped, so it can't be accessed again.
unsafe { self.value.get_mut().assume_init_drop() };
}
}
}
impl<T: fmt::Debug> fmt::Debug for OnceCell<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let mut d = f.debug_tuple("OnceCell");
match self.get() {
Some(v) => d.field(v),
None => d.field(&format_args!("<uninit>")),
};
d.finish()
}
}
impl<T: Clone> Clone for OnceCell<T> {
fn clone(&self) -> OnceCell<T> {
match self.get() {
Some(v) => OnceCell::from_value(v.clone()),
None => OnceCell::new(),
}
}
}
impl<T> From<T> for OnceCell<T> {
fn from(value: T) -> Self {
OnceCell::from_value(value)
}
}
impl<T: PartialEq> PartialEq for OnceCell<T> {
fn eq(&self, other: &Self) -> bool {
self.get() == other.get()
}
}
impl<T: Eq> Eq for OnceCell<T> {}
+202
View File
@@ -0,0 +1,202 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::sync::Arc;
use std::sync::atomic::AtomicBool;
use std::sync::atomic::Ordering;
use std::time::Duration;
use tokio::sync::Mutex;
use super::*;
use crate::latch::Latch;
struct Foo {
value: Arc<AtomicBool>,
}
impl Foo {
async fn new(value: Arc<AtomicBool>) -> Self {
// simulate some async initialization work
tokio::time::sleep(Duration::from_millis(100)).await;
Foo { value }
}
fn value(&self) -> bool {
self.value.load(Ordering::Acquire)
}
}
impl Drop for Foo {
fn drop(&mut self) {
self.value.store(true, Ordering::Release);
}
}
#[tokio::test]
async fn drop_cell() {
let dropped = Arc::new(AtomicBool::new(false));
{
let cell = OnceCell::new();
assert!(cell.get().is_none());
let state = dropped.clone();
cell.get_or_init(|| async { Foo::new(state).await }).await;
let foo = cell.get().unwrap();
assert!(!foo.value());
assert!(!dropped.load(Ordering::Acquire));
}
assert!(dropped.load(Ordering::Acquire));
}
#[test]
fn multi_init() {
let rt = tokio::runtime::Builder::new_multi_thread()
.worker_threads(4)
.build()
.unwrap();
static CELL: OnceCell<usize> = OnceCell::new();
rt.block_on(async {
const N: usize = 100;
let latch = Arc::new(Latch::new(N as u32));
let values = Arc::new(Mutex::new(vec![0; N]));
for i in 0..N {
let latch = latch.clone();
let values = values.clone();
rt.spawn(async move {
let result = CELL.get_or_init(move || async move { i + 1000 }).await;
let mut values = values.lock().await;
values[i] = *result;
latch.count_down();
});
}
latch.wait().await;
let cell_value = CELL.get().unwrap();
for (index, value) in values.lock().await.iter().enumerate() {
assert_eq!(*value, *cell_value, "mismatch at index {index}");
}
});
}
#[tokio::test]
async fn init_cancelled() {
static CELL: OnceCell<u8> = OnceCell::new();
let handle1 = tokio::spawn(async {
let fut = CELL.get_or_init(|| async {
tokio::time::sleep(Duration::from_millis(1000)).await;
1
});
let timeout = tokio::time::timeout(Duration::from_millis(1), fut).await;
assert!(timeout.is_err());
});
let handle2 = tokio::spawn(async {
tokio::time::sleep(Duration::from_millis(100)).await;
let value = CELL.get_or_init(|| async { 2 }).await;
assert_eq!(*value, 2);
});
handle1.await.unwrap();
handle2.await.unwrap();
}
#[tokio::test]
async fn init_error() {
{
static CELL: OnceCell<u8> = OnceCell::new();
let handle1 = tokio::spawn(async {
let result = CELL.get_or_try_init(|| async { Err(()) }).await;
assert!(result.is_err());
});
let handle2 = tokio::spawn(async {
tokio::time::sleep(Duration::from_millis(100)).await;
let value = CELL.get_or_try_init(|| async { Ok::<_, ()>(2) }).await;
assert_eq!(*value.unwrap(), 2);
});
handle1.await.unwrap();
handle2.await.unwrap();
}
{
static CELL: OnceCell<u8> = OnceCell::new();
let handle1 = tokio::spawn(async {
let value = CELL.get_or_try_init(|| async { Ok::<_, ()>(2) }).await;
assert_eq!(*value.unwrap(), 2);
});
let handle2 = tokio::spawn(async {
tokio::time::sleep(Duration::from_millis(100)).await;
let value = CELL.get_or_try_init(|| async { Err(()) }).await;
assert_eq!(*value.unwrap(), 2);
});
handle1.await.unwrap();
handle2.await.unwrap();
}
}
#[tokio::test]
async fn get_mut_or_init() {
let mut cell: OnceCell<u32> = OnceCell::new();
let v = cell
.get_mut_or_init(async || {
tokio::time::sleep(Duration::from_millis(1)).await;
41
})
.await;
*v += 1;
let v = tokio::spawn(async move { *cell.get_or_init(async || 0).await })
.await
.unwrap();
assert_eq!(v, 42);
}
#[tokio::test]
async fn get_mut_or_try_init() {
let mut cell: OnceCell<u32> = OnceCell::new();
let r = cell
.get_mut_or_try_init(async || {
tokio::time::sleep(Duration::from_millis(1)).await;
Err(())
})
.await;
assert!(r.is_err());
assert_eq!(cell.get_mut(), None);
let v = tokio::spawn(async move {
let v = cell
.get_mut_or_try_init(async || Ok::<_, ()>(10))
.await
.unwrap();
*v += 5;
*v
})
.await
.unwrap();
assert_eq!(v, 15);
}
+191
View File
@@ -0,0 +1,191 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::borrow::Borrow;
use std::collections::HashMap;
use std::hash::BuildHasher;
use std::hash::Hash;
use std::hash::RandomState;
use std::sync::Arc;
use crate::internal::Mutex;
use crate::once::OnceCell;
#[cfg(test)]
mod tests;
/// A hash map that runs computation only once for each key and stores the result.
///
/// Note that this always clones the value out of the underlying map. Because of this, it's common
/// to wrap the `V` in an `Arc<V>` to make cloning cheap.
#[derive(Debug)]
pub struct OnceMap<K, V, S = RandomState> {
map: Mutex<HashMap<K, Arc<OnceCell<V>>, S>>,
}
impl<K, V, S> Default for OnceMap<K, V, S>
where
K: Eq + Hash,
V: Clone,
S: BuildHasher + Clone + Default,
{
fn default() -> Self {
Self::with_hasher(S::default())
}
}
impl<K, V> OnceMap<K, V, RandomState>
where
K: Eq + Hash,
V: Clone,
{
/// Creates a new OnceMap with the default hasher.
pub fn new() -> Self {
Self {
map: Mutex::new(HashMap::new()),
}
}
/// Creates a new OnceMap with the default hasher and the specified capacity.
pub fn with_capacity(capacity: usize) -> Self {
Self {
map: Mutex::new(HashMap::with_capacity(capacity)),
}
}
}
impl<K, V, S> OnceMap<K, V, S>
where
K: Eq + Hash,
V: Clone,
S: BuildHasher + Clone,
{
/// Creates a new OnceMap with the given hasher.
pub fn with_hasher(hasher: S) -> Self {
Self {
map: Mutex::new(HashMap::with_hasher(hasher)),
}
}
/// Create a OnceMap with the specified capacity and hasher.
pub fn with_capacity_and_hasher(capacity: usize, hasher: S) -> Self {
Self {
map: Mutex::new(HashMap::with_capacity_and_hasher(capacity, hasher)),
}
}
/// Compute the value for the given key if absent.
///
/// If the value for the key is already being computed by another task, this task will wait for
/// the computation to finish and return the result.
pub async fn compute<F>(&self, key: K, func: F) -> V
where
F: AsyncFnOnce() -> V,
{
// 1. Get or create the OnceCell.
let cell = {
let mut map = self.map.lock();
map.entry(key)
.or_insert_with(|| Arc::new(OnceCell::new()))
.clone()
};
// 2. Try to initialize the cell.
// OnceCell::get_or_init guarantees that only one task executes the closure.
let res = cell.get_or_init(func).await;
res.clone()
}
/// Compute the value for the given key if absent.
///
/// If the value for the key is already being computed by another task, this task will wait for
/// the computation to finish and return the result.
///
/// If the computation fails, the error is returned and the value is not stored. Other tasks
/// waiting for the value will retry the computation.
pub async fn try_compute<E, F>(&self, key: K, func: F) -> Result<V, E>
where
F: AsyncFnOnce() -> Result<V, E>,
{
// 1. Get or create the OnceCell.
let cell = {
let mut map = self.map.lock();
map.entry(key)
.or_insert_with(|| Arc::new(OnceCell::new()))
.clone()
};
// 2. Try to initialize the cell.
// OnceCell::get_or_try_init guarantees that only one task executes the closure.
let res = cell.get_or_try_init(func).await?;
Ok(res.clone())
}
/// Get a clone of the value for the given key if exists.
pub fn get<Q>(&self, key: &Q) -> Option<V>
where
K: Borrow<Q>,
Q: Hash + Eq + ?Sized,
{
let map = self.map.lock();
let cell = map.get(key)?;
cell.get().cloned()
}
/// Remove the given key from the map.
///
/// If you need to get the value that has been removed, use the [`remove`] method instead.
///
/// [`remove`]: Self::remove
pub fn discard<Q>(&self, key: &Q)
where
K: Borrow<Q>,
Q: Hash + Eq + ?Sized,
{
let mut map = self.map.lock();
map.remove(key);
}
/// Remove the given key from the map and return a *clone* of the value if exists.
///
/// If you do not need to get the value that has been removed, use the [`discard`] method
/// instead.
///
/// [`discard`]: Self::discard
pub fn remove<Q>(&self, key: &Q) -> Option<V>
where
K: Borrow<Q>,
Q: Hash + Eq + ?Sized,
{
let cell = self.map.lock().remove(key)?;
cell.get().cloned()
}
}
impl<K, V, S> FromIterator<(K, V)> for OnceMap<K, V, S>
where
K: Eq + Hash + Clone,
V: Clone,
S: Default + BuildHasher + Clone,
{
fn from_iter<T: IntoIterator<Item = (K, V)>>(iter: T) -> Self {
Self {
map: Mutex::new(
iter.into_iter()
.map(|(k, v)| (k, Arc::new(OnceCell::from_value(v))))
.collect(),
),
}
}
}
+212
View File
@@ -0,0 +1,212 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::collections::hash_map::RandomState;
use std::sync::Arc;
use std::sync::atomic::AtomicBool;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::time::Duration;
use crate::once::OnceMap;
#[test]
fn test_default_and_constructors() {
let _map: OnceMap<String, i32> = OnceMap::default();
let _: OnceMap<String, i32> = OnceMap::new();
let _: OnceMap<String, i32> = OnceMap::with_capacity(10);
let _: OnceMap<String, i32> = OnceMap::with_hasher(RandomState::new());
let _: OnceMap<String, i32> = OnceMap::with_capacity_and_hasher(10, RandomState::new());
// Check capacity (indirectly via debug or just ensure it runs)
let map: OnceMap<String, i32> = OnceMap::with_capacity(100);
assert!(format!("{:?}", map).contains("OnceMap"));
}
#[tokio::test]
async fn test_compute() {
let map = OnceMap::new();
let v = map.compute("key", async || 1).await;
assert_eq!(v, 1);
let v = map.compute("key", async || 2).await;
assert_eq!(v, 1);
}
#[tokio::test]
async fn test_compute_concurrent() {
let map = Arc::new(OnceMap::new());
let cnt = Arc::new(AtomicUsize::new(0));
let mut handles = Vec::new();
for _ in 0..10 {
let map = map.clone();
let cnt = cnt.clone();
handles.push(tokio::spawn(async move {
map.compute("key", async move || {
cnt.fetch_add(1, Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(10)).await;
42
})
.await
}));
}
for h in handles {
assert_eq!(h.await.unwrap(), 42);
}
assert_eq!(cnt.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn test_try_compute() {
let map = OnceMap::new();
// Fail first
let res: Result<i32, &str> = map.try_compute("key", async || Err("fail")).await;
assert_eq!(res, Err("fail"));
// Success then
let res: Result<i32, &str> = map.try_compute("key", async || Ok::<i32, &str>(1)).await;
assert_eq!(res, Ok(1));
// Cached
let res: Result<i32, &str> = map.try_compute("key", async || Ok::<i32, &str>(2)).await;
assert_eq!(res, Ok(1));
}
#[tokio::test]
async fn test_try_compute_concurrent_failure_then_success() {
let map = Arc::new(OnceMap::new());
let success = Arc::new(AtomicBool::new(false));
let map_clone = map.clone();
let success_clone = success.clone();
// Spawn a task that fails
let t1 = tokio::spawn(async move {
map_clone
.try_compute("key", async move || {
tokio::time::sleep(Duration::from_millis(50)).await;
Err::<i32, _>("fail")
})
.await
});
// Spawn a task that succeeds, but starts slightly later/runs concurrent
let map_clone2 = map.clone();
let t2 = tokio::spawn(async move {
// Wait for t1 to start
tokio::time::sleep(Duration::from_millis(10)).await;
// This should block until t1 fails, then retry (conceptually)
map_clone2
.try_compute("key", async move || {
success_clone.store(true, Ordering::SeqCst);
Ok::<i32, &str>(1)
})
.await
});
let res1 = t1.await.unwrap();
assert_eq!(res1, Err("fail"));
let res2 = t2.await.unwrap();
assert_eq!(res2, Ok(1));
assert!(success.load(Ordering::SeqCst));
}
#[tokio::test]
async fn test_get_remove() {
let map = OnceMap::new();
assert_eq!(map.get("key"), None);
assert_eq!(map.remove("key"), None);
map.compute("key", async || 1).await;
assert_eq!(map.get("key"), Some(1));
let v = map.remove("key");
assert_eq!(v, Some(1));
assert_eq!(map.get("key"), None);
map.compute("key", async || 2).await;
map.discard("key");
assert_eq!(map.get("key"), None);
}
#[tokio::test]
async fn test_remove_while_computing() {
let map = Arc::new(OnceMap::new());
let map_clone = map.clone();
let t1 = tokio::spawn(async move {
map_clone
.compute("key", async || {
tokio::time::sleep(Duration::from_millis(100)).await;
1
})
.await
});
// Give t1 time to insert the cell and start "computing"
tokio::time::sleep(Duration::from_millis(20)).await;
// Remove should return None because value is not ready
// And it removes the cell from the map.
assert_eq!(map.remove("key"), None);
// t1 finishes. It returns 1.
assert_eq!(t1.await.unwrap(), 1);
// The map should be empty now (key was removed)
assert_eq!(map.get("key"), None);
}
#[tokio::test]
async fn test_get_while_computing() {
let map = Arc::new(OnceMap::new());
let map_clone = map.clone();
let t1 = tokio::spawn(async move {
map_clone
.compute("key", async || {
tokio::time::sleep(Duration::from_millis(50)).await;
1
})
.await
});
tokio::time::sleep(Duration::from_millis(10)).await;
assert_eq!(map.get("key"), None);
assert_eq!(t1.await.unwrap(), 1);
assert_eq!(map.get("key"), Some(1));
}
#[tokio::test]
async fn test_from_iter() {
let map: OnceMap<_, _> = vec![("a", 1), ("b", 2)].into_iter().collect();
assert_eq!(map.get("a"), Some(1));
assert_eq!(map.get("b"), Some(2));
assert_eq!(map.get("c"), None);
}
#[tokio::test]
async fn test_complex_key_value() {
#[derive(Hash, PartialEq, Eq, Clone, Debug)]
struct Key(i32);
let map = OnceMap::new();
let v = map.compute(Key(1), async || "value".to_string()).await;
assert_eq!(v, "value");
assert_eq!(map.get(&Key(1)), Some("value".to_string()));
}
+857
View File
@@ -0,0 +1,857 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
// This implementation is derived from the `oneshot` crate [1], with significant simplifications
// since mea needs not support synchronized receiving functions.
//
// [1] https://github.com/faern/oneshot/blob/83fd0864/src/lib.rs
//! A one-shot channel is used for sending a single message between
//! asynchronous tasks. The [`channel`] function is used to create a
//! [`Sender`] and [`Receiver`] handle pair that form the channel.
//!
//! The `Sender` handle is used by the producer to send the value.
//! The `Receiver` handle is used by the consumer to receive the value.
//!
//! Each handle can be used on separate tasks.
//!
//! Since the `send` method is not async, it can be used anywhere. This includes
//! sending between two runtimes, and using it from non-async code.
//!
//! # Examples
//!
//! ```
//! # #[tokio::main]
//! # async fn main() {
//! use mea::oneshot;
//!
//! let (tx, rx) = oneshot::channel();
//!
//! tokio::spawn(async move {
//! if let Err(_) = tx.send(3) {
//! println!("the receiver dropped");
//! }
//! });
//!
//! match rx.await {
//! Ok(v) => println!("got = {:?}", v),
//! Err(_) => println!("the sender dropped"),
//! }
//! # }
//! ```
//!
//! If the sender is dropped without sending, the receiver will fail with
//! [`RecvError`]:
//!
//! ```
//! # #[tokio::main]
//! # async fn main() {
//! use mea::oneshot;
//!
//! let (tx, rx) = oneshot::channel::<u32>();
//!
//! tokio::spawn(async move {
//! drop(tx);
//! });
//!
//! match rx.await {
//! Ok(_) => panic!("This doesn't happen"),
//! Err(_) => println!("the sender dropped"),
//! }
//! # }
//! ```
use std::any::type_name;
use std::cell::UnsafeCell;
use std::fmt;
use std::future::Future;
use std::future::IntoFuture;
use std::hint;
use std::mem;
use std::mem::MaybeUninit;
use std::pin::Pin;
use std::ptr;
use std::ptr::NonNull;
use std::sync::atomic::AtomicU8;
use std::sync::atomic::Ordering;
use std::sync::atomic::fence;
use std::task::Context;
use std::task::Poll;
use std::task::Waker;
#[cfg(test)]
mod tests;
/// Creates a new oneshot channel and returns the two endpoints, [`Sender`] and [`Receiver`].
pub fn channel<T>() -> (Sender<T>, Receiver<T>) {
let channel_ptr = NonNull::from(Box::leak(Box::new(Channel::new())));
(Sender { channel_ptr }, Receiver { channel_ptr })
}
/// Sends a value to the associated [`Receiver`].
pub struct Sender<T> {
channel_ptr: NonNull<Channel<T>>,
}
impl<T> fmt::Debug for Sender<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Sender").finish_non_exhaustive()
}
}
unsafe impl<T: Send> Send for Sender<T> {}
unsafe impl<T: Sync> Sync for Sender<T> {}
#[inline(always)]
fn sender_wake_up_receiver<T>(channel: &Channel<T>, state: u8) {
// Take the waker, but critically do not awake it. If we awake it now, the
// receiving thread could still observe the AWAKING state and re-await, meaning
// that after we change to the MESSAGE state, it would remain waiting indefinitely
// or until a spurious wakeup.
//
// SAFETY: at this point we are in the AWAKING state, and the receiving thread
// does not access the waker while in this state, nor does it free the channel
// allocation in this state. The caller's acquire ordering establishes a happens-before
// relationship with the writing of the waker.
let waker = unsafe { channel.take_waker() };
// ORDERING: this ordering serves two-fold: it synchronizes with the receiver's
// acquire fence after it observes this state, ensuring that both our read of the
// waker and write of the message happen-before the taking of the message and
// freeing of the channel. Furthermore, we need acquire ordering to ensure awaking
// the receiver happens after the channel state is updated.
channel.state.swap(state, Ordering::AcqRel);
// Note: it is possible that between the store above and this statement that
// the receiving thread is spuriously awakened, takes the message, and frees
// the channel allocation. However, we took ownership of the channel out of
// that allocation, and freeing the channel does not drop the waker since the
// waker is wrapped in MaybeUninit. Therefore, this data is valid regardless of
// whether the receiver has completed by this point.
waker.wake();
}
impl<T> Sender<T> {
/// Attempts to send a value on this channel, returning an error contains the message if it
/// could not be sent.
pub fn send(self, message: T) -> Result<(), SendError<T>> {
let channel_ptr = self.channel_ptr;
// Do not run the Drop implementation if send was called, any cleanup happens below.
mem::forget(self);
// SAFETY: The channel exists on the heap for the entire duration of this method, and we
// only ever acquire shared references to it. Note that if the receiver disconnects it
// does not free the channel.
let channel = unsafe { channel_ptr.as_ref() };
// Write the message into the channel on the heap.
//
// SAFETY: The receiver only ever accesses this memory location if we are in the MESSAGE
// state, and since we are responsible for setting that state, we can guarantee that we have
// exclusive access to this memory location to perform this write.
unsafe { channel.write_message(message) };
// Update the state to signal there is a message on the channel:
//
// * EMPTY + 1 = MESSAGE
// * RECEIVING + 1 = AWAKING
// * DISCONNECTED + 1 = EMPTY (invalid), however this state is never observed
//
// ORDERING: we need release ordering to allow the receiver to synchronize with our write
// of the message and with our final write of the state, in the case where the receiver
// becomes responsible for freeing the channel. We need acquire ordering in the RECEIVING
// and DISCONNECTED branches, as explained further down.
match channel.state.fetch_add(1, Ordering::AcqRel) {
// The receiver is alive and has not started waiting. Send done.
EMPTY => Ok(()),
// The receiver is waiting. Wake it up so it can return the message.
RECEIVING => {
sender_wake_up_receiver(channel, MESSAGE);
Ok(())
}
// The receiver was already dropped. The error is responsible for freeing the channel.
//
// SAFETY: The acquire ordering above synchronizes with the receiver's write of the
// DISCONNECTED state. Since the receiver disconnected it will no longer access
// `channel_ptr`, so we can transfer exclusive ownership of the channel's resources to
// the error.
// Moreover, since we just placed the message in the channel, the channel contains a
// valid message.
DISCONNECTED => Err(SendError { channel_ptr }),
state => unreachable!("unexpected channel state: {}", state),
}
}
/// Returns true if the associated [`Receiver`] has been dropped.
///
/// If true is returned, a future call to send is guaranteed to return an error.
pub fn is_closed(&self) -> bool {
// SAFETY: The channel exists on the heap for the entire duration of this method, and we
// only ever acquire shared references to it. Note that if the receiver disconnects it
// does not free the channel.
let channel = unsafe { self.channel_ptr.as_ref() };
// ORDERING: We *chose* a Relaxed ordering here as it sufficient to enforce the method's
// contract: "if true is returned, a future call to send is guaranteed to return an error."
//
// Once true has been observed, it will remain true. However, if false is observed,
// the receiver might have just disconnected but this thread has not observed it yet.
matches!(channel.state.load(Ordering::Relaxed), DISCONNECTED)
}
}
impl<T> Drop for Sender<T> {
fn drop(&mut self) {
// SAFETY: The receiver only ever frees the channel if we are in the MESSAGE or
// DISCONNECTED states.
//
// * If we are in the MESSAGE state, then we called mem::forget(self), so we should
// not be in this function call.
// * If we are in the DISCONNECTED state, then the receiver either received a MESSAGE
// so this statement is unreachable, or was dropped and observed that our side was still
// alive, and thus didn't free the channel.
let channel = unsafe { self.channel_ptr.as_ref() };
// Update the channel state to disconnected:
//
// * EMPTY ^ 001 = DISCONNECTED
// * RECEIVING ^ 001 = AWAKING
// * DISCONNECTED ^ 001 = EMPTY (invalid), but this state is never observed
//
// ORDERING: Release is required so that in the states where the receiver becomes
// responsible for deallocating the channel, they can synchronize with this final state
// write from us. Acquire is required by the branches below to synchronize with writes from
// the receiver.
match channel.state.fetch_xor(0b001, Ordering::AcqRel) {
// The receiver is not waiting, nor is it dropped. The receiver is responsible for
// deallocating the channel.
EMPTY => {}
// The receiver is waiting. Wake it up so it can detect that the channel disconnected.
RECEIVING => sender_wake_up_receiver(channel, DISCONNECTED),
// The receiver was already dropped. We are responsible for freeing the channel.
DISCONNECTED => {
// SAFETY: when the receiver switches the state to DISCONNECTED they have received
// the message or will no longer be trying to receive the message, and have
// observed that the sender is still alive, meaning that we are responsible for
// freeing the channel allocation. The acquire ordering above synchronizes with
// the receiver's final write of the state.
unsafe { dealloc(self.channel_ptr) };
}
state => unreachable!("unexpected channel state: {}", state),
}
}
}
/// Receives a value from the associated [`Sender`].
pub struct Receiver<T> {
channel_ptr: NonNull<Channel<T>>,
}
impl<T> fmt::Debug for Receiver<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Receiver").finish_non_exhaustive()
}
}
unsafe impl<T: Send> Send for Receiver<T> {}
// The Receiver can NOT be `Sync`! The current receive implementations that take `&self`
// assume no other receive operation runs in parallel.
impl<T> Unpin for Receiver<T> {}
impl<T> IntoFuture for Receiver<T> {
type Output = Result<T, RecvError>;
type IntoFuture = Recv<T>;
fn into_future(self) -> Self::IntoFuture {
let Receiver { channel_ptr } = self;
// Do not run our Drop implementation, since the receiver lives on as the new future.
mem::forget(self);
Recv { channel_ptr }
}
}
impl<T> Receiver<T> {
/// Returns true if the associated [`Sender`] was dropped before sending a message. Or if
/// the message has already been received.
///
/// If `true` is returned, all future calls to receive the message are guaranteed to return
/// [`RecvError`]. And future calls to this method is guaranteed to also return `true`.
pub fn is_closed(&self) -> bool {
// SAFETY: the existence of the `self` parameter serves as a certificate that the receiver
// is still alive, meaning that even if the sender was dropped then it would have observed
// the fact that we are still alive and left the responsibility of deallocating the
// channel to us, so `self.channel` is valid
let channel = unsafe { self.channel_ptr.as_ref() };
// ORDERING: We *chose* a Relaxed ordering here as it is sufficient to
// enforce the method's contract.
//
// Once true has been observed, it will remain true. However, if false is observed,
// the sender might have just disconnected but this thread has not observed it yet.
matches!(channel.state.load(Ordering::Relaxed), DISCONNECTED)
}
/// Returns true if there is a message in the channel, ready to be received.
///
/// If `true` is returned, the next call to receive the message is guaranteed to return
/// the message immediately.
pub fn has_message(&self) -> bool {
// SAFETY: the existence of the `self` parameter serves as a certificate that the receiver
// is still alive, meaning that even if the sender was dropped then it would have observed
// the fact that we are still alive and left the responsibility of deallocating the
// channel to us, so `self.channel` is valid
let channel = unsafe { self.channel_ptr.as_ref() };
// ORDERING: An acquire ordering is used to guarantee no subsequent loads is reordered
// before this one. This upholds the contract that if true is returned, the next call to
// receive the message is guaranteed to also observe the `MESSAGE` state and return the
// message immediately.
matches!(channel.state.load(Ordering::Acquire), MESSAGE)
}
/// Checks if there is a message in the channel without blocking. Returns:
///
/// * `Ok(message)` if there was a message in the channel.
/// * `Err(TryRecvError::Empty)` if the [`Sender`] is alive, but has not yet sent a message.
/// * `Err(TryRecvError::Disconnected)` if the [`Sender`] was dropped before sending anything or
/// if the message has already been extracted by a previous `try_recv` call.
///
/// If a message is returned, the channel is disconnected and any subsequent receive operation
/// using this receiver will return an error: [`TryRecvError::Disconnected`] for `try_recv`,
/// or [`RecvError::Disconnected`] for [`recv`](Receiver::into_future).
pub fn try_recv(&self) -> Result<T, TryRecvError> {
// SAFETY: The channel will not be freed while this method is still running.
let channel = unsafe { self.channel_ptr.as_ref() };
// ORDERING: Relaxed is fine since the only branch that needs synchronization is MESSAGE,
// and that branch has its own synchronization.
match channel.state.load(Ordering::Relaxed) {
MESSAGE => {
// It is okay to break up the load and store since once we are in the MESSAGE state,
// the sender no longer modifies the state
//
// ORDERING: at this point the sender has done its job and is no longer active, so
// we need not make any side effects visible to it.
channel.state.store(DISCONNECTED, Ordering::Relaxed);
// ORDERING: Synchronize with the sender's write of the message.
fence(Ordering::Acquire);
// SAFETY: we are in the MESSAGE state so the message is present and synchronized.
Ok(unsafe { channel.take_message() })
}
EMPTY => Err(TryRecvError::Empty),
DISCONNECTED => Err(TryRecvError::Disconnected),
state => unreachable!("unexpected channel state: {}", state),
}
}
}
impl<T> Drop for Receiver<T> {
fn drop(&mut self) {
// SAFETY: since the receiving side is still alive the sender would have observed that and
// left deallocating the channel allocation to us.
let channel = unsafe { self.channel_ptr.as_ref() };
// Set the channel state to disconnected and read what state the channel was in.
//
// ORDERING: Release is required so that in the states where the sender becomes responsible
// for deallocating the channel, they can synchronize with this final state write from us.
// Acquire is required by the branches below to synchronize with writes from the sender.
match channel.state.swap(DISCONNECTED, Ordering::AcqRel) {
// The sender has not sent anything, nor is it dropped. The sender is responsible for
// deallocating the channel.
EMPTY => {}
// The sender already sent something. We must drop it, and free the channel.
MESSAGE => {
// SAFETY: The MESSAGE state plus acquire ordering guarantees the sender has
// written a message and that it has a happens-before relationship with this drop.
unsafe { channel.drop_message() };
// SAFETY: The acquire ordering above synchronizes with the sender's final write
// of the state, so we can safely deallocate the channel.
unsafe { dealloc(self.channel_ptr) };
}
// The sender was already dropped. We are responsible for freeing the channel.
DISCONNECTED => {
// SAFETY: The acquire ordering above synchronizes with the sender's final write
// of the state, so we can safely deallocate the channel.
unsafe { dealloc(self.channel_ptr) };
}
// NOTE: the receiver, unless transformed into a future, will never see the
// RECEIVING or AWAKING states, so we can ignore them here.
state => unreachable!("unexpected channel state: {}", state),
}
}
}
/// A future that completes when the message is sent from the associated [`Sender`], or the
/// [`Sender`] is dropped before sending a message.
pub struct Recv<T> {
channel_ptr: NonNull<Channel<T>>,
}
impl<T> fmt::Debug for Recv<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Recv").finish_non_exhaustive()
}
}
unsafe impl<T: Send> Send for Recv<T> {}
fn recv_awaken<T>(channel: &Channel<T>) -> Poll<Result<T, RecvError>> {
loop {
hint::spin_loop();
// ORDERING: The MESSAGE branch below uses a dedicated fence to synchronize with the
// sender. Until then, we only need to observe the state change.
match channel.state.load(Ordering::Relaxed) {
AWAKING => {}
DISCONNECTED => break Poll::Ready(Err(RecvError::Disconnected)),
MESSAGE => {
// ORDERING: after publishing MESSAGE, the sender no longer uses the channel, so
// this state update only needs to be visible to this receiver.
channel.state.store(DISCONNECTED, Ordering::Relaxed);
// ORDERING: Synchronize with the sender's write of the message and final state.
fence(Ordering::Acquire);
// SAFETY: We observed the MESSAGE state and synchronized with the sender.
break Poll::Ready(Ok(unsafe { channel.take_message() }));
}
state => unreachable!("unexpected channel state: {}", state),
}
}
}
impl<T> Future for Recv<T> {
type Output = Result<T, RecvError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
// SAFETY: the existence of the `self` parameter serves as a certificate that the receiver
// is still alive, meaning that even if the sender was dropped then it would have observed
// the fact that we are still alive and left the responsibility of deallocating the
// channel to us, so `self.channel` is valid
let channel = unsafe { self.channel_ptr.as_ref() };
// ORDERING: Relaxed is fine since the branches that need synchronization use dedicated
// fences.
match channel.state.load(Ordering::Relaxed) {
// The sender is alive but has not sent anything yet.
EMPTY => {
let waker = cx.waker().clone();
// SAFETY: We can not be in the forbidden states, and no waker in the channel.
unsafe { channel.write_waker(waker) }
}
// The sender sent the message.
MESSAGE => {
// ORDERING: after publishing MESSAGE, the sender no longer uses the channel, so
// this state update only needs to be visible to this receiver.
channel.state.store(DISCONNECTED, Ordering::Relaxed);
// ORDERING: Synchronize with the sender's write of the message and final state.
fence(Ordering::Acquire);
// SAFETY: we are in the MESSAGE state and have synchronized with the sender.
Poll::Ready(Ok(unsafe { channel.take_message() }))
}
// We were polled again while waiting for the sender. Replace the waker with the new
// one.
RECEIVING => {
// ORDERING: Success synchronizes with the previous write_waker call before we
// drop the stored waker. Failure does not access the stored waker.
match channel.state.compare_exchange(
RECEIVING,
EMPTY,
Ordering::Acquire,
Ordering::Relaxed,
) {
// The state is EMPTY again.
Ok(_) => {
let waker = cx.waker().clone();
// SAFETY: The successful exchange makes the state EMPTY, so the sender
// cannot take the stored waker. The acquire ordering synchronizes with the
// waker write.
unsafe { channel.drop_waker() };
// SAFETY: We can not be in the forbidden states, and no waker in the
// channel.
unsafe { channel.write_waker(waker) }
}
// The sender sent the message while we prepared to replace the waker.
// We take the message and mark the channel disconnected.
// The sender has already taken the waker.
Err(MESSAGE) => {
// ORDERING: after publishing MESSAGE, the sender no longer uses the
// channel, so this state update only needs to be visible to this receiver.
channel.state.store(DISCONNECTED, Ordering::Relaxed);
// ORDERING: Synchronize with the sender's write of the message.
fence(Ordering::Acquire);
// SAFETY: The state tells us the sender has initialized the message, and
// the fence above synchronizes with that write.
Poll::Ready(Ok(unsafe { channel.take_message() }))
}
// The sender is currently waking us up.
Err(AWAKING) => recv_awaken(channel),
// The sender was dropped before sending anything while we prepared to park.
// The sender has taken the waker already.
Err(DISCONNECTED) => Poll::Ready(Err(RecvError::Disconnected)),
Err(state) => unreachable!("unexpected channel state: {}", state),
}
}
// The sender has observed the RECEIVING state and is currently reading the waker from
// a previous poll. We need to loop here until we observe the MESSAGE or DISCONNECTED
// state. We busy loop here since we know the sender is done very soon.
AWAKING => recv_awaken(channel),
// The sender was dropped before sending anything.
DISCONNECTED => Poll::Ready(Err(RecvError::Disconnected)),
state => unreachable!("unexpected channel state: {}", state),
}
}
}
impl<T> Drop for Recv<T> {
fn drop(&mut self) {
// SAFETY: since the receiving side is still alive the sender would have observed that and
// left deallocating the channel allocation to us.
let channel = unsafe { self.channel_ptr.as_ref() };
loop {
// ORDERING: MESSAGE and DISCONNECTED synchronize with the sender's state writes.
match channel.state.load(Ordering::Acquire) {
// The sender has not sent anything, nor is it dropped. Mark the receiver as
// dropped; the sender is responsible for deallocating the channel.
EMPTY => {
if channel
.state
.compare_exchange(EMPTY, DISCONNECTED, Ordering::Release, Ordering::Relaxed)
.is_ok()
{
break;
}
}
// The sender already sent something. We must drop it, and free the channel.
MESSAGE => {
// SAFETY: The MESSAGE state plus acquire ordering guarantees the sender has
// written a message and that it has a happens-before relationship with this
// drop.
unsafe { channel.drop_message() };
// SAFETY: The acquire load above synchronizes with the sender's final write of
// the state, so we can safely deallocate the channel.
unsafe { dealloc(self.channel_ptr) };
break;
}
// This receiver was previously polled, but was not polled to completion. Move away
// from RECEIVING before dropping the waker so the sender cannot take the same
// waker.
//
// A successful exchange creates a short EMPTY window before the next iteration can
// mark DISCONNECTED. This branch owns and drops the stored waker first. A sender
// that observes EMPTY does not touch the waker. It either stores MESSAGE and
// leaves the message and allocation to this loop, or stores DISCONNECTED and
// leaves the allocation to this loop. If this loop marks DISCONNECTED first, the
// sender observes DISCONNECTED and owns any send error cleanup.
RECEIVING => {
if channel
.state
.compare_exchange(RECEIVING, EMPTY, Ordering::Acquire, Ordering::Relaxed)
.is_ok()
{
// SAFETY: The successful exchange makes the state EMPTY, so the sender
// cannot take the stored waker. The acquire ordering synchronizes with the
// waker write.
unsafe { channel.drop_waker() };
}
}
// The sender has observed RECEIVING and is taking the waker. Wait until it stores
// MESSAGE or DISCONNECTED.
AWAKING => {
hint::spin_loop();
}
// The sender was already dropped, or this future was previously polled to
// completion. We are responsible for freeing the channel.
DISCONNECTED => {
// SAFETY: When DISCONNECTED comes from the sender, the acquire load
// synchronizes with the sender's state write. When it comes from our own
// completed poll, the message has already been taken.
unsafe { dealloc(self.channel_ptr) };
break;
}
state => unreachable!("unexpected channel state: {}", state),
}
}
}
}
/// Internal channel data structure.
///
/// The [`channel`] method allocates and puts one instance of this struct on the heap for each
/// oneshot channel instance. The struct holds:
///
/// * The current state of the channel.
/// * The message in the channel. This memory is uninitialized until the message is sent.
/// * The waker instance for the task that is currently receiving on this channel. This memory is
/// uninitialized until the receiver starts receiving.
struct Channel<T> {
state: AtomicU8,
message: UnsafeCell<MaybeUninit<T>>,
waker: UnsafeCell<MaybeUninit<Waker>>,
}
impl<T> Channel<T> {
const fn new() -> Self {
Self {
state: AtomicU8::new(EMPTY),
message: UnsafeCell::new(MaybeUninit::uninit()),
waker: UnsafeCell::new(MaybeUninit::uninit()),
}
}
#[inline(always)]
unsafe fn message(&self) -> &T {
// SAFETY: The caller guarantees that no other thread will access the message field.
let message_container = unsafe { &*self.message.get() };
// SAFETY: The caller guarantees that the message has been initialized.
unsafe { message_container.assume_init_ref() }
}
#[inline(always)]
unsafe fn write_message(&self, message: T) {
unsafe {
let slot = &mut *self.message.get();
slot.as_mut_ptr().write(message);
}
}
#[inline(always)]
unsafe fn drop_message(&self) {
unsafe {
let slot = &mut *self.message.get();
slot.assume_init_drop();
}
}
#[inline(always)]
unsafe fn take_message(&self) -> T {
unsafe { ptr::read(self.message.get()).assume_init() }
}
/// # Safety
///
/// * The `waker` field must not have a waker stored when calling this method.
/// * The `state` must not be in the RECEIVING or AWAKING state when calling this method.
unsafe fn write_waker(&self, waker: Waker) -> Poll<Result<T, RecvError>> {
// Write the waker instance to the channel.
//
// SAFETY: we are not yet in the RECEIVING state, meaning that the sender will not
// try to access the waker until it sees the state set to RECEIVING below.
unsafe {
let slot = &mut *self.waker.get();
slot.as_mut_ptr().write(waker);
}
// ORDERING: we use release ordering on success so the sender can synchronize with
// our write of the waker. We use relaxed ordering on failure since the sender does
// not need to synchronize with our write and the individual match arms handle any
// additional synchronization
match self
.state
.compare_exchange(EMPTY, RECEIVING, Ordering::Release, Ordering::Relaxed)
{
// We stored our waker, now we return and let the sender wake us up.
Ok(_) => Poll::Pending,
// The sender sent the message while we prepared to await.
// We take the message and mark the channel disconnected.
Err(MESSAGE) => {
// SAFETY: We wrote a waker above. The sender cannot have observed the RECEIVING
// state, so it has not accessed the waker. We must drop it.
unsafe { self.drop_waker() };
// ORDERING: sender does not exist, so this update only needs to be visible to
// us.
self.state.store(DISCONNECTED, Ordering::Relaxed);
// ORDERING: Synchronize with writing message. This branch is unlikely to be
// taken, so it is likely more efficient to use a fence here instead of AcqRel
// ordering on the compare_exchange operation.
fence(Ordering::Acquire);
// SAFETY: The MESSAGE state tells us there is a correctly initialized message,
// and the fence above synchronizes with that write.
Poll::Ready(Ok(unsafe { self.take_message() }))
}
// The sender was dropped before sending anything while we prepared to await.
Err(DISCONNECTED) => {
// SAFETY: We wrote a waker above. The sender cannot have observed the RECEIVING
// state, so it has not accessed the waker. We must drop it.
unsafe { self.drop_waker() };
Poll::Ready(Err(RecvError::Disconnected))
}
Err(state) => unreachable!("unexpected channel state: {}", state),
}
}
#[inline(always)]
unsafe fn drop_waker(&self) {
unsafe {
let slot = &mut *self.waker.get();
slot.assume_init_drop();
}
}
#[inline(always)]
unsafe fn take_waker(&self) -> Waker {
unsafe { ptr::read(self.waker.get()).assume_init() }
}
}
unsafe fn dealloc<T>(channel: NonNull<Channel<T>>) {
unsafe { drop(Box::from_raw(channel.as_ptr())) }
}
/// An error returned when trying to send on a closed channel. Returned from
/// [`Sender::send`] if the corresponding [`Receiver`] has already been dropped.
///
/// The message that could not be sent can be retrieved again with [`SendError::into_inner`].
pub struct SendError<T> {
channel_ptr: NonNull<Channel<T>>,
}
// SAFETY: The SendError only contains a pointer to the channel. The constructor (if used
// correctly) guarantees exclusive ownership and access to the underlying channel. Since
// the message is Send (`T: Send`) it is safe to extract it or drop it via the SendError
// on any thread.
unsafe impl<T: Send> Send for SendError<T> {}
// SAFETY: Same basic safety as described in the Send impl above. Plus the fact that `T`
// is `Sync` allows the SendError to be shared between threads and hand out `&T` references
// as well.
unsafe impl<T: Sync> Sync for SendError<T> {}
impl<T> SendError<T> {
/// Get a reference to the message that failed to be sent.
pub fn as_inner(&self) -> &T {
// SAFETY: we have exclusive ownership of the channel and require that the message has
// been initialized upon construction.
unsafe { self.channel_ptr.as_ref().message() }
}
/// Consumes the error and returns the message that failed to be sent.
pub fn into_inner(self) -> T {
let channel_ptr = self.channel_ptr;
// Do not run destructor if we consumed ourselves. Freeing happens below.
mem::forget(self);
// SAFETY: we have ownership of the channel
let channel: &Channel<T> = unsafe { channel_ptr.as_ref() };
// SAFETY: we know that the message is initialized according to the safety requirements of
// `new`
let message = unsafe { channel.take_message() };
// SAFETY: we own the channel
unsafe { dealloc(channel_ptr) };
message
}
}
impl<T> Drop for SendError<T> {
fn drop(&mut self) {
// SAFETY: there is a properly initialized message
unsafe { self.channel_ptr.as_ref().drop_message() };
// SAFETY: we own the channel
unsafe { dealloc(self.channel_ptr) };
}
}
impl<T> fmt::Display for SendError<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("sending on a closed channel")
}
}
impl<T> fmt::Debug for SendError<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "SendError<{}>(..)", type_name::<T>())
}
}
impl<T> std::error::Error for SendError<T> {}
/// Error returned by [`Receiver::try_recv`].
#[derive(Debug, Clone, Eq, PartialEq)]
pub enum TryRecvError {
/// This channel is currently empty, but the sender has not yet disconnected, so data may yet
/// become available.
Empty,
/// The sender has become disconnected, and there will never be any more data received on it.
Disconnected,
}
impl fmt::Display for TryRecvError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(match self {
TryRecvError::Empty => "receiving on an empty channel",
TryRecvError::Disconnected => "receiving on a closed channel",
})
}
}
impl std::error::Error for TryRecvError {}
/// An error returned when awaiting the message via [`Receiver`].
///
/// This error indicates that the corresponding [`Sender`] was dropped before sending any message.
/// Note that if a message was already received (e.g., via [`Receiver::try_recv`]), subsequent
/// `try_recv` calls will return [`TryRecvError::Disconnected`] instead.
#[derive(Debug, Clone, Eq, PartialEq)]
pub enum RecvError {
/// The sender has become disconnected, and there will never be any more data received on it.
Disconnected,
}
impl fmt::Display for RecvError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("receiving on a closed channel")
}
}
impl std::error::Error for RecvError {}
/// The initial channel state. Active while both endpoints are still alive, no message has been
/// sent, and the receiver is not receiving.
const EMPTY: u8 = 0b011;
/// A message has been sent to the channel, but the receiver has not yet read it.
const MESSAGE: u8 = 0b100;
/// No message has yet been sent on the channel, but the receiver future ([`Recv`]) is currently
/// receiving.
const RECEIVING: u8 = 0b000;
/// A message is sending to the channel, or the channel is closing. The receiver future ([`Recv`])
/// is currently being awakened.
const AWAKING: u8 = 0b001;
/// The channel has been closed. This means that either the sender or receiver has been dropped,
/// or the message sent to the channel has already been received.
const DISCONNECTED: u8 = 0b010;
+508
View File
@@ -0,0 +1,508 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::future::Future;
use std::future::IntoFuture;
use std::hint::spin_loop;
use std::mem;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::AtomicU32;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::task::Context;
use std::task::Poll;
use std::task::RawWaker;
use std::task::RawWakerVTable;
use std::task::Waker;
use std::time::Duration;
use std::time::Instant;
use crate::oneshot;
use crate::oneshot::TryRecvError;
struct DropCounterHandle(Arc<AtomicUsize>);
impl DropCounterHandle {
pub fn count(&self) -> usize {
self.0.load(Ordering::SeqCst)
}
}
struct DropCounter<T> {
drop_count: Arc<AtomicUsize>,
value: Option<T>,
}
impl<T> DropCounter<T> {
fn new(value: T) -> (Self, DropCounterHandle) {
let drop_count = Arc::new(AtomicUsize::new(0));
(
Self {
drop_count: drop_count.clone(),
value: Some(value),
},
DropCounterHandle(drop_count),
)
}
fn value(&self) -> &T {
self.value.as_ref().unwrap()
}
}
impl<T> Drop for DropCounter<T> {
fn drop(&mut self) {
self.drop_count.fetch_add(1, Ordering::SeqCst);
}
}
#[tokio::test]
async fn send_before_await() {
let (sender, receiver) = oneshot::channel();
assert!(sender.send(19i128).is_ok());
assert_eq!(receiver.await, Ok(19i128));
}
#[tokio::test]
async fn await_with_dropped_sender() {
let (sender, receiver) = oneshot::channel::<u128>();
drop(sender);
receiver.await.unwrap_err();
}
#[tokio::test]
async fn await_before_send() {
let (sender, receiver) = oneshot::channel();
let (message, counter) = DropCounter::new(79u128);
let t = tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(10)).await;
sender.send(message)
});
let returned_message = receiver.await.unwrap();
assert_eq!(counter.count(), 0);
assert_eq!(*returned_message.value(), 79u128);
drop(returned_message);
assert_eq!(counter.count(), 1);
t.await.unwrap().unwrap();
}
#[tokio::test]
async fn await_before_send_then_drop_sender() {
let (sender, receiver) = oneshot::channel::<u128>();
let t = tokio::spawn(async {
tokio::time::sleep(Duration::from_millis(10)).await;
drop(sender);
});
assert!(receiver.await.is_err());
t.await.unwrap();
}
#[tokio::test]
async fn poll_receiver_then_drop_it() {
let (sender, receiver) = oneshot::channel::<()>();
// This will poll the receiver and then give up after 100 ms.
tokio::time::timeout(Duration::from_millis(100), receiver)
.await
.unwrap_err();
// Make sure the receiver has been dropped by the runtime.
assert!(sender.send(()).is_err());
}
#[tokio::test]
async fn recv_within_select() {
let (tx, rx) = oneshot::channel::<&'static str>();
let mut interval = tokio::time::interval(Duration::from_millis(10));
let handle = tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(100)).await;
tx.send("shut down").unwrap();
});
let mut recv = rx.into_future();
loop {
tokio::select! {
_ = interval.tick() => println!("another 10ms"),
msg = &mut recv => {
println!("Got message: {}", msg.unwrap());
break;
}
}
}
handle.await.unwrap();
}
#[test]
fn try_recv_success_then_disconnected() {
let (tx, rx) = oneshot::channel::<i32>();
tx.send(10).unwrap();
assert_eq!(rx.try_recv(), Ok(10));
assert_eq!(rx.try_recv(), Err(TryRecvError::Disconnected));
}
#[test]
fn try_recv_empty_with_live_sender() {
let (_tx, rx) = oneshot::channel::<()>();
assert_eq!(rx.try_recv(), Err(TryRecvError::Empty));
}
#[test]
fn try_recv_disconnected_after_drop() {
let (tx, rx) = oneshot::channel::<()>();
drop(tx);
assert_eq!(rx.try_recv(), Err(TryRecvError::Disconnected));
}
#[derive(Default)]
pub struct WakerHandle {
clone_count: AtomicU32,
drop_count: AtomicU32,
wake_count: AtomicU32,
}
impl WakerHandle {
pub fn clone_count(&self) -> u32 {
self.clone_count.load(Ordering::Relaxed)
}
pub fn drop_count(&self) -> u32 {
self.drop_count.load(Ordering::Relaxed)
}
pub fn wake_count(&self) -> u32 {
self.wake_count.load(Ordering::Relaxed)
}
}
fn waker() -> (Waker, Arc<WakerHandle>) {
let waker_handle = Arc::new(WakerHandle::default());
let waker_handle_ptr = Arc::into_raw(waker_handle.clone());
let raw_waker = RawWaker::new(waker_handle_ptr as *const _, waker_vtable());
(unsafe { Waker::from_raw(raw_waker) }, waker_handle)
}
fn waker_vtable() -> &'static RawWakerVTable {
&RawWakerVTable::new(clone_raw, wake_raw, wake_by_ref_raw, drop_raw)
}
unsafe fn clone_raw(data: *const ()) -> RawWaker {
let handle: Arc<WakerHandle> = unsafe { Arc::from_raw(data as *const _) };
handle.clone_count.fetch_add(1, Ordering::Relaxed);
mem::forget(handle.clone());
mem::forget(handle);
RawWaker::new(data, waker_vtable())
}
unsafe fn wake_raw(data: *const ()) {
let handle: Arc<WakerHandle> = unsafe { Arc::from_raw(data as *const _) };
handle.wake_count.fetch_add(1, Ordering::Relaxed);
handle.drop_count.fetch_add(1, Ordering::Relaxed);
}
unsafe fn wake_by_ref_raw(data: *const ()) {
let handle: Arc<WakerHandle> = unsafe { Arc::from_raw(data as *const _) };
handle.wake_count.fetch_add(1, Ordering::Relaxed);
mem::forget(handle)
}
unsafe fn drop_raw(data: *const ()) {
let handle: Arc<WakerHandle> = unsafe { Arc::from_raw(data as *const _) };
handle.drop_count.fetch_add(1, Ordering::Relaxed);
drop(handle)
}
#[test]
fn poll_then_send() {
let (sender, receiver) = oneshot::channel::<u128>();
let mut receiver = receiver.into_future();
let (waker, waker_handle) = waker();
let mut context = Context::from_waker(&waker);
assert_eq!(Pin::new(&mut receiver).poll(&mut context), Poll::Pending);
assert_eq!(waker_handle.clone_count(), 1);
assert_eq!(waker_handle.drop_count(), 0);
assert_eq!(waker_handle.wake_count(), 0);
sender.send(1234).unwrap();
assert_eq!(waker_handle.clone_count(), 1);
assert_eq!(waker_handle.drop_count(), 1);
assert_eq!(waker_handle.wake_count(), 1);
assert_eq!(
Pin::new(&mut receiver).poll(&mut context),
Poll::Ready(Ok(1234))
);
assert_eq!(waker_handle.clone_count(), 1);
assert_eq!(waker_handle.drop_count(), 1);
assert_eq!(waker_handle.wake_count(), 1);
}
#[test]
fn poll_with_different_wakers() {
let (sender, receiver) = oneshot::channel::<u128>();
let mut receiver = receiver.into_future();
let (waker1, waker_handle1) = waker();
let mut context1 = Context::from_waker(&waker1);
assert_eq!(Pin::new(&mut receiver).poll(&mut context1), Poll::Pending);
assert_eq!(waker_handle1.clone_count(), 1);
assert_eq!(waker_handle1.drop_count(), 0);
assert_eq!(waker_handle1.wake_count(), 0);
let (waker2, waker_handle2) = waker();
let mut context2 = Context::from_waker(&waker2);
assert_eq!(Pin::new(&mut receiver).poll(&mut context2), Poll::Pending);
assert_eq!(waker_handle1.clone_count(), 1);
assert_eq!(waker_handle1.drop_count(), 1);
assert_eq!(waker_handle1.wake_count(), 0);
assert_eq!(waker_handle2.clone_count(), 1);
assert_eq!(waker_handle2.drop_count(), 0);
assert_eq!(waker_handle2.wake_count(), 0);
// Sending should cause the waker from the latest poll to be woken up
sender.send(1234).unwrap();
assert_eq!(waker_handle1.clone_count(), 1);
assert_eq!(waker_handle1.drop_count(), 1);
assert_eq!(waker_handle1.wake_count(), 0);
assert_eq!(waker_handle2.clone_count(), 1);
assert_eq!(waker_handle2.drop_count(), 1);
assert_eq!(waker_handle2.wake_count(), 1);
}
#[test]
fn poll_with_different_wakers_across_threads() {
let (sender, receiver) = oneshot::channel::<u128>();
let mut receiver = receiver.into_future();
let (waker1, waker_handle1) = waker();
let mut context1 = Context::from_waker(&waker1);
assert_eq!(Pin::new(&mut receiver).poll(&mut context1), Poll::Pending);
assert_eq!(waker_handle1.clone_count(), 1);
assert_eq!(waker_handle1.drop_count(), 0);
assert_eq!(waker_handle1.wake_count(), 0);
let receiver_thread = spawn_named("receiver", move || {
let (waker2, waker_handle2) = waker();
let mut context2 = Context::from_waker(&waker2);
assert_eq!(Pin::new(&mut receiver).poll(&mut context2), Poll::Pending);
assert_eq!(waker_handle2.clone_count(), 1);
assert_eq!(waker_handle2.drop_count(), 0);
assert_eq!(waker_handle2.wake_count(), 0);
drop(receiver);
assert_eq!(waker_handle2.drop_count(), 1);
});
receiver_thread.join().unwrap();
assert_eq!(waker_handle1.drop_count(), 1);
assert!(sender.is_closed());
}
#[test]
fn drop_pending_receiver_closes_channel_and_drops_waker() {
let (sender, receiver) = oneshot::channel::<u128>();
let mut receiver = receiver.into_future();
let (waker, waker_handle) = waker();
let mut context = Context::from_waker(&waker);
assert_eq!(Pin::new(&mut receiver).poll(&mut context), Poll::Pending);
assert_eq!(waker_handle.clone_count(), 1);
assert_eq!(waker_handle.drop_count(), 0);
assert_eq!(waker_handle.wake_count(), 0);
drop(receiver);
assert_eq!(waker_handle.drop_count(), 1);
assert_eq!(waker_handle.wake_count(), 0);
assert!(sender.is_closed());
let error = sender.send(1234).unwrap_err();
assert_eq!(*error.as_inner(), 1234);
}
#[test]
fn poll_then_drop_receiver_during_send() {
let (sender, receiver) = oneshot::channel::<u128>();
let mut receiver = receiver.into_future();
let (waker, _waker_handle) = waker();
let mut context = Context::from_waker(&waker);
// Put the channel into the receiving state
assert_eq!(Pin::new(&mut receiver).poll(&mut context), Poll::Pending);
// Spawn a separate thread that sends in parallel
let t = std::thread::spawn(move || {
let _ = sender.send(1234);
});
// Drop the receiver.
drop(receiver);
// The send operation should also not have panicked
t.join().unwrap();
}
#[test]
fn dropping_sender_disconnects_async_receiver() {
let (sender, receiver) = oneshot::channel::<()>();
assert!(!sender.is_closed());
assert!(!receiver.is_closed());
drop(sender);
assert!(receiver.is_closed());
}
#[test]
fn async_receiver_has_message() {
let (sender, receiver) = oneshot::channel();
assert!(!receiver.has_message());
assert!(sender.send(19i128).is_ok());
assert!(receiver.has_message());
}
#[test]
fn concurrent_send_and_try_recv_to_completion() {
let (sender, receiver) = oneshot::channel::<i32>();
let receiver_thread = spawn_named("receiver", move || {
spin_until("message from sender", || match receiver.try_recv() {
Ok(999) => Some(()),
Ok(value) => panic!("unexpected value: {value}"),
Err(TryRecvError::Empty) => None,
Err(TryRecvError::Disconnected) => panic!("unexpected disconnect"),
});
});
let sender_thread = spawn_named("sender", move || {
sender.send(999).unwrap();
});
receiver_thread.join().unwrap();
sender_thread.join().unwrap();
}
#[test]
fn concurrent_drop_sender_and_try_recv_to_completion() {
let (sender, receiver) = oneshot::channel::<i32>();
let receiver_thread = spawn_named("receiver", move || {
spin_until("sender disconnect", || match receiver.try_recv() {
Ok(value) => panic!("unexpected value: {value}"),
Err(TryRecvError::Empty) => None,
Err(TryRecvError::Disconnected) => Some(()),
});
});
let sender_thread = spawn_named("sender", move || {
drop(sender);
});
receiver_thread.join().unwrap();
sender_thread.join().unwrap();
}
#[test]
fn concurrent_send_and_poll_to_completion() {
let (sender, receiver) = oneshot::channel::<i32>();
let receiver_thread = spawn_named("receiver", move || {
let mut receiver = receiver.into_future();
let (waker, _waker_handle) = waker();
let mut context = Context::from_waker(&waker);
spin_until("poll ready with message", || {
match Pin::new(&mut receiver).poll(&mut context) {
Poll::Ready(Ok(999)) => Some(()),
Poll::Ready(result) => panic!("unexpected result: {result:?}"),
Poll::Pending => None,
}
});
});
let sender_thread = spawn_named("sender", move || {
sender.send(999).unwrap();
});
receiver_thread.join().unwrap();
sender_thread.join().unwrap();
}
#[test]
fn concurrent_drop_sender_and_poll_to_completion() {
let (sender, receiver) = oneshot::channel::<i32>();
let receiver_thread = spawn_named("receiver", move || {
let mut receiver = receiver.into_future();
let (waker, _waker_handle) = waker();
let mut context = Context::from_waker(&waker);
spin_until("poll ready with disconnect", || {
match Pin::new(&mut receiver).poll(&mut context) {
Poll::Ready(Err(oneshot::RecvError::Disconnected)) => Some(()),
Poll::Ready(result) => panic!("unexpected result: {result:?}"),
Poll::Pending => None,
}
});
});
let sender_thread = spawn_named("sender", move || {
drop(sender);
});
receiver_thread.join().unwrap();
sender_thread.join().unwrap();
}
fn spawn_named<F>(name: &str, f: F) -> std::thread::JoinHandle<()>
where
F: FnOnce() + Send + 'static,
{
std::thread::Builder::new()
.name(name.to_string())
.spawn(f)
.unwrap()
}
fn spin_until<F>(label: &str, mut f: F)
where
F: FnMut() -> Option<()>,
{
let deadline = Instant::now() + Duration::from_secs(5);
let mut spins = 0usize;
loop {
if f().is_some() {
break;
}
assert!(Instant::now() < deadline, "timed out waiting for {label}");
if spins % 64 == 0 {
std::thread::yield_now();
} else {
spin_loop();
}
spins += 1;
}
}
+265
View File
@@ -0,0 +1,265 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::fmt;
use std::marker::PhantomData;
use std::ops::Deref;
use std::ptr::NonNull;
use crate::internal;
/// RAII structure used to release the shared read access of a lock when dropped, for a mapped
/// component of the locked data.
///
/// This structure is created by the [`map`] and [`filter_map`] methods on [`RwLockReadGuard`]. It
/// allows you to hold a read lock on a subfield of the protected data, enabling more fine-grained
/// access control while maintaining the same locking semantics.
///
/// As long as you have this guard, you have shared read access to the underlying `T`. The guard
/// internally keeps a reference to the original rwlock's semaphore, so the original lock is
/// maintained until this guard is dropped.
///
/// `MappedRwLockReadGuard` implements [`Send`] and [`Sync`] when `T: Sync`, allowing it to be
/// used across task boundaries and shared between threads safely. Note that [`Send`] does not
/// require `T: Send` because the read guard only borrows the data rather than owning it.
///
/// [`map`]: crate::rwlock::RwLockReadGuard::map
/// [`filter_map`]: crate::rwlock::RwLockReadGuard::filter_map
/// [`RwLockReadGuard`]: crate::rwlock::RwLockReadGuard
///
/// See the [module level documentation](crate::rwlock) for more.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::rwlock::RwLock;
/// use mea::rwlock::RwLockReadGuard;
///
/// #[derive(Debug)]
/// struct User {
/// id: u32,
/// profile: UserProfile,
/// }
///
/// #[derive(Debug)]
/// struct UserProfile {
/// email: String,
/// name: String,
/// }
///
/// let user = User {
/// id: 1,
/// profile: UserProfile {
/// email: "user@example.com".to_owned(),
/// name: "Alice".to_owned(),
/// },
/// };
///
/// let rwlock = RwLock::new(user);
/// let guard = rwlock.read().await;
/// let profile_guard = RwLockReadGuard::map(guard, |user| &user.profile);
///
/// // Now we can only access the user's profile
/// assert_eq!(profile_guard.email, "user@example.com");
/// # }
/// ```
#[must_use = "if unused the RwLock will immediately unlock"]
pub struct MappedRwLockReadGuard<'a, T: ?Sized> {
d: NonNull<T>,
s: &'a internal::Semaphore,
variance: PhantomData<fn() -> T>,
}
// SAFETY: MappedRwLockReadGuard is Send when T: Sync. We don't require T: Send because
// the guard RwLockReadGuard doesn't transfer ownership of T - it only holds a shared reference.
// When moved to another thread, the guard maintains the read lock and the new thread
// can safely access &T (which is allowed since T: Sync). The semaphore reference
// and NonNull pointer are both safe to transfer between threads.
unsafe impl<T: ?Sized + Sync> Send for MappedRwLockReadGuard<'_, T> {}
// SAFETY: `&MappedRwLockReadGuard` can be shared between threads if `T: Sync`.
// Accessing the guard only provides a `&T`, which is safe to share concurrently when `T: Sync`.
unsafe impl<T: ?Sized + Sync> Sync for MappedRwLockReadGuard<'_, T> {}
impl<'a, T: ?Sized> MappedRwLockReadGuard<'a, T> {
pub(crate) fn new(d: NonNull<T>, s: &'a internal::Semaphore) -> Self {
Self {
d,
s,
variance: PhantomData,
}
}
}
impl<T: ?Sized> Drop for MappedRwLockReadGuard<'_, T> {
fn drop(&mut self) {
self.s.release(1);
}
}
impl<T: ?Sized + fmt::Debug> fmt::Debug for MappedRwLockReadGuard<'_, T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Debug::fmt(&**self, f)
}
}
impl<T: ?Sized + fmt::Display> fmt::Display for MappedRwLockReadGuard<'_, T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Display::fmt(&**self, f)
}
}
impl<T: ?Sized> Deref for MappedRwLockReadGuard<'_, T> {
type Target = T;
fn deref(&self) -> &Self::Target {
// SAFETY: we hold the read lock and the NonNull pointer is valid for the guard's lifetime
unsafe { self.d.as_ref() }
}
}
impl<'a, T: ?Sized> MappedRwLockReadGuard<'a, T> {
/// Makes a new [`MappedRwLockReadGuard`] for a component of the locked data.
///
/// This operation cannot fail as the `MappedRwLockReadGuard` passed in already locked the
/// rwlock.
///
/// This is an associated function that needs to be used as `MappedRwLockReadGuard::map(...)`.
///
/// A method would interfere with methods of the same name on the contents of the locked data.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::rwlock::MappedRwLockReadGuard;
/// use mea::rwlock::RwLock;
/// use mea::rwlock::RwLockReadGuard;
///
/// #[derive(Debug)]
/// struct User {
/// id: u32,
/// profile: UserProfile,
/// }
///
/// #[derive(Debug)]
/// struct UserProfile {
/// email: String,
/// name: String,
/// }
///
/// let user = User {
/// id: 1,
/// profile: UserProfile {
/// email: "user@example.com".to_owned(),
/// name: "Alice".to_owned(),
/// },
/// };
///
/// let rwlock = RwLock::new(user);
/// let guard = rwlock.read().await;
/// // First map to the profile field
/// let profile_guard = RwLockReadGuard::map(guard, |user| &user.profile);
/// // Then map to the email field specifically
/// let email_guard = MappedRwLockReadGuard::map(profile_guard, |profile| &profile.email);
///
/// assert_eq!(&*email_guard, "user@example.com");
/// # }
/// ```
pub fn map<U, F>(orig: Self, f: F) -> MappedRwLockReadGuard<'a, U>
where
F: FnOnce(&T) -> &U,
U: ?Sized,
{
// SAFETY: orig.d is a valid NonNull<T> pointer that was created from a valid reference
// when the original MappedRwLockReadGuard was constructed. The guard guarantees shared
// access to the data through the rwlock, so dereferencing is safe.
let d = NonNull::from(f(unsafe { orig.d.as_ref() }));
let orig = std::mem::ManuallyDrop::new(orig);
MappedRwLockReadGuard::new(d, orig.s)
}
/// Attempts to make a new [`MappedRwLockReadGuard`] for a component of the locked data. The
/// original guard is returned if the closure returns `None`.
///
/// This operation cannot fail as the `MappedRwLockReadGuard` passed in already locked the
/// rwlock.
///
/// This is an associated function that needs to be used as
/// `MappedRwLockReadGuard::filter_map(...)`.
///
/// A method would interfere with methods of the same name on the contents of the locked data.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::rwlock::MappedRwLockReadGuard;
/// use mea::rwlock::RwLock;
/// use mea::rwlock::RwLockReadGuard;
///
/// #[derive(Debug)]
/// struct Person {
/// name: String,
/// email: Option<String>,
/// }
///
/// let person = Person {
/// name: "Alice".to_owned(),
/// email: Some("alice@example.com".to_owned()),
/// };
///
/// let rwlock = RwLock::new(person);
/// let guard = rwlock.read().await;
/// let name_guard = RwLockReadGuard::map(guard, |person| &person.name);
///
/// // Try to map to the email if it exists
/// let person_guard = rwlock.read().await;
/// let email_result = MappedRwLockReadGuard::filter_map(
/// RwLockReadGuard::map(person_guard, |person| &person.email),
/// |email_opt| email_opt.as_ref(),
/// );
///
/// match email_result {
/// Ok(email_guard) => {
/// assert_eq!(&*email_guard, "alice@example.com");
/// }
/// Err(_original_guard) => {
/// // Email was None, original guard is returned
/// println!("No email available");
/// }
/// }
/// # }
/// ```
pub fn filter_map<U, F>(orig: Self, f: F) -> Result<MappedRwLockReadGuard<'a, U>, Self>
where
F: FnOnce(&T) -> Option<&U>,
U: ?Sized,
{
// SAFETY: orig.d is a valid NonNull<T> pointer that was created from a valid reference
// when the original MappedRwLockReadGuard was constructed. The guard guarantees shared
// access to the data through the rwlock, so dereferencing is safe.
match f(unsafe { orig.d.as_ref() }) {
Some(d) => {
let d = NonNull::from(d);
let orig = std::mem::ManuallyDrop::new(orig);
Ok(MappedRwLockReadGuard::new(d, orig.s))
}
None => Err(orig),
}
}
}
@@ -0,0 +1,346 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::fmt;
use std::marker::PhantomData;
use std::mem::ManuallyDrop;
use std::ops::Deref;
use std::ops::DerefMut;
use std::ptr::NonNull;
use crate::internal;
use crate::rwlock::MappedRwLockReadGuard;
/// RAII structure used to release the exclusive write access of a lock when dropped, for a mapped
/// component of the locked data.
///
/// This structure is created by the [`map`] and [`filter_map`] methods on [`RwLockWriteGuard`]. It
/// allows you to hold a write lock on a subfield of the protected data, enabling more fine-grained
/// access control while maintaining the same locking semantics.
///
/// As long as you have this guard, you have exclusive write access to the underlying `T`. The guard
/// internally keeps a reference to the original rwlock's semaphore and tracks the number of permits
/// acquired, so the original lock is maintained until this guard is dropped.
///
/// `MappedRwLockWriteGuard` implements [`Send`] when the underlying data type implements [`Send`],
/// and implements [`Sync`] when the underlying data type implements both [`Send`] and [`Sync`],
/// allowing it to be used across task boundaries and shared between threads safely.
///
/// [`map`]: crate::rwlock::RwLockWriteGuard::map
/// [`filter_map`]: crate::rwlock::RwLockWriteGuard::filter_map
/// [`RwLockWriteGuard`]: crate::rwlock::RwLockWriteGuard
///
/// See the [module level documentation](crate::rwlock) for more.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::rwlock::RwLock;
/// use mea::rwlock::RwLockWriteGuard;
///
/// #[derive(Debug)]
/// struct User {
/// id: u32,
/// profile: UserProfile,
/// }
///
/// #[derive(Debug)]
/// struct UserProfile {
/// email: String,
/// name: String,
/// }
///
/// let user = User {
/// id: 1,
/// profile: UserProfile {
/// email: "user@example.com".to_owned(),
/// name: "Alice".to_owned(),
/// },
/// };
///
/// let rwlock = RwLock::new(user);
/// let mut guard = rwlock.write().await;
/// let mut profile_guard = RwLockWriteGuard::map(guard, |user| &mut user.profile);
///
/// // Now we can only access and modify the user's profile
/// profile_guard.email = "newemail@example.com".to_owned();
/// assert_eq!(profile_guard.email, "newemail@example.com");
/// # }
/// ```
#[must_use = "if unused the RwLock will immediately unlock"]
pub struct MappedRwLockWriteGuard<'a, T: ?Sized> {
d: NonNull<T>,
s: &'a internal::Semaphore,
permits_acquired: usize,
variance: PhantomData<&'a mut T>,
}
// SAFETY: A `&MappedRwLockWriteGuard` can be safely shared between threads because it provides
// exclusive access to the data, and the `T: Send + Sync` bound prevents data races.
unsafe impl<T: ?Sized + Send + Sync> Sync for MappedRwLockWriteGuard<'_, T> {}
// SAFETY: `MappedRwLockWriteGuard` owns the lock and can be safely sent to another thread.
// The `T: Send` bound ensures that the data can be safely accessed by the new thread,
// and the guard's lifetime guarantees that the data remains valid.
unsafe impl<T: ?Sized + Send> Send for MappedRwLockWriteGuard<'_, T> {}
impl<'a, T: ?Sized> MappedRwLockWriteGuard<'a, T> {
pub(crate) fn new(d: NonNull<T>, s: &'a internal::Semaphore, permits_acquired: usize) -> Self {
Self {
d,
s,
permits_acquired,
variance: PhantomData,
}
}
}
impl<T: ?Sized> Drop for MappedRwLockWriteGuard<'_, T> {
fn drop(&mut self) {
self.s.release(self.permits_acquired);
}
}
impl<T: ?Sized + fmt::Debug> fmt::Debug for MappedRwLockWriteGuard<'_, T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Debug::fmt(&**self, f)
}
}
impl<T: ?Sized + fmt::Display> fmt::Display for MappedRwLockWriteGuard<'_, T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Display::fmt(&**self, f)
}
}
impl<T: ?Sized> Deref for MappedRwLockWriteGuard<'_, T> {
type Target = T;
fn deref(&self) -> &Self::Target {
// SAFETY: we hold the write lock and the NonNull pointer is valid for the guard's lifetime
unsafe { self.d.as_ref() }
}
}
impl<T: ?Sized> DerefMut for MappedRwLockWriteGuard<'_, T> {
fn deref_mut(&mut self) -> &mut Self::Target {
// SAFETY: we hold the write lock and the NonNull pointer is valid for the guard's lifetime
unsafe { self.d.as_mut() }
}
}
impl<'a, T: ?Sized> MappedRwLockWriteGuard<'a, T> {
/// Makes a new [`MappedRwLockWriteGuard`] for a component of the locked data.
///
/// This operation cannot fail as the `MappedRwLockWriteGuard` passed in already locked the
/// rwlock.
///
/// This is an associated function that needs to be used as `MappedRwLockWriteGuard::map(...)`.
///
/// A method would interfere with methods of the same name on the contents of the locked
/// data.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::rwlock::MappedRwLockWriteGuard;
/// use mea::rwlock::RwLock;
/// use mea::rwlock::RwLockWriteGuard;
///
/// #[derive(Debug)]
/// struct User {
/// id: u32,
/// profile: UserProfile,
/// }
///
/// #[derive(Debug)]
/// struct UserProfile {
/// email: String,
/// name: String,
/// }
///
/// let user = User {
/// id: 1,
/// profile: UserProfile {
/// email: "user@example.com".to_owned(),
/// name: "Alice".to_owned(),
/// },
/// };
///
/// let rwlock = RwLock::new(user);
/// let mut guard = rwlock.write().await;
/// // First map to the profile field
/// let mut profile_guard = RwLockWriteGuard::map(guard, |user| &mut user.profile);
/// // Then map to the email field specifically
/// let mut email_guard = MappedRwLockWriteGuard::map(profile_guard, |profile| &mut profile.email);
///
/// *email_guard = "newemail@example.com".to_owned();
/// assert_eq!(&*email_guard, "newemail@example.com");
/// # }
/// ```
pub fn map<U, F>(mut orig: Self, f: F) -> MappedRwLockWriteGuard<'a, U>
where
F: FnOnce(&mut T) -> &mut U,
U: ?Sized,
{
// SAFETY: orig.d is a valid NonNull<T> pointer that was created from a valid reference
// when the original MappedRwLockWriteGuard was constructed. The guard guarantees exclusive
// access to the data through the rwlock, so dereferencing is safe.
let d = NonNull::from(f(unsafe { orig.d.as_mut() }));
let permits_acquired = orig.permits_acquired;
let orig = ManuallyDrop::new(orig);
MappedRwLockWriteGuard::new(d, orig.s, permits_acquired)
}
/// Attempts to make a new [`MappedRwLockWriteGuard`] for a component of the locked data. The
/// original guard is returned if the closure returns `None`.
///
/// This operation cannot fail as the `MappedRwLockWriteGuard` passed in already locked the
/// rwlock.
///
/// This is an associated function that needs to be used as
/// `MappedRwLockWriteGuard::filter_map(...)`.
///
/// A method would interfere with methods of the same name on the contents of the locked data.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::rwlock::MappedRwLockWriteGuard;
/// use mea::rwlock::RwLock;
/// use mea::rwlock::RwLockWriteGuard;
///
/// #[derive(Debug)]
/// struct Document {
/// title: String,
/// content: String,
/// metadata: Option<Metadata>,
/// }
///
/// #[derive(Debug)]
/// struct Metadata {
/// author: String,
/// version: Option<u32>,
/// }
///
/// let doc = Document {
/// title: "My Document".to_owned(),
/// content: "Initial content".to_owned(),
/// metadata: Some(Metadata {
/// author: "Alice".to_owned(),
/// version: Some(1),
/// }),
/// };
///
/// let rwlock = RwLock::new(doc);
/// let mut guard = rwlock.write().await;
///
/// // First map to the metadata field
/// let meta_guard = RwLockWriteGuard::map(guard, |doc| &mut doc.metadata);
///
/// // Try to map to the version number if metadata and version both exist
/// let version_result = MappedRwLockWriteGuard::filter_map(meta_guard, |meta_opt| {
/// meta_opt.as_mut()?.version.as_mut()
/// });
/// match version_result {
/// Ok(mut version_guard) => {
/// *version_guard += 1; // Increment version
/// assert_eq!(*version_guard, 2);
/// }
/// Err(_) => {
/// // Handle case where metadata or version doesn't exist
/// println!("No version to update");
/// }
/// }
/// # }
/// ```
pub fn filter_map<U, F>(mut orig: Self, f: F) -> Result<MappedRwLockWriteGuard<'a, U>, Self>
where
F: FnOnce(&mut T) -> Option<&mut U>,
U: ?Sized,
{
// SAFETY: orig.d is a valid NonNull<T> pointer that was created from a valid reference
// when the original MappedRwLockWriteGuard was constructed. The guard guarantees exclusive
// access to the data through the rwlock, so dereferencing is safe.
match f(unsafe { orig.d.as_mut() }) {
Some(d) => {
let d = NonNull::from(d);
let permits_acquired = orig.permits_acquired;
let orig = ManuallyDrop::new(orig);
Ok(MappedRwLockWriteGuard::new(d, orig.s, permits_acquired))
}
None => Err(orig),
}
}
/// Atomically downgrades the write lock to a read lock while preserving the mapping.
///
/// This method changes the lock from exclusive mode to shared mode atomically,
/// preventing other writers from acquiring the lock in between.
///
/// The returned `MappedRwLockReadGuard` preserves the original mapping to the specific
/// component of the data.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::rwlock::RwLock;
/// use mea::rwlock::RwLockWriteGuard;
///
/// #[derive(Debug)]
/// struct Counter {
/// value: i32,
/// name: String,
/// }
///
/// let lock = Arc::new(RwLock::new(Counter {
/// value: 0,
/// name: "counter".to_owned(),
/// }));
///
/// let write_guard = lock.write().await;
/// let mut value_write_guard = RwLockWriteGuard::map(write_guard, |counter| &mut counter.value);
/// *value_write_guard = 42;
///
/// let value_read_guard = value_write_guard.downgrade();
/// assert_eq!(*value_read_guard, 42);
///
/// assert!(lock.try_write().is_none());
///
/// drop(value_read_guard);
/// assert!(lock.try_write().is_some());
/// # }
/// ```
pub fn downgrade(self) -> MappedRwLockReadGuard<'a, T> {
// Prevent the original write guard from running its Drop implementation,
// which would release all permits. This must be done BEFORE any operation
// that might panic to ensure panic safety.
let guard = ManuallyDrop::new(self);
// Release max_readers - 1 permits to convert the write lock to a read lock.
guard.s.release(guard.permits_acquired - 1);
// Create the mapped read guard with 1 permit (standard for read locks)
MappedRwLockReadGuard::new(guard.d, guard.s)
}
}
+207
View File
@@ -0,0 +1,207 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//! A reader-writer lock that allows multiple readers or a single writer at a time.
//!
//! This type of lock allows a number of readers or at most one writer at any point in time. The
//! write portion of this lock typically allows modification of the underlying data (exclusive
//! access) and the read portion of this lock typically allows for read-only access (shared access).
//!
//! In comparison, a [`Mutex`] does not distinguish between readers or writers that acquire the
//! lock, therefore causing any tasks waiting for the lock to become available to yield. An RwLock
//! will allow any number of readers to acquire the lock as long as a writer is not holding the
//! lock.
//!
//! The priority policy of Tokio's read-write lock is fair (or [write-preferring]), in order to
//! ensure that readers cannot starve writers. Fairness is ensured using a first-in, first-out queue
//! for the tasks awaiting the lock; if a task that wishes to acquire the write lock is at the head
//! of the queue, read locks will not be given out until the write lock has been released. This is
//! in contrast to the Rust standard library's `std::sync::RwLock`, where the priority policy is
//! dependent on the operating system's implementation.
//!
//! The type parameter `T` represents the data that this lock protects. It is required that `T`
//! satisfies [`Send`] to be shared across threads. The RAII guards returned from the locking
//! methods implement [`Deref`] (and [`DerefMut`] for the `write` method) to allow access to the
//! content of the lock.
//!
//! # Examples
//!
//! ```
//! # #[tokio::main]
//! # async fn main() {
//! use mea::rwlock::RwLock;
//!
//! let lock = RwLock::new(5);
//!
//! // many reader locks can be held at once
//! {
//! let r1 = lock.read().await;
//! let r2 = lock.read().await;
//! assert_eq!(*r1, 5);
//! assert_eq!(*r2, 5);
//! } // read locks are dropped at this point
//!
//! // only one write lock may be held, however
//! {
//! let mut w = lock.write().await;
//! *w += 1;
//! assert_eq!(*w, 6);
//! } // write lock is dropped here
//!
//! # }
//! ```
//!
//! [`Mutex`]: crate::mutex::Mutex
//! [`Deref`]: std::ops::Deref
//! [`DerefMut`]: std::ops::DerefMut
//! [write-preferring]: https://en.wikipedia.org/wiki/Readers%E2%80%93writer_lock#Priority_policies
use std::cell::UnsafeCell;
use std::fmt;
use std::num::NonZeroUsize;
use crate::internal::Semaphore;
mod mapped_read_guard;
pub use mapped_read_guard::MappedRwLockReadGuard;
mod mapped_write_guard;
pub use mapped_write_guard::MappedRwLockWriteGuard;
mod owned_mapped_read_guard;
pub use owned_mapped_read_guard::OwnedMappedRwLockReadGuard;
mod owned_mapped_write_guard;
pub use owned_mapped_write_guard::OwnedMappedRwLockWriteGuard;
mod owned_read_guard;
pub use owned_read_guard::OwnedRwLockReadGuard;
mod owned_write_guard;
pub use owned_write_guard::OwnedRwLockWriteGuard;
mod read_guard;
pub use read_guard::RwLockReadGuard;
mod write_guard;
pub use write_guard::RwLockWriteGuard;
#[cfg(test)]
mod test;
/// A reader-writer lock that allows multiple readers or a single writer at a time.
///
/// See the [module level documentation](self) for more.
pub struct RwLock<T: ?Sized> {
/// Maximum number of concurrent readers.
///
/// This is ensured to be non-zero.
max_readers: usize,
/// Semaphore to coordinate read and write access to T
s: Semaphore,
/// The inner data.
c: UnsafeCell<T>,
}
unsafe impl<T: ?Sized + Send> Send for RwLock<T> {}
unsafe impl<T: ?Sized + Send + Sync> Sync for RwLock<T> {}
impl<T> From<T> for RwLock<T> {
fn from(t: T) -> Self {
Self::new(t)
}
}
impl<T: Default> Default for RwLock<T> {
fn default() -> Self {
Self::new(T::default())
}
}
impl<T: ?Sized + fmt::Debug> fmt::Debug for RwLock<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let mut d = f.debug_struct("RwLock");
match self.try_read() {
Some(inner) => d.field("data", &&*inner),
None => d.field("data", &format_args!("<locked>")),
};
d.finish()
}
}
impl<T> RwLock<T> {
/// Creates a new reader-writer lock in an unlocked state ready for use.
///
/// # Examples
///
/// ```
/// use mea::rwlock::RwLock;
///
/// let rwlock = RwLock::new(5);
/// ```
pub const fn new(t: T) -> RwLock<T> {
// large enough while not touch the edge
RwLock::with_max_readers(t, NonZeroUsize::new(usize::MAX >> 1).unwrap())
}
/// Creates a new reader-writer lock in an unlocked state, and allows a maximum of
/// `max_readers` concurrent readers.
///
/// This method is typically used for debugging and testing purposes.
///
/// # Examples
///
/// ```
/// use std::num::NonZeroUsize;
///
/// use mea::rwlock::RwLock;
///
/// let max_readers = NonZeroUsize::new(1024).expect("max_readers must be non-zero");
/// let rwlock = RwLock::with_max_readers(5, max_readers);
/// ```
pub const fn with_max_readers(t: T, max_readers: NonZeroUsize) -> RwLock<T> {
let max_readers = max_readers.get();
let s = Semaphore::new(max_readers);
let c = UnsafeCell::new(t);
RwLock { max_readers, c, s }
}
/// Consumes the lock, returning the underlying data.
///
/// # Examples
///
/// ```
/// use mea::rwlock::RwLock;
///
/// let lock = RwLock::new(1);
/// let n = lock.into_inner();
/// assert_eq!(n, 1);
/// ```
pub fn into_inner(self) -> T {
self.c.into_inner()
}
}
impl<T: ?Sized> RwLock<T> {
/// Returns a mutable reference to the underlying data.
///
/// Since this call borrows the `RwLock` mutably, no actual locking needs to take place: the
/// mutable borrow statically guarantees no locks exist.
///
/// # Examples
///
/// ```
/// use mea::rwlock::RwLock;
///
/// let mut lock = RwLock::new(1);
/// let n = lock.get_mut();
/// *n = 2;
/// ```
pub fn get_mut(&mut self) -> &mut T {
self.c.get_mut()
}
}
@@ -0,0 +1,302 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::fmt;
use std::marker::PhantomData;
use std::mem::ManuallyDrop;
use std::ops::Deref;
use std::ptr::NonNull;
use std::sync::Arc;
use crate::rwlock::RwLock;
/// Owned RAII structure used to release the shared read access of a lock when dropped, for a mapped
/// component of the locked data.
///
/// This guard is only available from a [`RwLock`] that is wrapped in an [`Arc`]. It is similar to
/// [`MappedRwLockReadGuard`], except that rather than borrowing the `RwLock`, it clones the `Arc`,
/// incrementing the reference count. This means that unlike `MappedRwLockReadGuard`, it will have
/// the `'static` lifetime.
///
/// As long as you have this guard, you have shared read access to the underlying `U`. The guard
/// internally keeps an `Arc` reference to the original rwlock, so the original lock is
/// maintained until this guard is dropped.
///
/// `OwnedMappedRwLockReadGuard` implements [`Send`] and [`Sync`]
/// when the underlying data type supports these traits, allowing it to be used across task
/// boundaries and shared between threads safely.
///
/// [`map`]: crate::rwlock::OwnedRwLockReadGuard::map
/// [`filter_map`]: crate::rwlock::OwnedRwLockReadGuard::filter_map
/// [`OwnedRwLockReadGuard`]: crate::rwlock::OwnedRwLockReadGuard
/// [`MappedRwLockReadGuard`]: crate::rwlock::MappedRwLockReadGuard
///
/// See the [module level documentation](crate::rwlock) for more.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::rwlock::OwnedRwLockReadGuard;
/// use mea::rwlock::RwLock;
///
/// #[derive(Debug)]
/// struct User {
/// id: u32,
/// profile: UserProfile,
/// }
///
/// #[derive(Debug)]
/// struct UserProfile {
/// email: String,
/// name: String,
/// }
///
/// let user = User {
/// id: 1,
/// profile: UserProfile {
/// email: "user@example.com".to_owned(),
/// name: "Alice".to_owned(),
/// },
/// };
///
/// let rwlock = Arc::new(RwLock::new(user));
/// let guard = rwlock.read_owned().await;
/// let profile_guard = OwnedRwLockReadGuard::map(guard, |user| &user.profile);
///
/// // Now we can only access the user's profile
/// assert_eq!(profile_guard.email, "user@example.com");
/// # }
/// ```
#[must_use = "if unused the RwLock will immediately unlock"]
pub struct OwnedMappedRwLockReadGuard<T: ?Sized, U: ?Sized> {
// This Arc acts as an ownership certificate, ensuring the RwLock remains valid
// and the lock is not released
lock: Arc<RwLock<T>>,
// This NonNull pointer precisely points to the subfield U, telling us which
// memory location we can operate on
d: NonNull<U>,
variance: PhantomData<fn() -> U>,
}
// SAFETY: Arc<RwLock<T>> is Send when T: Send + Sync, and we only provide shared access (&U)
// through deref(), so U: Sync is sufficient for safe cross-thread transfer.
unsafe impl<T: ?Sized + Send + Sync, U: ?Sized + Sync> Send for OwnedMappedRwLockReadGuard<T, U> {}
// SAFETY: OwnedMappedRwLockReadGuard can be safely shared between threads when T: Send + Sync and
// U: Sync. Multiple threads can hold &OwnedMappedRwLockReadGuard and call deref() concurrently,
// which only returns &U.
unsafe impl<T: ?Sized + Send + Sync, U: ?Sized + Sync> Sync for OwnedMappedRwLockReadGuard<T, U> {}
impl<T: ?Sized, U: ?Sized> OwnedMappedRwLockReadGuard<T, U> {
pub(crate) fn new(d: NonNull<U>, lock: Arc<RwLock<T>>) -> Self {
Self {
d,
lock,
variance: PhantomData,
}
}
}
impl<T: ?Sized, U: ?Sized> Drop for OwnedMappedRwLockReadGuard<T, U> {
fn drop(&mut self) {
self.lock.s.release(1);
}
}
impl<T: ?Sized, U: ?Sized + fmt::Debug> fmt::Debug for OwnedMappedRwLockReadGuard<T, U> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Debug::fmt(&**self, f)
}
}
impl<T: ?Sized, U: ?Sized + fmt::Display> fmt::Display for OwnedMappedRwLockReadGuard<T, U> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Display::fmt(&**self, f)
}
}
impl<T: ?Sized, U: ?Sized> Deref for OwnedMappedRwLockReadGuard<T, U> {
type Target = U;
fn deref(&self) -> &Self::Target {
// SAFETY: we hold the read lock and the NonNull pointer is valid for the guard's lifetime
unsafe { self.d.as_ref() }
}
}
impl<T: ?Sized, U: ?Sized> OwnedMappedRwLockReadGuard<T, U> {
/// Makes a new [`OwnedMappedRwLockReadGuard`] for a component of the locked data.
///
/// This operation cannot fail as the `OwnedMappedRwLockReadGuard` passed in already locked the
/// rwlock.
///
/// This is an associated function that needs to be used as
/// `OwnedMappedRwLockReadGuard::map(...)`.
///
/// A method would interfere with methods of the same name on the contents of the locked data.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::rwlock::OwnedMappedRwLockReadGuard;
/// use mea::rwlock::OwnedRwLockReadGuard;
/// use mea::rwlock::RwLock;
///
/// #[derive(Debug)]
/// struct ServerStats {
/// uptime: u64,
/// connection_info: ConnectionInfo,
/// }
///
/// #[derive(Debug)]
/// struct ConnectionInfo {
/// active_connections: u32,
/// max_connections: u32,
/// }
///
/// let stats = ServerStats {
/// uptime: 86400, // 1 day in seconds
/// connection_info: ConnectionInfo {
/// active_connections: 150,
/// max_connections: 1000,
/// },
/// };
///
/// let rwlock = Arc::new(RwLock::new(stats));
/// let guard = rwlock.read_owned().await;
/// // Map to connection info for cross-task monitoring
/// let conn_guard = OwnedRwLockReadGuard::map(guard, |stats| &stats.connection_info);
/// // Further map to active connections count
/// let active_guard = OwnedMappedRwLockReadGuard::map(conn_guard, |conn| &conn.active_connections);
///
/// assert_eq!(*active_guard, 150);
/// # }
/// ```
pub fn map<V, F>(orig: Self, f: F) -> OwnedMappedRwLockReadGuard<T, V>
where
F: FnOnce(&U) -> &V,
V: ?Sized,
{
// SAFETY: orig.d is a valid NonNull<U> pointer that was created from a valid reference
// when the original OwnedMappedRwLockReadGuard was constructed. The guard guarantees shared
// access to the data through the rwlock, so dereferencing is safe.
let d = NonNull::from(f(unsafe { orig.d.as_ref() }));
let orig = ManuallyDrop::new(orig);
// SAFETY: The original guard is wrapped in `ManuallyDrop` and will not be dropped.
// This allows us to safely move the `Arc` out of it and transfer ownership to the new
// guard.
let lock = unsafe { std::ptr::read(&orig.lock) };
OwnedMappedRwLockReadGuard::new(d, lock)
}
/// Attempts to make a new [`OwnedMappedRwLockReadGuard`] for a component of the locked data.
/// The original guard is returned if the closure returns `None`.
///
/// This operation cannot fail as the `OwnedMappedRwLockReadGuard` passed in already locked the
/// rwlock.
///
/// This is an associated function that needs to be used as
/// `OwnedMappedRwLockReadGuard::filter_map(...)`.
///
/// A method would interfere with methods of the same name on the contents of the locked data.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::collections::HashMap;
/// use std::sync::Arc;
///
/// use mea::rwlock::OwnedMappedRwLockReadGuard;
/// use mea::rwlock::OwnedRwLockReadGuard;
/// use mea::rwlock::RwLock;
///
/// #[derive(Debug)]
/// struct Cache {
/// entries: HashMap<String, CacheEntry>,
/// stats: CacheStats,
/// }
///
/// #[derive(Debug)]
/// struct CacheEntry {
/// data: String,
/// metadata: Option<String>,
/// }
///
/// #[derive(Debug)]
/// struct CacheStats {
/// hits: u64,
/// }
///
/// let mut entries = HashMap::new();
/// entries.insert(
/// "key1".to_owned(),
/// CacheEntry {
/// data: "cached_data".to_owned(),
/// metadata: Some("important".to_owned()),
/// },
/// );
///
/// let cache = Cache {
/// entries,
/// stats: CacheStats { hits: 42 },
/// };
///
/// let rwlock = Arc::new(RwLock::new(cache));
/// let guard = rwlock.read_owned().await;
///
/// // Map to a specific cache entry for cross-task reading
/// let entry_guard = OwnedRwLockReadGuard::map(guard, |cache| cache.entries.get("key1").unwrap());
///
/// // Try to map to the metadata if it exists
/// let metadata_guard =
/// OwnedMappedRwLockReadGuard::filter_map(entry_guard, |entry| entry.metadata.as_ref())
/// .expect("entry should have metadata");
///
/// assert_eq!(&*metadata_guard, "important");
/// # }
/// ```
pub fn filter_map<V, F>(orig: Self, f: F) -> Result<OwnedMappedRwLockReadGuard<T, V>, Self>
where
F: FnOnce(&U) -> Option<&V>,
V: ?Sized,
{
// SAFETY: orig.d is a valid NonNull<U> pointer that was created from a valid reference
// when the original OwnedMappedRwLockReadGuard was constructed. The guard guarantees shared
// access to the data through the rwlock, so dereferencing is safe.
match f(unsafe { orig.d.as_ref() }) {
Some(d) => {
let d = NonNull::from(d);
let orig = ManuallyDrop::new(orig);
// SAFETY: The original guard is wrapped in `ManuallyDrop` and will not be dropped.
// This allows us to safely move the `Arc` out of it and transfer ownership to the
// new guard.
let lock = unsafe { std::ptr::read(&orig.lock) };
Ok(OwnedMappedRwLockReadGuard::new(d, lock))
}
None => Err(orig),
}
}
}
@@ -0,0 +1,376 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::fmt;
use std::marker::PhantomData;
use std::mem::ManuallyDrop;
use std::ops::Deref;
use std::ops::DerefMut;
use std::ptr::NonNull;
use std::sync::Arc;
use crate::rwlock::OwnedMappedRwLockReadGuard;
use crate::rwlock::RwLock;
/// Owned RAII structure used to release the exclusive write access of a lock when dropped, for a
/// mapped component of the locked data.
///
/// This guard is only available from a [`RwLock`] that is wrapped in an [`Arc`]. It is similar to
/// [`MappedRwLockWriteGuard`], except that rather than borrowing the `RwLock`, it clones the `Arc`,
/// incrementing the reference count. This means that unlike `MappedRwLockWriteGuard`, it will have
/// the `'static` lifetime.
///
/// As long as you have this guard, you have exclusive write access to the underlying `T`. The guard
/// internally keeps an `Arc` reference to the original rwlock and tracks the number of permits
/// acquired, so the original lock is maintained until this guard is dropped.
///
/// `OwnedMappedRwLockWriteGuard` implements [`Send`] and [`Sync`]
/// when the underlying data type supports these traits, allowing it to be used across task
/// boundaries and shared between threads safely.
///
/// [`map`]: crate::rwlock::OwnedRwLockWriteGuard::map
/// [`filter_map`]: crate::rwlock::OwnedRwLockWriteGuard::filter_map
/// [`OwnedRwLockWriteGuard`]: crate::rwlock::OwnedRwLockWriteGuard
/// [`MappedRwLockWriteGuard`]: crate::rwlock::MappedRwLockWriteGuard
///
/// See the [module level documentation](crate::rwlock) for more.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::rwlock::OwnedRwLockWriteGuard;
/// use mea::rwlock::RwLock;
///
/// #[derive(Debug)]
/// struct User {
/// id: u32,
/// profile: UserProfile,
/// }
///
/// #[derive(Debug)]
/// struct UserProfile {
/// email: String,
/// name: String,
/// }
///
/// let user = User {
/// id: 1,
/// profile: UserProfile {
/// email: "user@example.com".to_owned(),
/// name: "Alice".to_owned(),
/// },
/// };
///
/// let rwlock = Arc::new(RwLock::new(user));
/// let mut guard = rwlock.write_owned().await;
/// let mut profile_guard = OwnedRwLockWriteGuard::map(guard, |user| &mut user.profile);
///
/// // Now we can only access and modify the user's profile
/// profile_guard.email = "newemail@example.com".to_owned();
/// assert_eq!(profile_guard.email, "newemail@example.com");
/// # }
/// ```
#[must_use = "if unused the RwLock will immediately unlock"]
pub struct OwnedMappedRwLockWriteGuard<T: ?Sized, U: ?Sized> {
d: NonNull<U>,
lock: Arc<RwLock<T>>,
permits_acquired: usize,
variance: PhantomData<*mut U>,
}
// SAFETY: Sharing &Guard across threads is safe when T: Send + Sync and U: Sync.
// Arc<RwLock<T>> requires T: Send + Sync for thread safety.
// &Guard only provides &U (via Deref), so U: Sync ensures safe concurrent access.
unsafe impl<T: ?Sized + Send + Sync, U: ?Sized + Sync> Sync for OwnedMappedRwLockWriteGuard<T, U> {}
// SAFETY: Sending Guard across threads is safe when T: Send + Sync and U: Send.
// Arc<RwLock<T>> requires T: Send + Sync to be Send.
// Guard transfers exclusive access to U, so U: Send ensures safe access from new thread.
unsafe impl<T: ?Sized + Send + Sync, U: ?Sized + Send> Send for OwnedMappedRwLockWriteGuard<T, U> {}
impl<T: ?Sized, U: ?Sized> OwnedMappedRwLockWriteGuard<T, U> {
pub(crate) fn new(d: NonNull<U>, lock: Arc<RwLock<T>>, permits_acquired: usize) -> Self {
Self {
d,
lock,
permits_acquired,
variance: PhantomData,
}
}
}
impl<T: ?Sized, U: ?Sized> Drop for OwnedMappedRwLockWriteGuard<T, U> {
fn drop(&mut self) {
self.lock.s.release(self.permits_acquired);
}
}
impl<T: ?Sized, U: ?Sized + fmt::Debug> fmt::Debug for OwnedMappedRwLockWriteGuard<T, U> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Debug::fmt(&**self, f)
}
}
impl<T: ?Sized, U: ?Sized + fmt::Display> fmt::Display for OwnedMappedRwLockWriteGuard<T, U> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Display::fmt(&**self, f)
}
}
impl<T: ?Sized, U: ?Sized> Deref for OwnedMappedRwLockWriteGuard<T, U> {
type Target = U;
fn deref(&self) -> &Self::Target {
// SAFETY: we hold the write lock and the NonNull pointer is valid for the guard's lifetime
unsafe { self.d.as_ref() }
}
}
impl<T: ?Sized, U: ?Sized> DerefMut for OwnedMappedRwLockWriteGuard<T, U> {
fn deref_mut(&mut self) -> &mut Self::Target {
// SAFETY: we hold the write lock and the NonNull pointer is valid for the guard's lifetime
unsafe { self.d.as_mut() }
}
}
impl<T: ?Sized, U: ?Sized> OwnedMappedRwLockWriteGuard<T, U> {
/// Makes a new [`OwnedMappedRwLockWriteGuard`] for a component of the locked data.
///
/// This operation cannot fail as the `OwnedMappedRwLockWriteGuard` passed in already locked the
/// rwlock.
///
/// This is an associated function that needs to be used as
/// `OwnedMappedRwLockWriteGuard::map(...)`.
///
/// A method would interfere with methods of the same name on the contents of the locked data.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::rwlock::OwnedMappedRwLockWriteGuard;
/// use mea::rwlock::OwnedRwLockWriteGuard;
/// use mea::rwlock::RwLock;
///
/// #[derive(Debug)]
/// struct User {
/// id: u32,
/// profile: UserProfile,
/// }
///
/// #[derive(Debug)]
/// struct UserProfile {
/// email: String,
/// name: String,
/// }
///
/// let user = User {
/// id: 1,
/// profile: UserProfile {
/// email: "user@example.com".to_owned(),
/// name: "Alice".to_owned(),
/// },
/// };
///
/// let rwlock = Arc::new(RwLock::new(user));
/// let mut guard = rwlock.write_owned().await;
/// // First map to the profile field
/// let mut profile_guard = OwnedRwLockWriteGuard::map(guard, |user| &mut user.profile);
/// // Then map to the email field specifically
/// let mut email_guard =
/// OwnedMappedRwLockWriteGuard::map(profile_guard, |profile| &mut profile.email);
///
/// *email_guard = "newemail@example.com".to_owned();
/// assert_eq!(&*email_guard, "newemail@example.com");
/// # }
/// ```
pub fn map<V, F>(mut orig: Self, f: F) -> OwnedMappedRwLockWriteGuard<T, V>
where
F: FnOnce(&mut U) -> &mut V,
V: ?Sized,
{
// SAFETY: orig.d is a valid NonNull<U> pointer that was created from a valid reference
// when the original OwnedMappedRwLockWriteGuard was constructed. The guard guarantees
// exclusive access to the data through the rwlock, so dereferencing is safe.
let d = NonNull::from(f(unsafe { orig.d.as_mut() }));
let orig = ManuallyDrop::new(orig);
let permits_acquired = orig.permits_acquired;
// SAFETY: The original guard is wrapped in `ManuallyDrop` and will not be dropped.
// This allows us to safely move the `Arc` out of it and transfer ownership to the new
// guard.
let lock = unsafe { std::ptr::read(&orig.lock) };
OwnedMappedRwLockWriteGuard::new(d, lock, permits_acquired)
}
/// Attempts to make a new [`OwnedMappedRwLockWriteGuard`] for a component of the locked data.
/// The original guard is returned if the closure returns `None`.
///
/// This operation cannot fail as the `OwnedMappedRwLockWriteGuard` passed in already locked the
/// rwlock.
///
/// This is an associated function that needs to be used as
/// `OwnedMappedRwLockWriteGuard::filter_map(...)`.
///
/// A method would interfere with methods of the same name on the contents of the locked data.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::rwlock::OwnedMappedRwLockWriteGuard;
/// use mea::rwlock::OwnedRwLockWriteGuard;
/// use mea::rwlock::RwLock;
///
/// #[derive(Debug)]
/// struct AppState {
/// user_count: u64,
/// metrics: Option<Metrics>,
/// }
///
/// #[derive(Debug)]
/// struct Metrics {
/// requests_per_second: f64,
/// error_rate: f64,
/// }
///
/// let state = AppState {
/// user_count: 100,
/// metrics: Some(Metrics {
/// requests_per_second: 150.5,
/// error_rate: 0.01,
/// }),
/// };
///
/// let rwlock = Arc::new(RwLock::new(state));
/// let guard = rwlock.write_owned().await;
///
/// // First, map to the `metrics` field, which is an Option.
/// // This gives us an OwnedMappedRwLockWriteGuard<AppState, Option<Metrics>>
/// let metrics_opt_guard = OwnedRwLockWriteGuard::map(guard, |state| &mut state.metrics);
///
/// // Now, on the mapped guard, try to map into the Option.
/// // This is the correct usage of OwnedMappedRwLockWriteGuard::filter_map.
/// let metrics_result =
/// OwnedMappedRwLockWriteGuard::filter_map(metrics_opt_guard, |metrics_opt| {
/// metrics_opt.as_mut()
/// });
///
/// match metrics_result {
/// Ok(mut metrics_guard) => {
/// // Update metrics across tasks
/// metrics_guard.requests_per_second = 200.0;
/// metrics_guard.error_rate = 0.005;
/// assert_eq!(metrics_guard.requests_per_second, 200.0);
/// }
/// Err(_original_guard) => {
/// // Metrics not available, original guard is returned
/// println!("Metrics not enabled");
/// }
/// }
/// # }
/// ```
pub fn filter_map<V, F>(mut orig: Self, f: F) -> Result<OwnedMappedRwLockWriteGuard<T, V>, Self>
where
F: FnOnce(&mut U) -> Option<&mut V>,
V: ?Sized,
{
// SAFETY: orig.d is a valid NonNull<U> pointer that was created from a valid reference
// when the original OwnedMappedRwLockWriteGuard was constructed. The guard guarantees
// exclusive access to the data through the rwlock, so dereferencing is safe.
match f(unsafe { orig.d.as_mut() }) {
Some(d) => {
let d = NonNull::from(d);
let orig = ManuallyDrop::new(orig);
let permits_acquired = orig.permits_acquired;
// SAFETY: The original guard is wrapped in `ManuallyDrop` and will not be dropped.
// This allows us to safely move the `Arc` out of it and transfer ownership to the
// new guard.
let lock = unsafe { std::ptr::read(&orig.lock) };
Ok(OwnedMappedRwLockWriteGuard::new(d, lock, permits_acquired))
}
None => Err(orig),
}
}
/// Atomically downgrades the write lock to a read lock while preserving the mapping.
///
/// This method changes the lock from exclusive mode to shared mode atomically,
/// preventing other writers from acquiring the lock in between.
///
/// The returned `OwnedMappedRwLockReadGuard` preserves the original mapping and
/// has a `'static` lifetime.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::rwlock::OwnedRwLockWriteGuard;
/// use mea::rwlock::RwLock;
///
/// #[derive(Debug)]
/// struct Database {
/// connection_count: u32,
/// status: String,
/// }
///
/// let db = Arc::new(RwLock::new(Database {
/// connection_count: 0,
/// status: "idle".to_owned(),
/// }));
///
/// let write_guard = db.clone().write_owned().await;
/// let mut count_write_guard =
/// OwnedRwLockWriteGuard::map(write_guard, |db| &mut db.connection_count);
/// *count_write_guard = 5;
///
/// let count_read_guard = count_write_guard.downgrade();
/// assert_eq!(*count_read_guard, 5);
///
/// assert!(db.clone().try_write_owned().is_none());
///
/// drop(count_read_guard);
/// assert!(db.clone().try_write_owned().is_some());
/// # }
/// ```
pub fn downgrade(self) -> OwnedMappedRwLockReadGuard<T, U> {
// Prevent the original write guard from running its Drop implementation,
// which would release all permits. This must be done BEFORE any operation
// that might panic to ensure panic safety.
let guard = ManuallyDrop::new(self);
// Release max_readers - 1 permits to convert the write lock to a read lock.
guard.lock.s.release(guard.permits_acquired - 1);
// SAFETY: The `guard` is wrapped in `ManuallyDrop`, so its destructor will not be run.
// We can safely move the `Arc` out of the guard, as the guard is not used after this.
// This is a standard way to transfer ownership from a `ManuallyDrop` wrapper.
let lock = unsafe { std::ptr::read(&guard.lock) };
OwnedMappedRwLockReadGuard::new(guard.d, lock)
}
}
+264
View File
@@ -0,0 +1,264 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::fmt;
use std::ops::Deref;
use std::sync::Arc;
use crate::rwlock::OwnedMappedRwLockReadGuard;
use crate::rwlock::RwLock;
impl<T: ?Sized> RwLock<T> {
/// Locks this `RwLock` with shared read access, causing the current task to yield until the
/// lock has been acquired.
///
/// The calling task will yield until there are no writers which hold the lock. There may be
/// other readers inside the lock when the task resumes.
///
/// This method is identical to [`RwLock::read`], except that the returned guard references the
/// `RwLock` with an [`Arc`] rather than by borrowing it. Therefore, the `RwLock` must be
/// wrapped in an `Arc` to call this method, and the guard will live for the `'static` lifetime,
/// as it keeps the `RwLock` alive by holding an `Arc`.
///
/// Note that under the priority policy of [`RwLock`], read locks are not granted until prior
/// write locks, to prevent starvation. Therefore, deadlock may occur if a read lock is held
/// by the current task, a write lock attempt is made, and then a subsequent read lock attempt
/// is made by the current task.
///
/// Returns an RAII guard which will drop this read access of the `RwLock` when dropped.
///
/// # Cancel safety
///
/// This method uses a queue to fairly distribute locks in the order they were requested.
/// Cancelling a call to `read_owned` makes you lose your place in the queue.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::rwlock::RwLock;
///
/// let lock = Arc::new(RwLock::new(1));
/// let lock_clone = lock.clone();
///
/// let n = lock.read_owned().await;
/// assert_eq!(*n, 1);
///
/// tokio::spawn(async move {
/// // while the outer read lock is held, we acquire a read lock, too
/// let r = lock_clone.read_owned().await;
/// assert_eq!(*r, 1);
/// })
/// .await
/// .unwrap();
/// # }
/// ```
pub async fn read_owned(self: Arc<Self>) -> OwnedRwLockReadGuard<T> {
self.s.acquire(1).await;
OwnedRwLockReadGuard { lock: self }
}
/// Attempts to acquire this `RwLock` with shared read access.
///
/// If the access couldn't be acquired immediately, returns `None`. Otherwise, an RAII guard is
/// returned which will release read access when dropped.
///
/// This method is identical to [`RwLock::try_read`], except that the returned guard references
/// the `RwLock` with an [`Arc`] rather than by borrowing it. Therefore, the `RwLock` must
/// be wrapped in an `Arc` to call this method, and the guard will live for the `'static`
/// lifetime, as it keeps the `RwLock` alive by holding an `Arc`.
///
/// # Examples
///
/// ```
/// use std::sync::Arc;
///
/// use mea::rwlock::RwLock;
///
/// let lock = Arc::new(RwLock::new(1));
///
/// let v = lock.clone().try_read_owned().unwrap();
/// assert_eq!(*v, 1);
/// drop(v);
///
/// let v = lock.try_write().unwrap();
/// assert!(lock.clone().try_read_owned().is_none());
/// ```
pub fn try_read_owned(self: Arc<Self>) -> Option<OwnedRwLockReadGuard<T>> {
if self.s.try_acquire(1) {
Some(OwnedRwLockReadGuard { lock: self })
} else {
None
}
}
}
/// Owned RAII structure used to release the shared read access of a lock when dropped.
///
/// This structure is created by the [`RwLock::read`] method.
///
/// See the [module level documentation](crate::rwlock) for more.
#[must_use = "if unused the RwLock will immediately unlock"]
pub struct OwnedRwLockReadGuard<T: ?Sized> {
pub(super) lock: Arc<RwLock<T>>,
}
unsafe impl<T: ?Sized + Sync> Send for OwnedRwLockReadGuard<T> {}
unsafe impl<T: ?Sized + Send + Sync> Sync for OwnedRwLockReadGuard<T> {}
impl<T: ?Sized> Drop for OwnedRwLockReadGuard<T> {
fn drop(&mut self) {
self.lock.s.release(1);
}
}
impl<T: ?Sized + fmt::Debug> fmt::Debug for OwnedRwLockReadGuard<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Debug::fmt(&**self, f)
}
}
impl<T: ?Sized + fmt::Display> fmt::Display for OwnedRwLockReadGuard<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Display::fmt(&**self, f)
}
}
impl<T: ?Sized> Deref for OwnedRwLockReadGuard<T> {
type Target = T;
fn deref(&self) -> &Self::Target {
unsafe { &*self.lock.c.get() }
}
}
impl<T: ?Sized> OwnedRwLockReadGuard<T> {
/// Makes a new [`OwnedMappedRwLockReadGuard`] for a component of the locked
/// data.
///
/// This operation cannot fail as the `OwnedRwLockReadGuard` passed in already locked the
/// rwlock.
///
/// This is an associated function that needs to be used as `OwnedRwLockReadGuard::map(...)`.
///
/// A method would interfere with methods of the same name on the contents of the locked data.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::rwlock::OwnedRwLockReadGuard;
/// use mea::rwlock::RwLock;
///
/// #[derive(Debug)]
/// struct Foo {
/// a: u32,
/// b: String,
/// }
///
/// let rwlock = Arc::new(RwLock::new(Foo {
/// a: 1,
/// b: "hello".to_owned(),
/// }));
///
/// let guard = rwlock.read_owned().await;
/// let mapped_guard = OwnedRwLockReadGuard::map(guard, |foo| &foo.a);
///
/// assert_eq!(*mapped_guard, 1);
/// # }
/// ```
pub fn map<U, F>(orig: Self, f: F) -> OwnedMappedRwLockReadGuard<T, U>
where
F: FnOnce(&T) -> &U,
U: ?Sized,
{
// SAFETY: orig.lock.c.get() is a valid pointer to T that was created when the lock was
// acquired. The guard guarantees shared access to the data through the rwlock, so
// dereferencing is safe.
let d = std::ptr::NonNull::from(f(unsafe { &*orig.lock.c.get() }));
let orig = std::mem::ManuallyDrop::new(orig);
// Safely extract the Arc from the guard
let lock = unsafe { std::ptr::read(&orig.lock) };
OwnedMappedRwLockReadGuard::new(d, lock)
}
/// Attempts to make a new [`OwnedMappedRwLockReadGuard`] for a component of the locked data.
/// The original guard is returned if the closure returns `None`.
///
/// This operation cannot fail as the `OwnedRwLockReadGuard` passed in already locked the
/// rwlock.
///
/// This is an associated function that needs to be used as
/// `OwnedRwLockReadGuard::filter_map(...)`.
///
/// A method would interfere with methods of the same name on the contents of the locked data.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::rwlock::OwnedRwLockReadGuard;
/// use mea::rwlock::RwLock;
///
/// #[derive(Debug)]
/// struct Foo {
/// a: u32,
/// b: String,
/// }
///
/// let rwlock = Arc::new(RwLock::new(Foo {
/// a: 1,
/// b: "hello".to_owned(),
/// }));
///
/// let guard = rwlock.read_owned().await;
/// let mapped_guard =
/// OwnedRwLockReadGuard::filter_map(guard, |foo| if foo.a > 0 { Some(&foo.b) } else { None })
/// .expect("should have mapped");
///
/// assert_eq!(&*mapped_guard, "hello");
/// # }
/// ```
pub fn filter_map<U, F>(orig: Self, f: F) -> Result<OwnedMappedRwLockReadGuard<T, U>, Self>
where
F: FnOnce(&T) -> Option<&U>,
U: ?Sized,
{
// SAFETY: orig.lock.c.get() is a valid pointer to T that was created when the lock was
// acquired. The guard guarantees shared access to the data through the rwlock, so
// dereferencing is safe.
match f(unsafe { &*orig.lock.c.get() }) {
Some(d) => {
let d = std::ptr::NonNull::from(d);
let orig = std::mem::ManuallyDrop::new(orig);
// Safely extract the Arc from the guard
let lock = unsafe { std::ptr::read(&orig.lock) };
Ok(OwnedMappedRwLockReadGuard::new(d, lock))
}
None => Err(orig),
}
}
}
+324
View File
@@ -0,0 +1,324 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::fmt;
use std::mem::ManuallyDrop;
use std::ops::Deref;
use std::ops::DerefMut;
use std::ptr::NonNull;
use std::sync::Arc;
use crate::rwlock::OwnedMappedRwLockWriteGuard;
use crate::rwlock::OwnedRwLockReadGuard;
use crate::rwlock::RwLock;
impl<T: ?Sized> RwLock<T> {
/// Locks this `RwLock` with exclusive write access, causing the current task to yield until the
/// lock has been acquired.
///
/// The calling task will yield while other writers or readers currently have access to the
/// lock.
///
/// This method is identical to [`RwLock::write`], except that the returned guard references the
/// `RwLock` with an [`Arc`] rather than by borrowing it. Therefore, the `RwLock` must be
/// wrapped in an `Arc` to call this method, and the guard will live for the `'static` lifetime,
/// as it keeps the `RwLock` alive by holding an `Arc`.
///
/// Returns an RAII guard which will drop the write access of this `RwLock` when dropped.
///
/// # Cancel safety
///
/// This method uses a queue to fairly distribute locks in the order they were requested.
/// Cancelling a call to `write_owned` makes you lose your place in the queue.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::rwlock::RwLock;
///
/// let lock = Arc::new(RwLock::new(1));
/// let mut n = lock.write_owned().await;
/// *n = 2;
/// # }
/// ```
pub async fn write_owned(self: Arc<Self>) -> OwnedRwLockWriteGuard<T> {
self.s.acquire(self.max_readers).await;
OwnedRwLockWriteGuard {
permits_acquired: self.max_readers,
lock: self,
}
}
/// Attempts to acquire this `RwLock` with exclusive write access.
///
/// If the access couldn't be acquired immediately, returns `None`. Otherwise, an RAII guard is
/// returned which will release write access when dropped.
///
/// This method is identical to [`RwLock::try_write`], except that the returned guard references
/// the `RwLock` with an [`Arc`] rather than by borrowing it. Therefore, the `RwLock` must
/// be wrapped in an `Arc` to call this method, and the guard will live for the `'static`
/// lifetime, as it keeps the `RwLock` alive by holding an `Arc`.
///
/// # Examples
///
/// ```
/// use std::sync::Arc;
///
/// use mea::rwlock::RwLock;
///
/// let lock = Arc::new(RwLock::new(1));
///
/// let v = lock.try_read().unwrap();
/// assert!(lock.clone().try_write_owned().is_none());
/// drop(v);
///
/// let mut v = lock.try_write_owned().unwrap();
/// *v = 2;
/// ```
pub fn try_write_owned(self: Arc<Self>) -> Option<OwnedRwLockWriteGuard<T>> {
if self.s.try_acquire(self.max_readers) {
Some(OwnedRwLockWriteGuard {
permits_acquired: self.max_readers,
lock: self,
})
} else {
None
}
}
}
/// Owned RAII structure used to release the exclusive write access of a lock when dropped.
///
/// This structure is created by the [`RwLock::write`] method.
///
/// See the [module level documentation](crate::rwlock) for more.
#[must_use = "if unused the RwLock will immediately unlock"]
pub struct OwnedRwLockWriteGuard<T: ?Sized> {
pub(super) permits_acquired: usize,
pub(super) lock: Arc<RwLock<T>>,
}
unsafe impl<T: ?Sized + Send + Sync> Send for OwnedRwLockWriteGuard<T> {}
unsafe impl<T: ?Sized + Send + Sync> Sync for OwnedRwLockWriteGuard<T> {}
impl<T: ?Sized> Drop for OwnedRwLockWriteGuard<T> {
fn drop(&mut self) {
self.lock.s.release(self.permits_acquired);
}
}
impl<T: ?Sized + fmt::Debug> fmt::Debug for OwnedRwLockWriteGuard<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Debug::fmt(&**self, f)
}
}
impl<T: ?Sized + fmt::Display> fmt::Display for OwnedRwLockWriteGuard<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Display::fmt(&**self, f)
}
}
impl<T: ?Sized> Deref for OwnedRwLockWriteGuard<T> {
type Target = T;
fn deref(&self) -> &Self::Target {
unsafe { &*self.lock.c.get() }
}
}
impl<T: ?Sized> DerefMut for OwnedRwLockWriteGuard<T> {
fn deref_mut(&mut self) -> &mut Self::Target {
unsafe { &mut *self.lock.c.get() }
}
}
impl<T: ?Sized> OwnedRwLockWriteGuard<T> {
/// Makes a new [`OwnedMappedRwLockWriteGuard`] for a component of the locked
/// data.
///
/// This operation cannot fail as the `OwnedRwLockWriteGuard` passed in already locked the
/// rwlock.
///
/// This is an associated function that needs to be used as `OwnedRwLockWriteGuard::map(...)`.
///
/// A method would interfere with methods of the same name on the contents of the locked data.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::rwlock::OwnedRwLockWriteGuard;
/// use mea::rwlock::RwLock;
///
/// #[derive(Debug)]
/// struct Foo {
/// a: u32,
/// b: String,
/// }
///
/// let rwlock = Arc::new(RwLock::new(Foo {
/// a: 1,
/// b: "hello".to_owned(),
/// }));
///
/// let mut guard = rwlock.write_owned().await;
/// let mut mapped_guard = OwnedRwLockWriteGuard::map(guard, |foo| &mut foo.b);
///
/// mapped_guard.push_str(" world");
/// assert_eq!(&*mapped_guard, "hello world");
/// # }
/// ```
pub fn map<U, F>(orig: Self, f: F) -> OwnedMappedRwLockWriteGuard<T, U>
where
F: FnOnce(&mut T) -> &mut U,
U: ?Sized,
{
// SAFETY: We have exclusive write access to the data through the rwlock.
// The data pointer is valid for the lifetime of the guard.
let d = NonNull::from(f(unsafe { &mut *orig.lock.c.get() }));
let orig = ManuallyDrop::new(orig);
let permits_acquired = orig.permits_acquired;
// SAFETY: The original guard is wrapped in `ManuallyDrop` and will not be dropped.
// This allows us to safely move the `Arc` out of it and transfer ownership to the new
// guard.
let lock = unsafe { std::ptr::read(&orig.lock) };
OwnedMappedRwLockWriteGuard::new(d, lock, permits_acquired)
}
/// Attempts to make a new [`OwnedMappedRwLockWriteGuard`] for a component of the
/// locked data. The original guard is returned if the closure returns `None`.
///
/// This operation cannot fail as the `OwnedRwLockWriteGuard` passed in already locked the
/// rwlock.
///
/// This is an associated function that needs to be used as
/// `OwnedRwLockWriteGuard::filter_map(...)`.
///
/// A method would interfere with methods of the same name on the contents of the locked data.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::rwlock::OwnedRwLockWriteGuard;
/// use mea::rwlock::RwLock;
///
/// #[derive(Debug)]
/// struct Foo {
/// a: u32,
/// b: String,
/// }
///
/// let rwlock = Arc::new(RwLock::new(Foo {
/// a: 1,
/// b: "hello".to_owned(),
/// }));
///
/// let mut guard = rwlock.write_owned().await;
/// let mut mapped_guard = OwnedRwLockWriteGuard::filter_map(guard, |foo| {
/// if foo.b.len() > 3 {
/// Some(&mut foo.b)
/// } else {
/// None
/// }
/// })
/// .expect("should have mapped");
///
/// mapped_guard.push_str(" world");
/// assert_eq!(&*mapped_guard, "hello world");
/// # }
/// ```
pub fn filter_map<U, F>(orig: Self, f: F) -> Result<OwnedMappedRwLockWriteGuard<T, U>, Self>
where
F: FnOnce(&mut T) -> Option<&mut U>,
U: ?Sized,
{
// SAFETY: We have exclusive write access to the data through the rwlock.
// The data pointer is valid for the lifetime of the guard.
let d = match f(unsafe { &mut *orig.lock.c.get() }) {
Some(d) => NonNull::from(d),
None => return Err(orig),
};
let orig = ManuallyDrop::new(orig);
let permits_acquired = orig.permits_acquired;
// SAFETY: The original guard is wrapped in `ManuallyDrop` and will not be dropped.
// This allows us to safely move the `Arc` out of it and transfer ownership to the new
// guard.
let lock = unsafe { std::ptr::read(&orig.lock) };
Ok(OwnedMappedRwLockWriteGuard::new(d, lock, permits_acquired))
}
/// Atomically downgrades the write lock to a read lock.
///
/// This method changes the lock from exclusive mode to shared mode atomically,
/// preventing other writers from acquiring the lock in between.
///
/// The returned `OwnedRwLockReadGuard` has a `'static` lifetime, as it keeps
/// the `RwLock` alive by holding an `Arc`.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::rwlock::RwLock;
///
/// let lock = Arc::new(RwLock::new(1));
///
/// let mut write_guard = lock.clone().write_owned().await;
/// *write_guard = 42;
///
/// let read_guard = write_guard.downgrade();
/// assert_eq!(*read_guard, 42);
///
/// assert!(lock.clone().try_write_owned().is_none());
///
/// drop(read_guard);
/// assert!(lock.clone().try_write_owned().is_some());
/// # }
/// ```
pub fn downgrade(self) -> OwnedRwLockReadGuard<T> {
// Prevent the original write guard from running its Drop implementation,
// which would release all permits. This must be done BEFORE any operation
// that might panic to ensure panic safety.
let guard = ManuallyDrop::new(self);
// Release max_readers - 1 permits to convert the write lock to a read lock.
// The remaining 1 permit is kept for the read lock.
guard.lock.s.release(guard.permits_acquired - 1);
// SAFETY: The `guard` is wrapped in `ManuallyDrop`, so its destructor will not be run.
// We can safely move the `Arc` out of the guard, as the guard is not used after this.
// This is a standard way to transfer ownership from a `ManuallyDrop` wrapper.
let lock = unsafe { std::ptr::read(&guard.lock) };
OwnedRwLockReadGuard { lock }
}
}
+221
View File
@@ -0,0 +1,221 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::fmt;
use std::mem::ManuallyDrop;
use std::ops::Deref;
use std::ptr::NonNull;
use crate::rwlock::MappedRwLockReadGuard;
use crate::rwlock::RwLock;
impl<T: ?Sized> RwLock<T> {
/// Locks this `RwLock` with shared read access, causing the current task to yield until the
/// lock has been acquired.
///
/// The calling task will yield until there are no writers which hold the lock. There may be
/// other readers inside the lock when the task resumes.
///
/// Note that under the priority policy of [`RwLock`], read locks are not granted until prior
/// write locks, to prevent starvation. Therefore, deadlock may occur if a read lock is held
/// by the current task, a write lock attempt is made, and then a subsequent read lock attempt
/// is made by the current task.
///
/// Returns an RAII guard which will drop this read access of the `RwLock` when dropped.
///
/// # Cancel safety
///
/// This method uses a queue to fairly distribute locks in the order they were requested.
/// Cancelling a call to `read` makes you lose your place in the queue.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::rwlock::RwLock;
///
/// let lock = Arc::new(RwLock::new(1));
/// let lock_clone = lock.clone();
///
/// let n = lock.read().await;
/// assert_eq!(*n, 1);
///
/// tokio::spawn(async move {
/// // while the outer read lock is held, we acquire a read lock, too
/// let r = lock_clone.read().await;
/// assert_eq!(*r, 1);
/// })
/// .await
/// .unwrap();
/// # }
/// ```
pub async fn read(&self) -> RwLockReadGuard<'_, T> {
self.s.acquire(1).await;
RwLockReadGuard { lock: self }
}
/// Attempts to acquire this `RwLock` with shared read access.
///
/// If the access couldn't be acquired immediately, returns `None`. Otherwise, an RAII guard is
/// returned which will release read access when dropped.
///
/// # Examples
///
/// ```
/// use std::sync::Arc;
///
/// use mea::rwlock::RwLock;
///
/// let lock = Arc::new(RwLock::new(1));
///
/// let v = lock.try_read().unwrap();
/// assert_eq!(*v, 1);
/// drop(v);
///
/// let v = lock.try_write().unwrap();
/// assert!(lock.try_read().is_none());
/// ```
pub fn try_read(&self) -> Option<RwLockReadGuard<'_, T>> {
if self.s.try_acquire(1) {
Some(RwLockReadGuard { lock: self })
} else {
None
}
}
}
/// RAII structure used to release the shared read access of a lock when dropped.
///
/// This structure is created by the [`RwLock::read`] method.
///
/// See the [module level documentation](crate::rwlock) for more.
#[must_use = "if unused the RwLock will immediately unlock"]
pub struct RwLockReadGuard<'a, T: ?Sized> {
pub(super) lock: &'a RwLock<T>,
}
unsafe impl<T: ?Sized + Sync> Send for RwLockReadGuard<'_, T> {}
unsafe impl<T: ?Sized + Send + Sync> Sync for RwLockReadGuard<'_, T> {}
impl<T: ?Sized> Drop for RwLockReadGuard<'_, T> {
fn drop(&mut self) {
self.lock.s.release(1);
}
}
impl<T: ?Sized + fmt::Debug> fmt::Debug for RwLockReadGuard<'_, T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Debug::fmt(&**self, f)
}
}
impl<T: ?Sized + fmt::Display> fmt::Display for RwLockReadGuard<'_, T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Display::fmt(&**self, f)
}
}
impl<T: ?Sized> Deref for RwLockReadGuard<'_, T> {
type Target = T;
fn deref(&self) -> &Self::Target {
unsafe { &*self.lock.c.get() }
}
}
impl<'a, T: ?Sized> RwLockReadGuard<'a, T> {
/// Makes a new [`MappedRwLockReadGuard`] for a component of the locked data.
///
/// This operation cannot fail as the `RwLockReadGuard` passed in already locked the rwlock.
///
/// This is an associated function that needs to be used as `RwLockReadGuard::map(...)`.
///
/// A method would interfere with methods of the same name on the contents of the locked data.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::rwlock::RwLock;
/// use mea::rwlock::RwLockReadGuard;
///
/// #[derive(Debug, Clone)]
/// struct Foo(String);
///
/// let rwlock = RwLock::new(Foo("hello".to_owned()));
///
/// let guard = rwlock.read().await;
/// let mapped_guard = RwLockReadGuard::map(guard, |f| &f.0);
///
/// assert_eq!(&*mapped_guard, "hello");
/// # }
/// ```
pub fn map<U, F>(orig: Self, f: F) -> MappedRwLockReadGuard<'a, U>
where
F: FnOnce(&T) -> &U,
U: ?Sized,
{
let d = NonNull::from(f(&*orig));
let orig = ManuallyDrop::new(orig);
MappedRwLockReadGuard::new(d, &orig.lock.s)
}
/// Attempts to make a new [`MappedRwLockReadGuard`] for a component of the
/// locked data. The original guard is returned if the closure returns `None`.
///
/// This operation cannot fail as the `RwLockReadGuard` passed in already locked the rwlock.
///
/// This is an associated function that needs to be used as `RwLockReadGuard::filter_map(...)`.
///
/// A method would interfere with methods of the same name on the contents of the locked data.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::rwlock::RwLock;
/// use mea::rwlock::RwLockReadGuard;
///
/// #[derive(Debug, Clone)]
/// struct Foo(String);
///
/// let rwlock = RwLock::new(Foo("hello".to_owned()));
///
/// let guard = rwlock.read().await;
/// let mapped_guard =
/// RwLockReadGuard::filter_map(guard, |f| if f.0.len() > 3 { Some(&f.0) } else { None })
/// .expect("should have mapped");
///
/// assert_eq!(&*mapped_guard, "hello");
/// # }
/// ```
pub fn filter_map<U, F>(orig: Self, f: F) -> Result<MappedRwLockReadGuard<'a, U>, Self>
where
F: FnOnce(&T) -> Option<&U>,
U: ?Sized,
{
match f(&*orig) {
Some(d) => {
let d = NonNull::from(d);
let orig = ManuallyDrop::new(orig);
Ok(MappedRwLockReadGuard::new(d, &orig.lock.s))
}
None => Err(orig),
}
}
}
File diff suppressed because it is too large Load Diff
+286
View File
@@ -0,0 +1,286 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::fmt;
use std::mem::ManuallyDrop;
use std::ops::Deref;
use std::ops::DerefMut;
use std::ptr::NonNull;
use crate::rwlock::MappedRwLockWriteGuard;
use crate::rwlock::RwLock;
use crate::rwlock::RwLockReadGuard;
impl<T: ?Sized> RwLock<T> {
/// Locks this `RwLock` with exclusive write access, causing the current task to yield until the
/// lock has been acquired.
///
/// The calling task will yield while other writers or readers currently have access to the
/// lock.
///
/// Returns an RAII guard which will drop the write access of this `RwLock` when dropped.
///
/// # Cancel safety
///
/// This method uses a queue to fairly distribute locks in the order they were requested.
/// Cancelling a call to `write` makes you lose your place in the queue.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::rwlock::RwLock;
///
/// let lock = RwLock::new(1);
/// let mut n = lock.write().await;
/// *n = 2;
/// # }
/// ```
pub async fn write(&self) -> RwLockWriteGuard<'_, T> {
self.s.acquire(self.max_readers).await;
RwLockWriteGuard {
permits_acquired: self.max_readers,
lock: self,
}
}
/// Attempts to acquire this `RwLock` with exclusive write access.
///
/// If the access couldn't be acquired immediately, returns `None`. Otherwise, an RAII guard is
/// returned which will release write access when dropped.
///
/// # Examples
///
/// ```
/// use std::sync::Arc;
///
/// use mea::rwlock::RwLock;
///
/// let lock = Arc::new(RwLock::new(1));
///
/// let v = lock.try_read().unwrap();
/// assert!(lock.try_write().is_none());
/// drop(v);
///
/// let mut v = lock.try_write().unwrap();
/// *v = 2;
/// ```
pub fn try_write(&self) -> Option<RwLockWriteGuard<'_, T>> {
if self.s.try_acquire(self.max_readers) {
Some(RwLockWriteGuard {
permits_acquired: self.max_readers,
lock: self,
})
} else {
None
}
}
}
/// RAII structure used to release the exclusive write access of a lock when dropped.
///
/// This structure is created by the [`RwLock::write`] method.
///
/// See the [module level documentation](crate::rwlock) for more.
#[must_use = "if unused the RwLock will immediately unlock"]
pub struct RwLockWriteGuard<'a, T: ?Sized> {
pub(super) permits_acquired: usize,
pub(super) lock: &'a RwLock<T>,
}
unsafe impl<T: ?Sized + Send + Sync> Send for RwLockWriteGuard<'_, T> {}
unsafe impl<T: ?Sized + Send + Sync> Sync for RwLockWriteGuard<'_, T> {}
impl<T: ?Sized> Drop for RwLockWriteGuard<'_, T> {
fn drop(&mut self) {
self.lock.s.release(self.permits_acquired);
}
}
impl<T: ?Sized + fmt::Debug> fmt::Debug for RwLockWriteGuard<'_, T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Debug::fmt(&**self, f)
}
}
impl<T: ?Sized + fmt::Display> fmt::Display for RwLockWriteGuard<'_, T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Display::fmt(&**self, f)
}
}
impl<T: ?Sized> Deref for RwLockWriteGuard<'_, T> {
type Target = T;
fn deref(&self) -> &Self::Target {
unsafe { &*self.lock.c.get() }
}
}
impl<T: ?Sized> DerefMut for RwLockWriteGuard<'_, T> {
fn deref_mut(&mut self) -> &mut Self::Target {
unsafe { &mut *self.lock.c.get() }
}
}
impl<'a, T: ?Sized> RwLockWriteGuard<'a, T> {
/// Makes a new [`MappedRwLockWriteGuard`] for a component of the locked data.
///
/// This operation cannot fail as the `RwLockWriteGuard` passed in already locked the rwlock.
///
/// This is an associated function that needs to be used as `RwLockWriteGuard::map(...)`.
///
/// A method would interfere with methods of the same name on the contents of the locked data.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::rwlock::RwLock;
/// use mea::rwlock::RwLockWriteGuard;
///
/// #[derive(Debug)]
/// struct Foo {
/// a: u32,
/// b: String,
/// }
///
/// let rwlock = RwLock::new(Foo {
/// a: 1,
/// b: "hello".to_owned(),
/// });
///
/// let mut guard = rwlock.write().await;
/// let mut mapped_guard = RwLockWriteGuard::map(guard, |foo| &mut foo.a);
///
/// *mapped_guard = 42;
/// assert_eq!(*mapped_guard, 42);
/// # }
/// ```
pub fn map<U, F>(orig: Self, f: F) -> MappedRwLockWriteGuard<'a, U>
where
F: FnOnce(&mut T) -> &mut U,
U: ?Sized,
{
let d = NonNull::from(f(unsafe { &mut *orig.lock.c.get() }));
let permits_acquired = orig.permits_acquired;
let orig = ManuallyDrop::new(orig);
MappedRwLockWriteGuard::new(d, &orig.lock.s, permits_acquired)
}
/// Attempts to make a new [`MappedRwLockWriteGuard`] for a component of the
/// locked data. The original guard is returned if the closure returns `None`.
///
/// This operation cannot fail as the `RwLockWriteGuard` passed in already locked the rwlock.
///
/// This is an associated function that needs to be used as `RwLockWriteGuard::filter_map(...)`.
///
/// A method would interfere with methods of the same name on the contents of the locked data.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use mea::rwlock::RwLock;
/// use mea::rwlock::RwLockWriteGuard;
///
/// #[derive(Debug)]
/// struct Foo {
/// a: u32,
/// b: String,
/// }
///
/// let rwlock = RwLock::new(Foo {
/// a: 11,
/// b: "ok".to_owned(),
/// });
///
/// let mut guard = rwlock.write().await;
/// let mut mapped_guard =
/// RwLockWriteGuard::filter_map(
/// guard,
/// |foo| {
/// if foo.a > 10 { Some(&mut foo.a) } else { None }
/// },
/// )
/// .expect("should have mapped");
///
/// *mapped_guard = 12;
/// assert_eq!(*mapped_guard, 12);
/// # }
/// ```
pub fn filter_map<U, F>(orig: Self, f: F) -> Result<MappedRwLockWriteGuard<'a, U>, Self>
where
F: FnOnce(&mut T) -> Option<&mut U>,
U: ?Sized,
{
match f(unsafe { &mut *orig.lock.c.get() }) {
Some(d) => {
let d = NonNull::from(d);
let permits_acquired = orig.permits_acquired;
let orig = ManuallyDrop::new(orig);
Ok(MappedRwLockWriteGuard::new(
d,
&orig.lock.s,
permits_acquired,
))
}
None => Err(orig),
}
}
/// Atomically downgrades the write lock to a read lock.
///
/// This method changes the lock from exclusive mode to shared mode atomically,
/// preventing other writers from acquiring the lock in between.
///
/// This is more efficient than dropping the write guard and acquiring a new read guard.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::rwlock::RwLock;
///
/// let lock = Arc::new(RwLock::new(1));
///
/// let mut write_guard = lock.write().await;
/// *write_guard = 2;
///
/// let read_guard = write_guard.downgrade();
/// assert_eq!(*read_guard, 2);
///
/// assert!(lock.try_write().is_none());
///
/// drop(read_guard);
/// assert!(lock.try_write().is_some());
/// # }
/// ```
pub fn downgrade(self) -> RwLockReadGuard<'a, T> {
// Prevent the original write guard from running its Drop implementation,
// which would release all permits. This must be done BEFORE any operation
// that might panic to ensure panic safety.
let guard = ManuallyDrop::new(self);
// Release max_readers - 1 permits to convert the write lock to a read lock.
// The remaining 1 permit is kept for the read lock.
guard.lock.s.release(guard.permits_acquired - 1);
RwLockReadGuard { lock: guard.lock }
}
}
+671
View File
@@ -0,0 +1,671 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//! An async counting semaphore for controlling access to a set of resources.
//!
//! A semaphore maintains a set of permits. Permits are used to synchronize access
//! to a pool of resources. Each [`acquire`] call blocks until a permit is available,
//! and then takes one permit. Each [`release`] call adds a new permit, potentially
//! releasing a blocked acquirer.
//!
//! Semaphores are often used to restrict the number of tasks that can access some
//! (physical or logical) resource. For example, here is a class that uses a
//! semaphore to control access to a pool of connections:
//!
//! # Examples
//!
//! ## Basic usage
//!
//! ```
//! # #[tokio::main]
//! # async fn main() {
//! use mea::semaphore::Semaphore;
//!
//! let semaphore = Semaphore::new(3);
//! let a_permit = semaphore.acquire(1).await;
//! let two_permits = semaphore.acquire(2).await;
//!
//! assert_eq!(semaphore.available_permits(), 0);
//!
//! let permit_attempt = semaphore.try_acquire(1);
//! assert!(permit_attempt.is_none());
//! # }
//! ```
//!
//! ## Limit the number of simultaneously opened files in your program
//!
//! Most operating systems have limits on the number of open file
//! handles. Even in systems without explicit limits, resource constraints
//! implicitly set an upper bound on the number of open files. If your
//! program attempts to open a large number of files and exceeds this
//! limit, it will result in an error.
//!
//! This example uses a Semaphore with 100 permits. By acquiring a permit from
//! the Semaphore before accessing a file, you ensure that your program opens
//! no more than 100 files at a time. When trying to open the 101st
//! file, the program will wait until a permit becomes available before
//! proceeding to open another file.
//!
//! ```
//! use std::fs::File;
//! use std::io::Result;
//! use std::io::Write;
//! use std::sync::LazyLock;
//!
//! use mea::semaphore::Semaphore;
//!
//! static PERMITS: LazyLock<Semaphore> = LazyLock::new(|| Semaphore::new(100));
//!
//! async fn write_to_file(message: &[u8]) -> Result<()> {
//! let _permit = PERMITS.acquire(1).await;
//! let mut buffer = File::create("example.txt")?;
//! buffer.write_all(message)?;
//! Ok(()) // Permit goes out of scope here, and is available again for acquisition
//! }
//! ```
//!
//! [`acquire`]: Semaphore::acquire
//! [`release`]: Semaphore::release
use std::sync::Arc;
use crate::internal;
#[cfg(test)]
mod tests;
/// An async counting semaphore for controlling access to a set of resources.
///
/// See the [module level documentation](self) for more.
#[derive(Debug)]
pub struct Semaphore {
s: internal::Semaphore,
}
impl Semaphore {
/// Creates a new semaphore with the given number of permits.
///
/// # Examples
///
/// ```
/// use mea::semaphore::Semaphore;
///
/// let sem = Semaphore::new(5); // Creates a semaphore with 5 permits
/// ```
pub const fn new(permits: usize) -> Self {
Self {
s: internal::Semaphore::new(permits),
}
}
/// Returns the current number of permits available.
///
/// # Examples
///
/// ```
/// use mea::semaphore::Semaphore;
///
/// let sem = Semaphore::new(2);
/// assert_eq!(sem.available_permits(), 2);
///
/// let permit = sem.try_acquire(1).unwrap();
/// assert_eq!(sem.available_permits(), 1);
/// ```
pub fn available_permits(&self) -> usize {
self.s.available_permits()
}
/// Reduces the semaphore's permits by a maximum of `n`.
///
/// Returns the actual number of permits that were reduced. This may be less
/// than `n` if there are insufficient permits available.
///
/// This is useful when you want to permanently remove permits from the semaphore.
///
/// # Examples
///
/// ```
/// use mea::semaphore::Semaphore;
///
/// let sem = Semaphore::new(5);
/// assert_eq!(sem.forget(3), 3); // Removes 3 permits
/// assert_eq!(sem.available_permits(), 2);
///
/// // Trying to forget more permits than available
/// assert_eq!(sem.forget(3), 2); // Only removes remaining 2 permits
/// assert_eq!(sem.available_permits(), 0);
/// ```
pub fn forget(&self, n: usize) -> usize {
self.s.forget(n)
}
/// Reduces the semaphore's permits by exactly `n`.
///
/// If the semaphore has not enough permits, this would enqueue front an empty waiter to
/// consume the permits, which ensures the permits are reduced by exactly `n`.
///
/// This is useful when you want to permanently remove permits from the semaphore.
///
/// # Examples
///
/// ```
/// use mea::semaphore::Semaphore;
///
/// let sem = Semaphore::new(5);
/// sem.forget_exact(3); // Removes 3 permits
/// assert_eq!(sem.available_permits(), 2);
///
/// // Trying to forget more permits than available
/// sem.forget_exact(3); // Only removes remaining 2 permits
/// assert_eq!(sem.available_permits(), 0);
///
/// sem.release(5); // Adds 5 permits
/// assert_eq!(sem.available_permits(), 4); // Only 4 permits are available
/// ```
pub fn forget_exact(&self, n: usize) {
self.s.forget_exact(n);
}
/// Adds `n` new permits to the semaphore.
///
/// # Panics
///
/// Panics if adding the permits would cause the total number of permits to overflow.
///
/// # Examples
///
/// ```
/// use mea::semaphore::Semaphore;
///
/// let sem = Semaphore::new(0);
/// sem.release(2); // Adds 2 permits
/// assert_eq!(sem.available_permits(), 2);
/// ```
pub fn release(&self, permits: usize) {
self.s.release(permits);
}
/// Attempts to acquire `n` permits from the semaphore without blocking.
///
/// If the permits are successfully acquired, a [`SemaphorePermit`] is returned.
/// The permits will be automatically returned to the semaphore when the permit
/// is dropped, unless [`forget`] is called.
///
/// # Examples
///
/// ```
/// use mea::semaphore::Semaphore;
///
/// let sem = Semaphore::new(2);
///
/// // First acquisition succeeds
/// let permit1 = sem.try_acquire(1).unwrap();
/// assert_eq!(sem.available_permits(), 1);
///
/// // Second acquisition succeeds
/// let permit2 = sem.try_acquire(1).unwrap();
/// assert_eq!(sem.available_permits(), 0);
///
/// // Third acquisition fails
/// assert!(sem.try_acquire(1).is_none());
/// ```
///
/// [`forget`]: SemaphorePermit::forget
pub fn try_acquire(&self, permits: usize) -> Option<SemaphorePermit<'_>> {
if self.s.try_acquire(permits) {
Some(SemaphorePermit { sem: self, permits })
} else {
None
}
}
/// Attempts to acquire `n` permits from the semaphore without blocking.
///
/// This method performs as a combinator of [`Semaphore::try_acquire`] and
/// [`Semaphore::forget`].
pub fn try_acquire_and_forget(&self, permits: usize) -> bool {
self.s.try_acquire(permits)
}
/// Acquires `n` permits from the semaphore.
///
/// If the permits are not immediately available, this method will wait until they become
/// available. Returns a [`SemaphorePermit`] that will release the permits when dropped.
///
/// # Cancel safety
///
/// This method uses a queue to fairly distribute permits in the order they were requested.
/// Cancelling a call to `acquire` makes you lose your place in the queue.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::semaphore::Semaphore;
///
/// let sem = Arc::new(Semaphore::new(2));
/// let sem2 = sem.clone();
///
/// let handle = tokio::spawn(async move {
/// let permit = sem2.acquire(1).await;
/// // Do some work with the permit.
/// // Permit is automatically released when dropped.
/// });
///
/// let permit = sem.acquire(1).await;
/// // Do some work with the permit
/// drop(permit); // Explicitly release the permit
///
/// handle.await.unwrap();
/// # }
/// ```
pub async fn acquire(&self, permits: usize) -> SemaphorePermit<'_> {
self.s.acquire(permits).await;
SemaphorePermit { sem: self, permits }
}
/// Acquires `n` permits from the semaphore.
///
/// This method performs as a combinator of [`Semaphore::acquire`] and
/// [`Semaphore::forget`].
pub async fn acquire_and_forget(&self, permits: usize) {
self.s.acquire(permits).await;
}
/// Attempts to acquire `n` permits from the semaphore without blocking.
///
/// The semaphore must be wrapped in an [`Arc`] to call this method.
///
/// If the permits are successfully acquired, a [`OwnedSemaphorePermit`] is returned.
/// The permits will be automatically returned to the semaphore when the permit
/// is dropped, unless [`forget`] is called.
///
/// # Examples
///
/// ```
/// use std::sync::Arc;
///
/// use mea::semaphore::Semaphore;
///
/// let sem = Arc::new(Semaphore::new(2));
///
/// let p1 = sem.clone().try_acquire_owned(1).unwrap();
/// assert_eq!(sem.available_permits(), 1);
///
/// let p2 = sem.clone().try_acquire_owned(1).unwrap();
/// assert_eq!(sem.available_permits(), 0);
///
/// let p3 = sem.try_acquire_owned(1);
/// assert!(p3.is_none());
/// ```
///
/// [`forget`]: SemaphorePermit::forget
pub fn try_acquire_owned(self: Arc<Self>, permits: usize) -> Option<OwnedSemaphorePermit> {
if self.s.try_acquire(permits) {
Some(OwnedSemaphorePermit { sem: self, permits })
} else {
None
}
}
/// Attempts to acquire `n` permits from the semaphore without blocking.
///
/// The semaphore must be wrapped in an [`Arc`] to call this method.
///
/// This method performs as a combinator of [`Semaphore::try_acquire_owned`] and
/// [`Semaphore::forget`].
pub fn try_acquire_owned_and_forget(self: Arc<Self>, permits: usize) -> bool {
self.s.try_acquire(permits)
}
/// Acquires `n` permits from the semaphore.
///
/// The semaphore must be wrapped in an [`Arc`] to call this method.
///
/// If the permits are not immediately available, this method will wait until they become
/// available. Returns a [`OwnedSemaphorePermit`] that will release the permits when dropped.
///
/// # Cancel safety
///
/// This method uses a queue to fairly distribute permits in the order they were requested.
/// Cancelling a call to `acquire_owned` makes you lose your place in the queue.
///
/// # Examples
///
/// ```
/// # #[tokio::main]
/// # async fn main() {
/// use std::sync::Arc;
///
/// use mea::semaphore::Semaphore;
///
/// let sem = Arc::new(Semaphore::new(3));
/// let mut join_handles = Vec::new();
///
/// for _ in 0..5 {
/// let permit = sem.clone().acquire_owned(1).await;
/// join_handles.push(tokio::spawn(async move {
/// // perform task...
/// // explicitly own `permit` in the task
/// drop(permit);
/// }));
/// }
///
/// for handle in join_handles {
/// handle.await.unwrap();
/// }
/// # }
/// ```
pub async fn acquire_owned(self: Arc<Self>, permits: usize) -> OwnedSemaphorePermit {
self.s.acquire(permits).await;
OwnedSemaphorePermit { sem: self, permits }
}
/// Acquires `n` permits from the semaphore.
///
/// The semaphore must be wrapped in an [`Arc`] to call this method.
///
/// This method performs as a combinator of [`Semaphore::acquire_owned`] and
/// [`Semaphore::forget`].
pub async fn acquire_owned_and_forget(self: Arc<Self>, permits: usize) {
self.s.acquire(permits).await;
}
}
/// A permit from the semaphore.
///
/// This type is created by the [`acquire`] and [`try_acquire`] methods on [`Semaphore`].
/// When the permit is dropped, the permits will be returned to the semaphore unless
/// [`forget`] is called.
///
/// [`acquire`]: Semaphore::acquire
/// [`try_acquire`]: Semaphore::try_acquire
/// [`forget`]: SemaphorePermit::forget
#[must_use = "permits are released immediately when dropped"]
#[derive(Debug)]
pub struct SemaphorePermit<'a> {
sem: &'a Semaphore,
permits: usize,
}
impl SemaphorePermit<'_> {
/// Forgets the permit **without** releasing it back to the semaphore.
///
/// This can be used to permanently reduce the number of permits available
/// from a semaphore.
///
/// # Examples
///
/// ```
/// use std::sync::Arc;
///
/// use mea::semaphore::Semaphore;
///
/// let sem = Arc::new(Semaphore::new(10));
/// {
/// let permit = sem.try_acquire(5).unwrap();
/// assert_eq!(sem.available_permits(), 5);
/// permit.forget();
/// }
///
/// // Since we forgot the permit, available permits won't go back to
/// // its initial value even after the permit is dropped
/// assert_eq!(sem.available_permits(), 5);
/// ```
pub fn forget(mut self) {
self.permits = 0;
}
/// Merge two [`SemaphorePermit`] instances together, consuming `other`
/// without releasing the permits it holds.
///
/// Permits held by both `self` and `other` are released when `self` drops.
///
/// # Panics
///
/// This function panics if permits from different [`Semaphore`] instances
/// are merged.
///
/// # Examples
///
/// ```
/// use std::sync::Arc;
///
/// use mea::semaphore::Semaphore;
///
/// let sem = Arc::new(Semaphore::new(10));
/// let mut permit = sem.try_acquire(1).unwrap();
///
/// for _ in 0..9 {
/// let new_permit = sem.try_acquire(1).unwrap();
/// // Merge individual permits into a single one.
/// permit.merge(new_permit)
/// }
///
/// assert_eq!(sem.available_permits(), 0);
///
/// // Release all permits in a single batch.
/// drop(permit);
///
/// assert_eq!(sem.available_permits(), 10);
/// ```
#[track_caller]
pub fn merge(&mut self, mut other: Self) {
assert!(
std::ptr::eq(self.sem, other.sem),
"merging permits from different semaphore instances"
);
self.permits += other.permits;
other.permits = 0;
}
/// Splits `n` permits from `self` and returns a new [`SemaphorePermit`] instance that holds `n`
/// permits.
///
/// If there are insufficient permits, and it is impossible to reduce by `n`, returns `None`.
///
/// # Examples
///
/// ```
/// use std::sync::Arc;
///
/// use mea::semaphore::Semaphore;
///
/// let sem = Arc::new(Semaphore::new(3));
///
/// let mut p1 = sem.try_acquire(3).unwrap();
/// let p2 = p1.split(1).unwrap();
///
/// assert_eq!(p1.permits(), 2);
/// assert_eq!(p2.permits(), 1);
/// ```
pub fn split(&mut self, n: usize) -> Option<Self> {
if n > self.permits {
return None;
}
self.permits -= n;
Some(Self {
sem: self.sem,
permits: n,
})
}
/// Returns the number of permits this permit holds.
///
/// # Examples
///
/// ```
/// use mea::semaphore::Semaphore;
///
/// let sem = Semaphore::new(5);
/// let permit = sem.try_acquire(3).unwrap();
/// assert_eq!(permit.permits(), 3);
/// ```
pub fn permits(&self) -> usize {
self.permits
}
}
impl Drop for SemaphorePermit<'_> {
fn drop(&mut self) {
self.sem.release(self.permits);
}
}
/// An owned permit from the semaphore.
///
/// This type is created by the [`acquire_owned`] method.
///
/// [`acquire_owned`]: Semaphore::acquire_owned
#[must_use = "permits are released immediately when dropped"]
#[derive(Debug)]
pub struct OwnedSemaphorePermit {
sem: Arc<Semaphore>,
permits: usize,
}
impl OwnedSemaphorePermit {
/// Forgets the permit **without** releasing it back to the semaphore.
///
/// This can be used to permanently reduce the number of permits available
/// from a semaphore.
///
/// # Examples
///
/// ```
/// use std::sync::Arc;
///
/// use mea::semaphore::Semaphore;
///
/// let sem = Arc::new(Semaphore::new(10));
/// {
/// let permit = sem.try_acquire(5).unwrap();
/// assert_eq!(sem.available_permits(), 5);
/// permit.forget();
/// }
///
/// // Since we forgot the permit, available permits won't go back to
/// // its initial value even after the permit is dropped
/// assert_eq!(sem.available_permits(), 5);
/// ```
pub fn forget(mut self) {
self.permits = 0;
}
/// Merge two [`SemaphorePermit`] instances together, consuming `other`
/// without releasing the permits it holds.
///
/// Permits held by both `self` and `other` are released when `self` drops.
///
/// # Panics
///
/// This function panics if permits from different [`Semaphore`] instances
/// are merged.
///
/// # Examples
///
/// ```
/// use std::sync::Arc;
///
/// use mea::semaphore::Semaphore;
///
/// let sem = Arc::new(Semaphore::new(10));
/// let mut permit = sem.try_acquire(1).unwrap();
///
/// for _ in 0..9 {
/// let new_permit = sem.try_acquire(1).unwrap();
/// // Merge individual permits into a single one.
/// permit.merge(new_permit)
/// }
///
/// assert_eq!(sem.available_permits(), 0);
///
/// // Release all permits in a single batch.
/// drop(permit);
///
/// assert_eq!(sem.available_permits(), 10);
/// ```
#[track_caller]
pub fn merge(&mut self, mut other: Self) {
assert!(
Arc::ptr_eq(&self.sem, &other.sem),
"merging permits from different semaphore instances"
);
self.permits += other.permits;
other.permits = 0;
}
/// Splits `n` permits from `self` and returns a new [`OwnedSemaphorePermit`] instance that
/// holds `n` permits.
///
/// If there are insufficient permits, and it is impossible to reduce by `n`, returns `None`.
///
/// # Note
///
/// It will clone the owned `Arc<Semaphore>` to construct the new instance.
///
/// # Examples
///
/// ```
/// use std::sync::Arc;
///
/// use mea::semaphore::Semaphore;
///
/// let sem = Arc::new(Semaphore::new(3));
///
/// let mut p1 = sem.try_acquire_owned(3).unwrap();
/// let p2 = p1.split(1).unwrap();
///
/// assert_eq!(p1.permits(), 2);
/// assert_eq!(p2.permits(), 1);
/// ```
pub fn split(&mut self, n: usize) -> Option<Self> {
if n > self.permits {
return None;
}
self.permits -= n;
Some(Self {
sem: self.sem.clone(),
permits: n,
})
}
/// Returns the number of permits this permit holds.
///
/// # Examples
///
/// ```
/// use mea::semaphore::Semaphore;
///
/// let sem = Semaphore::new(5);
/// let permit = sem.try_acquire(3).unwrap();
/// assert_eq!(permit.permits(), 3);
/// ```
pub fn permits(&self) -> usize {
self.permits
}
}
impl Drop for OwnedSemaphorePermit {
fn drop(&mut self) {
self.sem.release(self.permits);
}
}
+204
View File
@@ -0,0 +1,204 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::future::Future;
use std::pin::pin;
use std::sync::Arc;
use std::task::Context;
use std::task::Waker;
use std::vec::Vec;
use super::*;
use crate::latch::Latch;
#[test]
fn no_permits() {
// this should not panic
Semaphore::new(0);
}
#[test]
fn try_acquire() {
let sem = Semaphore::new(1);
{
let p1 = sem.try_acquire(1);
assert!(p1.is_some());
let p2 = sem.try_acquire(1);
assert!(p2.is_none());
}
let p3 = sem.try_acquire(1);
assert!(p3.is_some());
}
#[tokio::test]
async fn acquire() {
let sem = Arc::new(Semaphore::new(1));
let p1 = sem.try_acquire(1).unwrap();
let sem_clone = sem.clone();
let j = tokio::spawn(async move {
let _p2 = sem_clone.acquire(1).await;
});
drop(p1);
j.await.unwrap();
}
#[tokio::test]
async fn add_permits() {
let sem = Arc::new(Semaphore::new(0));
let sem_clone = sem.clone();
let j = tokio::spawn(async move {
let _p2 = sem_clone.acquire(1).await;
});
sem.release(1);
j.await.unwrap();
}
#[test]
fn forget() {
let sem = Arc::new(Semaphore::new(1));
{
let p = sem.try_acquire(1).unwrap();
assert_eq!(sem.available_permits(), 0);
p.forget();
assert_eq!(sem.available_permits(), 0);
}
assert_eq!(sem.available_permits(), 0);
assert!(sem.try_acquire(1).is_none());
}
#[tokio::test]
async fn stress_test() {
let sem = Arc::new(Semaphore::new(5));
let mut join_handles = Vec::new();
for i in 0..100 {
let sem_clone = sem.clone();
join_handles.push(tokio::spawn(async move {
let _p = sem_clone.acquire(1).await;
tokio::time::sleep(std::time::Duration::from_millis(100 - i)).await;
}));
}
for j in join_handles {
j.await.unwrap();
}
// there should be exactly 5 semaphores available now
let _p1 = sem.try_acquire(1).unwrap();
let _p2 = sem.try_acquire(1).unwrap();
let _p3 = sem.try_acquire(1).unwrap();
let _p4 = sem.try_acquire(1).unwrap();
let _p5 = sem.try_acquire(1).unwrap();
assert!(sem.try_acquire(1).is_none());
}
#[test]
fn add_max_amount_permits() {
let s = Semaphore::new(0);
s.release(usize::MAX);
assert_eq!(s.available_permits(), usize::MAX);
}
#[test]
#[should_panic]
fn add_more_than_max_amount_permits1() {
let s = Semaphore::new(1);
s.release(usize::MAX);
}
#[test]
#[should_panic]
fn add_more_than_max_amount_permits2() {
let s = Semaphore::new(usize::MAX - 1);
s.release(1);
s.release(1);
}
#[test]
fn no_panic_at_max_permits() {
let _ = Semaphore::new(usize::MAX);
let s = Semaphore::new(usize::MAX - 1);
s.release(1);
}
#[test]
fn try_acquire_concurrently() {
let s = Semaphore::new(1);
let p1 = s.try_acquire(1).unwrap();
assert_eq!(s.available_permits(), 0);
let p2 = s.try_acquire(1);
assert!(p2.is_none());
assert_eq!(s.available_permits(), 0);
drop(p1);
assert_eq!(s.available_permits(), 1);
}
#[test]
fn acquire_then_drop() {
let waker = Waker::noop();
let mut context = Context::from_waker(waker);
let s = Semaphore::new(1);
let p1 = s.try_acquire(1).unwrap();
{
let p2 = s.acquire(1);
let poll = pin!(p2).poll(&mut context);
assert!(poll.is_pending());
}
drop(p1);
assert_eq!(s.available_permits(), 1);
}
#[test]
fn wake_then_drop() {
let waker = Waker::noop();
let mut context = Context::from_waker(waker);
let s = Semaphore::new(2);
let p1 = s.try_acquire(2).unwrap();
{
let p2 = s.acquire(1);
let p2 = pin!(p2);
assert!(p2.poll(&mut context).is_pending());
{
let p3 = s.acquire(1);
let p3 = pin!(p3);
assert!(p3.poll(&mut context).is_pending());
drop(p1);
}
}
assert_eq!(s.available_permits(), 2);
}
#[tokio::test]
async fn acquire_then_forget_exact() {
let s = Arc::new(Semaphore::new(5));
s.forget_exact(3);
assert_eq!(s.available_permits(), 2);
let acquired = Arc::new(Latch::new(1));
let acquired_clone = acquired.clone();
let s_clone = s.clone();
tokio::spawn(async move {
let _p = s_clone.acquire(3).await;
acquired_clone.count_down();
});
assert!(acquired.try_wait().is_err());
s.forget_exact(2);
s.release(2);
assert!(acquired.try_wait().is_err());
s.release(1);
acquired.wait().await;
assert_eq!(s.available_permits(), 3);
}
+164
View File
@@ -0,0 +1,164 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//! A composite synchronization primitive for managing shutdown signals.
//!
//! This module provides [`new_pair`] to create a pair of handles for managing shutdown signals:
//!
//! * [`ShutdownSend`] can send a shutdown signal, and can wait for all the tasks to finish.
//! * [`ShutdownRecv`] can wait for the shutdown signal, and should be dropped when the task is
//! done, which will notify the sender on all the tasks finished.
//! * [`ShutdownWatch`] can wait for the shutdown signal without blocking
//! [`ShutdownSend::await_shutdown`].
//!
//! Internally, the shutdown signal is implemented using a countdown latch, and the task completion
//! is tracked using a wait group. [`ShutdownSend`] is cloneable, allowing multiple sources to send
//! the shutdown signal; [`ShutdownRecv`] is also cloneable, allowing multiple tasks to wait for the
//! same shutdown signal.
//!
//! [`ShutdownSend::await_shutdown`] would block until all the tasks are done, i.e., all the
//! [`ShutdownRecv`]s dropped.
//!
//! # Examples
//!
//! ```
//! # #[tokio::main]
//! # async fn main() {
//! let (tx, rx) = mea::shutdown::new_pair();
//!
//! for i in 0..3 {
//! let rx = rx.clone();
//! tokio::spawn(async move {
//! println!("Task {} starting", i);
//! rx.is_shutdown().await;
//! println!("Task {} done", i);
//! });
//! }
//! drop(rx);
//!
//! tx.shutdown();
//! tx.await_shutdown().await;
//! # }
//! ```
use std::future::Future;
use std::future::IntoFuture;
use std::sync::Arc;
use crate::latch::Latch;
use crate::waitgroup::Wait;
use crate::waitgroup::WaitGroup;
#[cfg(test)]
mod tests;
/// Create a pair of handles for managing shutdown signals.
///
/// See the [module level documentation](self) for more.
pub fn new_pair() -> (ShutdownSend, ShutdownRecv) {
let latch = Arc::new(Latch::new(1));
let wg = WaitGroup::new();
let send = ShutdownSend {
latch: latch.clone(),
wait: wg.clone().into_future(),
};
let recv = ShutdownRecv { latch, wg };
(send, recv)
}
/// A handle for sending shutdown signals.
///
/// See the [module level documentation](self) for more.
#[derive(Debug, Clone)]
pub struct ShutdownSend {
latch: Arc<Latch>,
wait: Wait,
}
impl ShutdownSend {
/// Send a shutdown signal to all [`ShutdownRecv`] handles.
pub fn shutdown(&self) {
self.latch.count_down();
}
/// Wait for all [`ShutdownRecv`] handles to be dropped.
pub async fn await_shutdown(self) {
self.wait.await;
}
}
/// A handle for receiving shutdown signals.
///
/// See the [module level documentation](self) for more.
#[derive(Debug, Clone)]
pub struct ShutdownRecv {
latch: Arc<Latch>,
#[allow(dead_code)] // hold the wait group
wg: WaitGroup,
}
impl ShutdownRecv {
/// Returns a handle for watching the shutdown signal.
///
/// The returned handle does not block [`ShutdownSend::await_shutdown`].
pub fn watch(&self) -> ShutdownWatch {
ShutdownWatch {
latch: self.latch.clone(),
}
}
/// Returns whether the shutdown signal has been received.
pub fn is_shutdown_now(&self) -> bool {
self.latch.try_wait().is_ok()
}
/// Returns a future that resolves when the shutdown signal is received.
pub async fn is_shutdown(&self) {
self.latch.wait().await;
}
/// Returns an owned future that resolves when the shutdown signal is received.
///
/// The returned future has no lifetime constraints.
pub fn is_shutdown_owned(&self) -> impl Future<Output = ()> + 'static {
self.latch.clone().wait_owned()
}
}
/// A handle for watching shutdown signals without participating in shutdown completion.
///
/// See the [module level documentation](self) for more.
#[derive(Debug, Clone)]
pub struct ShutdownWatch {
latch: Arc<Latch>,
}
impl ShutdownWatch {
/// Returns whether the shutdown signal has been received.
pub fn is_shutdown_now(&self) -> bool {
self.latch.try_wait().is_ok()
}
/// Returns a future that resolves when the shutdown signal is received.
pub async fn is_shutdown(&self) {
self.latch.wait().await;
}
/// Returns an owned future that resolves when the shutdown signal is received.
///
/// The returned future has no lifetime constraints.
pub fn is_shutdown_owned(&self) -> impl Future<Output = ()> + 'static {
self.latch.clone().wait_owned()
}
}
+107
View File
@@ -0,0 +1,107 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::future::pending;
use super::*;
use crate::test_runtime;
#[test]
fn test_single_pair() {
let (tx, rx) = new_pair();
let handle = test_runtime().spawn(async move { rx.is_shutdown().await });
tx.shutdown();
pollster::block_on(tx.await_shutdown());
pollster::block_on(handle).unwrap();
}
#[test]
fn test_multiple_tasks() {
let (tx, rx) = new_pair();
for _i in 0..100 {
let rx = rx.clone();
test_runtime().spawn(async move { rx.is_shutdown().await });
}
drop(rx);
tx.shutdown();
pollster::block_on(tx.await_shutdown());
}
#[test]
fn test_multiple_senders() {
let (tx, rx) = new_pair();
for _i in 0..100 {
let rx = rx.clone();
test_runtime().spawn(async move { rx.is_shutdown().await });
}
drop(rx);
let tx_clone = tx.clone();
tx.shutdown();
pollster::block_on(tx.await_shutdown());
pollster::block_on(tx_clone.await_shutdown());
}
#[test]
fn test_is_shutdown_now() {
let (tx, rx) = new_pair();
assert!(!rx.is_shutdown_now());
tx.shutdown();
assert!(rx.is_shutdown_now());
}
#[test]
fn test_is_shutdown_owned_not_capture_self() {
struct State {
rx: ShutdownRecv,
}
async fn run_state(_state: &mut State) {
pending::<()>().await;
}
let (tx, rx) = new_pair();
let mut state = State { rx };
test_runtime().spawn(async move {
let is_shutdown = state.rx.is_shutdown_owned();
tokio::select! {
_ = is_shutdown => (),
_ = run_state(&mut state) => (),
}
});
tx.shutdown();
pollster::block_on(tx.await_shutdown());
}
#[test]
fn test_watch_does_not_block_shutdown() {
let (tx, rx) = new_pair();
let watch = rx.watch();
drop(rx);
tx.shutdown();
assert!(watch.is_shutdown_now());
pollster::block_on(tx.await_shutdown());
}
#[test]
fn test_watch_is_shutdown() {
let (tx, rx) = new_pair();
let watch = rx.watch();
let handle = test_runtime().spawn(async move { watch.is_shutdown().await });
drop(rx);
tx.shutdown();
pollster::block_on(tx.await_shutdown());
pollster::block_on(handle).unwrap();
}
+251
View File
@@ -0,0 +1,251 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//! Singleflight provides a duplicate function call suppression mechanism.
use std::borrow::Borrow;
use std::collections::HashMap;
use std::hash::BuildHasher;
use std::hash::Hash;
use std::hash::RandomState;
use std::sync::Arc;
use crate::internal::Mutex;
use crate::once::OnceCell;
#[cfg(test)]
mod tests;
/// Group represents a class of work and forms a namespace in which
/// units of work can be executed with duplicate suppression.
#[derive(Debug)]
pub struct Group<K, V, S = RandomState> {
map: Mutex<HashMap<K, Arc<OnceCell<V>>, S>>,
}
impl<K, V, S> Default for Group<K, V, S>
where
K: Eq + Hash + Clone,
V: Clone,
S: BuildHasher + Clone + Default,
{
fn default() -> Self {
Self::with_hasher(S::default())
}
}
impl<K, V> Group<K, V, RandomState>
where
K: Eq + Hash + Clone,
V: Clone,
{
/// Creates a new Group with the default hasher.
pub fn new() -> Self {
Self {
map: Mutex::new(HashMap::new()),
}
}
}
impl<K, V, S> Group<K, V, S>
where
K: Eq + Hash + Clone,
V: Clone,
S: BuildHasher + Clone,
{
/// Creates a new Group with the given hasher.
pub fn with_hasher(hasher: S) -> Self {
Self {
map: Mutex::new(HashMap::with_hasher(hasher)),
}
}
/// Executes and returns the results of the given function, making sure that only one execution
/// is in-flight for a given key at a time.
///
/// If a duplicate comes in, the duplicate caller waits for the original to complete and
/// receives the same results.
///
/// Once the function completes, the key, if not [`forgotten`], is removed from the group,
/// allowing future calls with the same key to execute the function again.
///
/// [`forgotten`]: Self::forget
///
/// # Examples
///
/// ```
/// use std::sync::Arc;
/// use std::sync::atomic::AtomicUsize;
/// use std::sync::atomic::Ordering;
/// use std::time::Duration;
///
/// use mea::singleflight::Group;
///
/// # #[tokio::main]
/// # async fn main() {
/// let group = Group::new();
/// let counter = Arc::new(AtomicUsize::new(0));
///
/// let c1 = counter.clone();
/// let fut1 = group.work("key", || async move {
/// c1.fetch_add(1, Ordering::SeqCst);
/// // simulate heavy work to avoid immediate completion
/// tokio::time::sleep(Duration::from_millis(100)).await;
/// "result"
/// });
///
/// let c2 = counter.clone();
/// let fut2 = group.work("key", || async move {
/// c2.fetch_add(1, Ordering::SeqCst);
/// // simulate heavy work to avoid immediate completion
/// tokio::time::sleep(Duration::from_millis(100)).await;
/// "result"
/// });
///
/// let (r1, r2) = tokio::join!(fut1, fut2);
///
/// assert_eq!(r1, "result");
/// assert_eq!(r2, "result");
/// assert_eq!(counter.load(Ordering::SeqCst), 1);
/// # }
/// ```
pub async fn work<F>(&self, key: K, func: F) -> V
where
F: AsyncFnOnce() -> V,
{
// 1. Get or create the OnceCell.
let cell = {
let mut map = self.map.lock();
map.entry(key.clone())
.or_insert_with(|| Arc::new(OnceCell::new()))
.clone()
};
// 2. Try to initialize the cell.
// OnceCell::get_or_init guarantees that only one task executes the closure.
let res = cell
.get_or_init(async || {
let result = func().await;
// Cleanup: remove the key from the map.
// We must ensure we remove the entry corresponding to *this* cell.
let mut map = self.map.lock();
if let Some(existing) = map.get(&key) {
// Check if the map still points to our cell.
if Arc::ptr_eq(&cell, existing) {
map.remove(&key);
}
}
result
})
.await;
res.clone()
}
/// Executes and returns the results of the given function, making sure that only one execution
/// is in-flight for a given key at a time.
///
/// If a duplicate comes in, the duplicate caller waits for the original to complete and
/// receives the same results.
///
/// If the computation fails, the error is returned for the caller. Other tasks waiting for the
/// result will retry the computation.
///
/// Once the function completes successfully, the key, if not [`forgotten`], is removed from
/// the group, allowing future calls with the same key to execute the function again.
///
/// [`forgotten`]: Self::forget
///
/// # Examples
///
/// ```
/// use std::sync::Arc;
/// use std::sync::atomic::AtomicUsize;
/// use std::sync::atomic::Ordering;
/// use std::time::Duration;
///
/// use mea::singleflight::Group;
///
/// # #[tokio::main]
/// # async fn main() {
/// let group = Group::new();
///
/// let fut1 = group.try_work("key", || async move {
/// // simulate heavy work to avoid immediate completion
/// tokio::time::sleep(Duration::from_millis(100)).await;
/// Err::<_, &'static str>("fut1")
/// });
///
/// let fut2 = group.try_work("key", || async move {
/// // simulate heavy work to avoid immediate completion
/// tokio::time::sleep(Duration::from_millis(200)).await;
/// Ok::<_, &'static str>("fut2")
/// });
///
/// let (r1, r2) = tokio::join!(fut1, fut2);
///
/// assert_eq!(r1, Err("fut1"));
/// assert_eq!(r2, Ok("fut2"));
/// # }
/// ```
pub async fn try_work<E, F>(&self, key: K, func: F) -> Result<V, E>
where
F: AsyncFnOnce() -> Result<V, E>,
{
// 1. Get or create the OnceCell.
let cell = {
let mut map = self.map.lock();
map.entry(key.clone())
.or_insert_with(|| Arc::new(OnceCell::new()))
.clone()
};
// 2. Try to initialize the cell.
// OnceCell::get_or_try_init guarantees that only one task executes the closure.
let res = cell
.get_or_try_init(async || {
let result = func().await?;
// Cleanup: remove the key from the map.
// We must ensure we remove the entry corresponding to *this* cell.
let mut map = self.map.lock();
if let Some(existing) = map.get(&key) {
// Check if the map still points to our cell.
if Arc::ptr_eq(&cell, existing) {
map.remove(&key);
}
}
Ok(result)
})
.await?;
Ok(res.clone())
}
/// Forgets about the given key.
///
/// Future calls to `work` for this key will call the function rather than waiting for an
/// earlier call to complete. Existing calls to `work` for this key are not affected.
pub fn forget<Q>(&self, key: &Q)
where
K: Borrow<Q>,
Q: Hash + Eq + ?Sized,
{
let mut map = self.map.lock();
map.remove(key);
}
}
+234
View File
@@ -0,0 +1,234 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::sync::Arc;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::time::Duration;
use crate::singleflight::Group;
#[tokio::test]
async fn test_simple() {
let group = Group::new();
let res = group.work("key", || async { "val" }).await;
assert_eq!(res, "val");
}
#[tokio::test]
async fn test_coalescing() {
let group = Arc::new(Group::new());
let counter = Arc::new(AtomicUsize::new(0));
let mut handles = Vec::new();
for _ in 0..10 {
let group = group.clone();
let counter = counter.clone();
handles.push(tokio::spawn(async move {
group
.work("key", || async move {
tokio::time::sleep(Duration::from_millis(100)).await;
counter.fetch_add(1, Ordering::SeqCst);
"val"
})
.await
}));
}
for handle in handles {
assert_eq!(handle.await.unwrap(), "val");
}
assert_eq!(counter.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn test_multiple_keys() {
let group = Arc::new(Group::new());
let counter = Arc::new(AtomicUsize::new(0));
let g1 = group.clone();
let c1 = counter.clone();
let h1 = tokio::spawn(async move {
g1.work("key1", || async move {
tokio::time::sleep(Duration::from_millis(50)).await;
c1.fetch_add(1, Ordering::SeqCst);
"val1"
})
.await
});
let g2 = group.clone();
let c2 = counter.clone();
let h2 = tokio::spawn(async move {
g2.work("key2", || async move {
tokio::time::sleep(Duration::from_millis(50)).await;
c2.fetch_add(1, Ordering::SeqCst);
"val2"
})
.await
});
assert_eq!(h1.await.unwrap(), "val1");
assert_eq!(h2.await.unwrap(), "val2");
assert_eq!(counter.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn test_forget() {
let group = Arc::new(Group::new());
let counter = Arc::new(AtomicUsize::new(0));
let g1 = group.clone();
let c1 = counter.clone();
let h1 = tokio::spawn(async move {
g1.work("key", || async move {
tokio::time::sleep(Duration::from_millis(100)).await;
c1.fetch_add(1, Ordering::SeqCst);
"val1"
})
.await
});
// Wait a bit to ensure the first call is established
tokio::time::sleep(Duration::from_millis(10)).await;
group.forget(&"key");
let g2 = group.clone();
let c2 = counter.clone();
let h2 = tokio::spawn(async move {
g2.work("key", || async move {
tokio::time::sleep(Duration::from_millis(100)).await;
c2.fetch_add(1, Ordering::SeqCst);
"val2"
})
.await
});
assert_eq!(h1.await.unwrap(), "val1");
assert_eq!(h2.await.unwrap(), "val2");
assert_eq!(counter.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn test_panic_safe() {
let group = Arc::new(Group::<&str, String>::new());
// Task that panics
let g1 = group.clone();
let h1 = tokio::spawn(async move {
g1.work("key", || async {
panic!("oops");
})
.await
});
// Wait for h1 to panic and exit
let err = h1.await.unwrap_err();
assert!(err.is_panic());
// Next task should succeed (new attempt)
let res = group.work("key", || async { "success".to_string() }).await;
assert_eq!(res, "success");
}
#[tokio::test]
async fn test_try_work_simple() {
let group = Group::new();
let res = group
.try_work("key", || async { Ok::<&str, ()>("val") })
.await;
assert_eq!(res, Ok("val"));
// Should be removed from map, so next call executes again
let res2 = group
.try_work("key", || async { Ok::<&str, ()>("val2") })
.await;
assert_eq!(res2, Ok("val2"));
}
#[tokio::test]
async fn test_try_work_coalescing() {
let group = Arc::new(Group::new());
let counter = Arc::new(AtomicUsize::new(0));
let mut handles = Vec::new();
for _ in 0..10 {
let group = group.clone();
let counter = counter.clone();
handles.push(tokio::spawn(async move {
group
.try_work("key", || async move {
tokio::time::sleep(Duration::from_millis(100)).await;
counter.fetch_add(1, Ordering::SeqCst);
Ok::<&str, ()>("val")
})
.await
}));
}
for handle in handles {
assert_eq!(handle.await.unwrap(), Ok("val"));
}
assert_eq!(counter.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn test_try_work_failure() {
let group = Group::new();
let res = group
.try_work("key", || async { Err::<&str, &str>("error") })
.await;
assert_eq!(res, Err("error"));
// Retry should work
let res2 = group
.try_work("key", || async { Ok::<&str, ()>("success") })
.await;
assert_eq!(res2, Ok("success"));
}
#[tokio::test]
async fn test_try_work_wait_and_retry() {
let group = Arc::new(Group::new());
let counter = Arc::new(AtomicUsize::new(0));
let g1 = group.clone();
let c1 = counter.clone();
let h1 = tokio::spawn(async move {
g1.try_work("key", || async move {
c1.fetch_add(1, Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(100)).await;
Err::<&str, &str>("fail")
})
.await
});
let g2 = group.clone();
let c2 = counter.clone();
let h2 = tokio::spawn(async move {
// Ensure h1 starts first
tokio::time::sleep(Duration::from_millis(10)).await;
g2.try_work("key", || async move {
c2.fetch_add(1, Ordering::SeqCst);
Ok::<&str, ()>("success")
})
.await
});
assert_eq!(h1.await.unwrap(), Err("fail"));
assert_eq!(h2.await.unwrap(), Ok("success"));
assert_eq!(counter.load(Ordering::SeqCst), 2);
}
+187
View File
@@ -0,0 +1,187 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//! A synchronization primitive for waiting on multiple tasks to complete.
//!
//! Similar to Go's WaitGroup, this type allows a task to wait for multiple other
//! tasks to finish. Each task holds a handle to the WaitGroup, and the main task
//! can wait for all handles to be dropped before proceeding.
//!
//! A WaitGroup waits for a collection of tasks to finish. The main task calls
//! [`clone()`] to create a new worker handle for each task, and can then wait
//! for all tasks to complete by calling `.await` on the WaitGroup.
//!
//! # Examples
//!
//! ```
//! # #[tokio::main]
//! # async fn main() {
//! use std::time::Duration;
//!
//! use mea::waitgroup::WaitGroup;
//! let wg = WaitGroup::new();
//!
//! for i in 0..3 {
//! let wg = wg.clone();
//! tokio::spawn(async move {
//! println!("Task {} starting", i);
//! tokio::time::sleep(Duration::from_millis(100)).await;
//! // wg is automatically decremented when dropped
//! drop(wg);
//! });
//! }
//!
//! // Wait for all tasks to complete
//! wg.await;
//! println!("All tasks completed");
//! # }
//! ```
//!
//! [`clone()`]: WaitGroup::clone
use std::fmt;
use std::future::Future;
use std::future::IntoFuture;
use std::pin::Pin;
use std::sync::Arc;
use std::task::Context;
use std::task::Poll;
use crate::internal::CountdownState;
#[cfg(test)]
mod tests;
/// A synchronization primitive for waiting on multiple tasks to complete.
///
/// See the [module level documentation](self) for more.
pub struct WaitGroup {
state: Arc<CountdownState>,
}
impl fmt::Debug for WaitGroup {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("WaitGroup").finish_non_exhaustive()
}
}
impl Default for WaitGroup {
fn default() -> Self {
Self::new()
}
}
impl WaitGroup {
/// Creates a new `WaitGroup`.
///
/// # Examples
///
/// ```
/// use mea::waitgroup::WaitGroup;
///
/// let wg = WaitGroup::new();
/// ```
pub fn new() -> Self {
Self {
state: Arc::new(CountdownState::new(1)),
}
}
}
impl Clone for WaitGroup {
/// Creates a new worker handle for the WaitGroup.
///
/// This increments the WaitGroup counter. The counter will be decremented
/// when the new handle is dropped.
fn clone(&self) -> Self {
let sync = self.state.clone();
let mut cnt = sync.state();
loop {
let new_cnt = cnt.saturating_add(1);
match sync.cas_state(cnt, new_cnt) {
Ok(_) => return Self { state: sync },
Err(x) => cnt = x,
}
}
}
}
impl Drop for WaitGroup {
fn drop(&mut self) {
if self.state.decrement(1) {
self.state.wake_all();
}
}
}
impl IntoFuture for WaitGroup {
type Output = ();
type IntoFuture = Wait;
/// Converts the WaitGroup into a future that completes when all tasks finish. This decreases
/// the WaitGroup counter.
fn into_future(self) -> Self::IntoFuture {
let state = self.state.clone();
drop(self);
Wait { idx: None, state }
}
}
/// A future that completes when all tasks in a WaitGroup have finished.
///
/// This type is created by either: (1) calling `.await` on a `WaitGroup`, or (2) cloning
/// itself, which does not increase the WaitGroup counter, but creates a new future that
/// will complete when the WaitGroup counter reaches zero.
#[must_use = "futures do nothing unless you `.await` or poll them"]
pub struct Wait {
idx: Option<usize>,
state: Arc<CountdownState>,
}
impl Clone for Wait {
/// Creates a new future that also completes when the WaitGroup counter reaches zero.
///
/// This does not increment the WaitGroup counter.
fn clone(&self) -> Self {
Wait {
idx: None,
state: self.state.clone(),
}
}
}
impl fmt::Debug for Wait {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Wait").finish_non_exhaustive()
}
}
impl Future for Wait {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let Self { idx, state } = self.get_mut();
// register waker if the counter is not zero
if state.spin_wait(16).is_err() {
state.register_waker(idx, cx);
// double check after register waker, to catch the update between two steps
if state.spin_wait(0).is_err() {
return Poll::Pending;
}
}
Poll::Ready(())
}
}
+81
View File
@@ -0,0 +1,81 @@
// Copyright 2024 tison <wander4096@gmail.com>
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
use std::future::IntoFuture;
use std::time::Duration;
use super::*;
use crate::test_runtime;
#[test]
fn test_wait_group_drop() {
let wg = WaitGroup::new();
for _i in 0..100 {
let w = wg.clone();
test_runtime().spawn(async move {
drop(w);
});
}
pollster::block_on(wg.into_future());
}
#[test]
fn test_wait_group_await() {
let wg = WaitGroup::new();
for _i in 0..100 {
let w = wg.clone();
test_runtime().spawn(async move {
w.await;
});
}
pollster::block_on(wg.into_future());
}
#[test]
fn test_wait_group_timeout() {
let wg = WaitGroup::new();
let _wg_clone = wg.clone();
let timeout = test_runtime().block_on(async move {
tokio::select! {
_ = tokio::time::sleep(Duration::from_millis(50)) => true ,
_ = wg => false,
}
});
assert!(timeout);
}
#[test]
fn test_wait_group_cancel() {
let wg = WaitGroup::new();
let wg_clone = wg.clone().into_future();
let wg_clone_2 = wg.clone().into_future();
test_runtime().block_on(async move {
tokio::select! {
_ = tokio::time::sleep(Duration::ZERO) => {},
_ = wg_clone => {}
}
});
let fut = test_runtime().spawn(async move {
wg_clone_2.await;
});
std::thread::sleep(Duration::from_millis(50));
drop(wg);
let timeout = test_runtime().block_on(async move {
tokio::select! {
_ = tokio::time::sleep(Duration::from_secs(60)) => true ,
_ = fut => false,
}
});
assert!(!timeout);
}